1414from matplotlib import cm
1515
1616from epde .integrate import SolverAdapter
17- # DeepXDEAdapter is imported lazily inside DeepXDEBasedFitness.apply() to
18- # avoid triggering deepxde's import-time backend banner when no DeepXDE
19- # solver is used (e.g. legacy L2/L2LR fitness paths).
17+
2018from epde .structure .main_structures import SoEq , Equation
2119from epde .operators .utils .template import CompoundOperator
2220import epde .globals as global_var
23- from sklearn .linear_model import LinearRegression , Ridge
24- from scipy .optimize import minimize
25- from epde .supplementary import minmax_normalize
2621from epde .supplementary import calculate_weights
2722
2823LOSS_NAN_VAL = 1e7
@@ -474,13 +469,18 @@ class DeepXDEBasedFitness(CompoundOperator):
474469
475470 def __init__ (self , param_keys : list ):
476471 super ().__init__ (param_keys )
477- self .adapter = None
478-
479- def set_adapter (self , config : dict = None , pretrained_net = None ):
480- if self .adapter is None :
481- from epde .integrate .deepxde_integration import DeepXDEAdapter
482- cfg = self .params .get ('deepxde_config' , {}) if config is None else config
483- self .adapter = DeepXDEAdapter (pretrained_net = pretrained_net , ** cfg )
472+ self ._solver = None
473+
474+ def _get_solver (self ):
475+ if self ._solver is None :
476+ from epde .solver .factory import SolverFactory
477+ solver_type = self .params .get ('solver_type' , 'deepxde' )
478+ solver_config = self .params .get ('solver_config' , {})
479+ # Обратная совместимость со старым параметром deepxde_config
480+ if 'deepxde_config' in self .params and solver_type == 'deepxde' :
481+ solver_config = self .params ['deepxde_config' ]
482+ self ._solver = SolverFactory .create (solver_type , ** solver_config )
483+ return self ._solver
484484
485485 def apply (self , objective , arguments : dict , force_out_of_place : bool = False ):
486486 self_args , subop_args = self .parse_suboperator_args (arguments = arguments )
@@ -489,13 +489,7 @@ def apply(self, objective, arguments: dict, force_out_of_place: bool = False):
489489 self .suboperators ['sparsity' ].apply (objective , subop_args .get ('sparsity' , {}))
490490 self .suboperators ['coeff_calc' ].apply (objective , subop_args .get ('coeff_calc' , {}))
491491
492- try :
493- pretrained_net = deepcopy (global_var .solution_guess_nn )
494- except :
495- pretrained_net = None
496- self .set_adapter (pretrained_net = pretrained_net )
497-
498- keys , grids = global_var .grid_cache .get_all (mode = 'numpy' )
492+ keys , grids = global_var .grid_cache .get_all (mode = 'numpy' , structural = True )
499493
500494 if isinstance (objective , SoEq ):
501495 data_list = []
@@ -507,15 +501,26 @@ def apply(self, objective, arguments: dict, force_out_of_place: bool = False):
507501 _ , target , _ = objective .evaluate (normalize = False , return_val = False )
508502 data_list = [target .reshape (- 1 )]
509503
504+ solver = self ._get_solver ()
505+
506+ print (f"[DEBUG] objective type: { type (objective )} " )
507+ if isinstance (objective , SoEq ):
508+ print (f"[DEBUG] vars_to_describe: { objective .vars_to_describe } " )
509+ print (f"[DEBUG] data_list length: { len (data_list )} " )
510+ for i , d in enumerate (data_list ):
511+ print (f"[DEBUG] data_list[{ i } ].shape: { d .shape } " )
512+
510513 try :
511- solution_list , loss = self . adapter .solve (equation_or_system = objective ,
512- grids = grids ,
513- data = data_list )
514+ solution_list , loss = solver .solve (equation_or_system = objective ,
515+ grids = grids ,
516+ data = data_list )
514517 if np .isnan (loss ):
515518 raise ValueError ("NaN loss" )
516519
517520 if isinstance (objective , SoEq ):
518- for idx , (var_name , eq ) in enumerate ({val : objective .vals [val ] for val in objective .vars_to_describe }.items ()):
521+ # Перебираем уравнения в порядке vars_to_describe
522+ for idx , var_name in enumerate (objective .vars_to_describe ):
523+ eq = objective .vals [var_name ]
519524 err = self ._compute_error (solution_list [idx ], data_list [idx ], eq )
520525 if force_out_of_place :
521526 pass
@@ -534,7 +539,9 @@ def apply(self, objective, arguments: dict, force_out_of_place: bool = False):
534539 objective .fitness_calculated = True
535540 self ._compute_stability_for_equation (objective )
536541 except Exception as e :
537- print (f'[DeepXDEBasedFitness] DeepXDE solve failed: { e } ' )
542+ print (f'[DeepXDEBasedFitness] Solver failed: { e } ' )
543+ import traceback
544+ traceback .print_exc ()
538545 fitness_value = 1e7
539546 if force_out_of_place :
540547 return fitness_value
@@ -566,7 +573,6 @@ def _compute_error(self, solution, data, eq):
566573 return err
567574
568575 def _compute_stability_for_equation (self , eq : Equation ):
569- # Повторно вычисляется evaluate
570576 _ , target , features = eq .evaluate (normalize = False , return_val = False )
571577 data_shape = global_var .grid_cache .inner_shape
572578 self .get_g_fun_vals ()
0 commit comments