@@ -62,12 +62,9 @@ def func(x):
6262 num_boundary = adapter .num_boundary ,
6363 num_test = 500 )
6464
65- layer_size = [1 ] + adapter .net + [len (var_names )]
66- net = dde .nn .FNN (layer_size , adapter .activation , adapter .kernel_initializer )
67- model = dde .Model (data_obj , net )
68- model .compile (adapter .optimizer , lr = adapter .lr )
65+ model = adapter ._get_or_create_model (data_obj , dim = 1 , var_count = len (var_names ))
6966 try :
70- losshistory , train_state = model .train (epochs = adapter .epochs )
67+ losshistory , train_state = model .train (iterations = adapter .iterations , verbose = adapter . verbose )
7168 final_loss = float (
7269 losshistory .loss_train [- 1 ][0 ]) if losshistory .loss_train else np .nan
7370 except Exception as e :
@@ -140,12 +137,9 @@ def func(x):
140137 num_initial = adapter .num_initial ,
141138 num_test = 500 )
142139
143- layer_size = [geomtime .dim ] + adapter .net + [len (var_names )]
144- net = dde .nn .FNN (layer_size , adapter .activation , adapter .kernel_initializer )
145- model = dde .Model (data_obj , net )
146- model .compile (adapter .optimizer , lr = adapter .lr )
140+ model = adapter ._get_or_create_model (data_obj , dim = 2 , var_count = len (var_names ))
147141 try :
148- losshistory , train_state = model .train (epochs = adapter .epochs ) # <-- ИСПРАВЛЕНО
142+ losshistory , train_state = model .train (iterations = adapter .iterations , verbose = adapter . verbose ) # <-- ИСПРАВЛЕНО
149143 final_loss = float (losshistory .loss_train [- 1 ][0 ]) if losshistory .loss_train else np .nan
150144 except Exception as e :
151145 print (f"Exception: { e } " )
@@ -233,12 +227,9 @@ def func(x):
233227 num_initial = adapter .num_initial ,
234228 num_test = 500 )
235229
236- layer_size = [geomtime .dim ] + adapter .net + [len (var_names )]
237- net = dde .nn .FNN (layer_size , adapter .activation , adapter .kernel_initializer )
238- model = dde .Model (data_obj , net )
239- model .compile (adapter .optimizer , lr = adapter .lr )
230+ model = adapter ._get_or_create_model (data_obj , dim = 3 , var_count = len (var_names ))
240231 try :
241- losshistory , train_state = model .train (epochs = adapter .epochs ) # <-- ИСПРАВЛЕНО
232+ losshistory , train_state = model .train (iterations = adapter .iterations , verbose = adapter . verbose ) # <-- ИСПРАВЛЕНО
242233 final_loss = float (losshistory .loss_train [- 1 ][0 ]) if losshistory .loss_train else np .nan
243234 except Exception as e :
244235 print (f"Exception: { e } " )
@@ -263,10 +254,12 @@ def __init__(self, pretrained_net=None, **config):
263254 self .num_domain = int (self .config .get ('num_domain' , 2000 ))
264255 self .num_boundary = int (self .config .get ('num_boundary' , 500 ))
265256 self .num_initial = int (self .config .get ('num_initial' , 500 ))
266- self .epochs = int (self .config .get ('epochs' , 10000 ))
257+ #self.epochs = int(self.config.get('epochs', 10000))
258+ self .iterations = int (self .config .get ('iterations' , 10000 ))
267259 # self.iterations = int(self.config.get('epochs', 5))
268260 self .bc_type = self .config .get ('bc_type' , 'Dirichlet' )
269261 self .fallback_bc_value = self .config .get ('fallback_bc_value' , 0.0 )
262+ self .verbose = config .get ('verbose' , False )
270263
271264 self .coordinate_mapping = self .config .get ('coordinate_mapping' , None )
272265 self .coord_names = None
@@ -278,6 +271,23 @@ def __init__(self, pretrained_net=None, **config):
278271 3 : Solver3D (),
279272 }
280273
274+ self ._model = None
275+
276+ def _get_or_create_model (self , data_obj , dim , var_count ):
277+ if self ._model is None :
278+ layer_size = [dim ] + self .net + [var_count ]
279+ net = dde .nn .FNN (layer_size , self .activation , self .kernel_initializer )
280+ model = dde .Model (data_obj , net )
281+ model .compile (self .optimizer , lr = self .lr , verbose = self .verbose )
282+ self ._model = model
283+ else :
284+ def reset_weights (m ):
285+ if hasattr (m , 'reset_parameters' ):
286+ m .reset_parameters ()
287+ self ._model .net .apply (reset_weights )
288+ self ._model .data = data_obj
289+ return self ._model
290+
281291 def _set_coordinate_info (self , coord_names ):
282292 self .coord_names = coord_names
283293 if self .coordinate_mapping is not None :
0 commit comments