1313import matplotlib .pyplot as plt
1414from matplotlib import cm
1515
16- from epde .integrate import SolverAdapter , DeepXDEAdapter
16+ from epde .integrate import SolverAdapter
1717from epde .structure .main_structures import SoEq , Equation
1818from epde .operators .utils .template import CompoundOperator
1919import epde .globals as global_var
20- from sklearn .linear_model import LinearRegression , Ridge
21- from scipy .optimize import minimize
22- from epde .supplementary import minmax_normalize
2320from epde .supplementary import calculate_weights
2421
2522LOSS_NAN_VAL = 1e7
@@ -391,13 +388,18 @@ class DeepXDEBasedFitness(CompoundOperator):
391388
392389 def __init__ (self , param_keys : list ):
393390 super ().__init__ (param_keys )
394- self .adapter = None
395-
396- def set_adapter (self , config : dict = None , pretrained_net = None ):
397- if self .adapter is None :
398- from epde .integrate .deepxde_integration import DeepXDEAdapter
399- cfg = self .params .get ('deepxde_config' , {}) if config is None else config
400- self .adapter = DeepXDEAdapter (pretrained_net = pretrained_net , ** cfg )
391+ self ._solver = None
392+
393+ def _get_solver (self ):
394+ if self ._solver is None :
395+ from epde .solver .factory import SolverFactory
396+ solver_type = self .params .get ('solver_type' , 'deepxde' )
397+ solver_config = self .params .get ('solver_config' , {})
398+ # Обратная совместимость со старым параметром deepxde_config
399+ if 'deepxde_config' in self .params and solver_type == 'deepxde' :
400+ solver_config = self .params ['deepxde_config' ]
401+ self ._solver = SolverFactory .create (solver_type , ** solver_config )
402+ return self ._solver
401403
402404 def apply (self , objective , arguments : dict , force_out_of_place : bool = False ):
403405 self_args , subop_args = self .parse_suboperator_args (arguments = arguments )
@@ -406,13 +408,7 @@ def apply(self, objective, arguments: dict, force_out_of_place: bool = False):
406408 self .suboperators ['sparsity' ].apply (objective , subop_args .get ('sparsity' , {}))
407409 self .suboperators ['coeff_calc' ].apply (objective , subop_args .get ('coeff_calc' , {}))
408410
409- try :
410- pretrained_net = deepcopy (global_var .solution_guess_nn )
411- except :
412- pretrained_net = None
413- self .set_adapter (pretrained_net = pretrained_net )
414-
415- keys , grids = global_var .grid_cache .get_all (mode = 'numpy' )
411+ keys , grids = global_var .grid_cache .get_all (mode = 'numpy' , structural = True )
416412
417413 if isinstance (objective , SoEq ):
418414 data_list = []
@@ -424,15 +420,26 @@ def apply(self, objective, arguments: dict, force_out_of_place: bool = False):
424420 _ , target , _ = objective .evaluate (normalize = False , return_val = False )
425421 data_list = [target .reshape (- 1 )]
426422
423+ solver = self ._get_solver ()
424+
425+ print (f"[DEBUG] objective type: { type (objective )} " )
426+ if isinstance (objective , SoEq ):
427+ print (f"[DEBUG] vars_to_describe: { objective .vars_to_describe } " )
428+ print (f"[DEBUG] data_list length: { len (data_list )} " )
429+ for i , d in enumerate (data_list ):
430+ print (f"[DEBUG] data_list[{ i } ].shape: { d .shape } " )
431+
427432 try :
428- solution_list , loss = self . adapter .solve (equation_or_system = objective ,
429- grids = grids ,
430- data = data_list )
433+ solution_list , loss = solver .solve (equation_or_system = objective ,
434+ grids = grids ,
435+ data = data_list )
431436 if np .isnan (loss ):
432437 raise ValueError ("NaN loss" )
433438
434439 if isinstance (objective , SoEq ):
435- for idx , (var_name , eq ) in enumerate ({val : objective .vals [val ] for val in objective .vars_to_describe }.items ()):
440+ # Перебираем уравнения в порядке vars_to_describe
441+ for idx , var_name in enumerate (objective .vars_to_describe ):
442+ eq = objective .vals [var_name ]
436443 err = self ._compute_error (solution_list [idx ], data_list [idx ], eq )
437444 if force_out_of_place :
438445 pass
@@ -451,7 +458,9 @@ def apply(self, objective, arguments: dict, force_out_of_place: bool = False):
451458 objective .fitness_calculated = True
452459 self ._compute_stability_for_equation (objective )
453460 except Exception as e :
454- print (f'[DeepXDEBasedFitness] DeepXDE solve failed: { e } ' )
461+ print (f'[DeepXDEBasedFitness] Solver failed: { e } ' )
462+ import traceback
463+ traceback .print_exc ()
455464 fitness_value = 1e7
456465 if force_out_of_place :
457466 return fitness_value
@@ -483,7 +492,6 @@ def _compute_error(self, solution, data, eq):
483492 return err
484493
485494 def _compute_stability_for_equation (self , eq : Equation ):
486- # Повторно вычисляется evaluate
487495 _ , target , features = eq .evaluate (normalize = False , return_val = False )
488496 data_shape = global_var .grid_cache .inner_shape
489497 self .get_g_fun_vals ()
0 commit comments