@@ -646,7 +646,7 @@ def create_pool(self, data: Union[np.ndarray, list, tuple], variable_names=['u',
646646 derivs = None , max_deriv_order = 1 , additional_tokens = [],
647647 data_fun_pow : int = 1 , deriv_fun_pow : int = 1 , grid : list = None ,
648648 data_nn : torch .nn .Sequential = None , fourier_layers : bool = True ,
649- fourier_params : dict = {'L' : [4 ,], 'M' : [3 ,]}):
649+ fourier_params : dict = {'L' : [4 ,], 'M' : [3 ,]}, ann_epochs_max = 1e5 ):
650650 '''
651651 Create pool of tokens to represent elementary functions, that can be included in equations.
652652
@@ -699,8 +699,8 @@ def create_pool(self, data: Union[np.ndarray, list, tuple], variable_names=['u',
699699 global_var .reset_data_repr_nn (data = data , derivs = base_derivs , train = False ,
700700 grids = grid , predefined_ann = data_nn , device = self ._device )
701701 else :
702- epochs_max = 1e4 # 1e4
703- global_var .reset_data_repr_nn (data = data , derivs = base_derivs , epochs_max = epochs_max ,
702+ # epochs_max = 1e5 # 1e4
703+ global_var .reset_data_repr_nn (data = data , derivs = base_derivs , epochs_max = ann_epochs_max ,
704704 grids = grid , predefined_ann = None , device = self ._device ,
705705 use_fourier = fourier_layers , fourier_params = fourier_params )
706706
@@ -754,7 +754,7 @@ def fit(self, data: Union[np.ndarray, list, tuple] = None, equation_terms_max_nu
754754 equation_factors_max_number = 1 , variable_names = ['u' ,], eq_sparsity_interval = (1e-4 , 2.5 ),
755755 derivs = None , max_deriv_order = 1 , additional_tokens = None , data_fun_pow : int = 1 , deriv_fun_pow : int = 1 ,
756756 optimizer : Union [SimpleOptimizer , MOEADDOptimizer ] = None , pool : TFPool = None ,
757- population : List [SoEq ] = None , data_nn = None ,
757+ population : List [SoEq ] = None , data_nn = None , ann_epochs_max = 1e5 ,
758758 fourier_layers : bool = False , fourier_params : dict = {'L' : [4 ,], 'M' : [3 ,]}):
759759 """
760760 Fit epde search algorithm to obtain differential equations, describing passed data.
@@ -827,7 +827,8 @@ def fit(self, data: Union[np.ndarray, list, tuple] = None, equation_terms_max_nu
827827 derivs = derivs , max_deriv_order = max_deriv_order ,
828828 additional_tokens = additional_tokens ,
829829 data_fun_pow = data_fun_pow , deriv_fun_pow = deriv_fun_pow ,
830- data_nn = data_nn , fourier_layers = fourier_layers , fourier_params = fourier_params )
830+ data_nn = data_nn , ann_epochs_max = ann_epochs_max ,
831+ fourier_layers = fourier_layers , fourier_params = fourier_params )
831832 else :
832833 self .pool = pool ; self .pool_params = cur_params
833834
0 commit comments