Skip to content

Commit 25c9e35

Browse files
committed
Numerical solver fix
1 parent 3fd8d02 commit 25c9e35

2 files changed

Lines changed: 13 additions & 9 deletions

File tree

epde/integrate/numeric_integration.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -51,8 +51,8 @@ def transform_term(term: Dict, deriv_key: list, var: int) -> Dict:
5151
term_filtered['term'] = [None,]
5252
term_filtered['pow'] = 0
5353
else:
54-
term_idx = [der_var for idx, der_var in enumerate(term_filtered['term'])
55-
if der_var == deriv_key and term_filtered['pow'][idx] == var][0]
54+
term_idx = [idx for idx, der_var in enumerate(term_filtered['term'])
55+
if der_var == deriv_key and term_filtered['var'][idx] == var][0]
5656
term_filtered['term'][term_idx] = [None,]
5757
term_filtered['pow'][term_idx] = 0
5858
return term_filtered

epde/integrate/pinn_integration.py

Lines changed: 11 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -120,17 +120,17 @@
120120

121121

122122
BASE_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

Comments
 (0)