120120
121121
122122BASE_TRAINING_PARAMS = {
123- 'epochs' : 3e4 , # 1e5
123+ 'epochs' : 1e2 , # 1e5
124124 'info_string_every' : 'None' , #1e4,
125125 'mixed_precision' : False ,
126126 'save_model' : False ,
127127 'model_name' : 'None'
128128 }
129129
130- def solver_formed_grid (training_grid = None , device = 'cpu' ):
130+ def solver_formed_grid (training_grid = None , grid_var_keys = None , device = 'cpu' ):
131131 if training_grid is None :
132132 keys , training_grid = global_var .grid_cache .get_all (mode = 'torch' )
133- else :
133+ elif grid_var_keys is None :
134134 keys , _ = global_var .grid_cache .get_all (mode = 'torch' )
135135
136136 assert len (keys ) == training_grid [0 ].ndim , 'Mismatching dimensionalities'
@@ -327,21 +327,24 @@ def create_domain(variables: List[str], grids : List[Union[np.ndarray, torch.Ten
327327 def solve_epde_system (self , system : Union [SoEq , dict ], grids : list = None , boundary_conditions = None ,
328328 mode = 'NN' , data = None , use_cache : bool = False , use_fourier : bool = False ,
329329 fourier_params : dict = None , use_adaptive_lambdas : bool = False ,
330- to_numpy : bool = False , * args , ** kwargs ):
330+ to_numpy : bool = False , grid_var_keys = None , * args , ** kwargs ):
331331 solver_device (device = self ._device )
332332
333333 if isinstance (system , SoEq ):
334334 system_interface = SystemSolverInterface (system_to_adapt = system )
335335 system_solver_forms = system_interface .form (grids = grids , mode = mode )
336336 elif isinstance (system , dict ):
337+ system_solver_forms = list (system .values ())
338+ elif isinstance (system , list ):
337339 system_solver_forms = system
338340 else :
339341 raise TypeError (f'Incorrect type of the equations passed into solver. Expected dict or SoEq, got { type (system )} .' )
340342
341343 if boundary_conditions is None :
344+ raise NotImplementedError ('TBD' )
342345 op_gen = PregenBOperator (system = system ,
343346 system_of_equation_solver_form = [sf_labeled [1 ] for sf_labeled
344- in system_solver_forms ])
347+ in system . values () ])
345348 op_gen .generate_default_bc (vals = data , grids = grids )
346349 boundary_conditions = op_gen .conditions
347350
@@ -352,11 +355,12 @@ def solve_epde_system(self, system: Union[SoEq, dict], grids: list=None, boundar
352355
353356 if grids is None :
354357 grid_var_keys , grids = global_var .grid_cache .get_all (mode = 'torch' )
355- else :
358+ elif grid_var_keys is None :
356359 grid_var_keys , _ = global_var .grid_cache .get_all (mode = 'torch' )
360+
357361 domain = self .create_domain (grid_var_keys , grids , self ._device )
358362
359- return self .solve (equations = [ form [ 1 ] for form in system_solver_forms ] , domain = domain ,
363+ return self .solve (equations = system_solver_forms , domain = domain ,
360364 boundary_conditions = bconds_combined , mode = mode , use_cache = use_cache ,
361365 use_fourier = use_fourier , fourier_params = fourier_params ,
362366 use_adaptive_lambdas = use_adaptive_lambdas , to_numpy = to_numpy )
0 commit comments