99import numpy as np
1010import torch
1111
12- from typing import Callable , Union , Dict , List
12+ from typing import Callable , Union , Dict , List , Tuple
1313from functools import singledispatchmethod , singledispatch
1414
1515from torch .nn import Sequential
@@ -327,24 +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 , grid_var_keys = None , * args , ** kwargs ):
330+ to_numpy : bool = False , grid_var_keys = None ,
331+ * args , ** kwargs ) -> Tuple [float , Union [torch .Tensor , np .ndarray ]]:
331332 solver_device (device = self ._device )
332333
333334 if isinstance (system , SoEq ):
334335 system_interface = SystemSolverInterface (system_to_adapt = system )
335336 system_solver_forms = system_interface .form (grids = grids , mode = mode )
336337 elif isinstance (system , dict ):
337- system_solver_forms = list (system .values ())
338+ system_solver_forms = list (system .values ()) # TODO: refactor instead of quickfixes
338339 elif isinstance (system , list ):
339340 system_solver_forms = system
340341 else :
341342 raise TypeError (f'Incorrect type of the equations passed into solver. Expected dict or SoEq, got { type (system )} .' )
342-
343+
343344 if boundary_conditions is None :
344- raise NotImplementedError ('TBD' )
345345 op_gen = PregenBOperator (system = system ,
346346 system_of_equation_solver_form = [sf_labeled [1 ] for sf_labeled
347- in system .values ()] )
347+ in system_solver_forms ]) # system.values .vals( )
348348 op_gen .generate_default_bc (vals = data , grids = grids )
349349 boundary_conditions = op_gen .conditions
350350
@@ -355,6 +355,7 @@ def solve_epde_system(self, system: Union[SoEq, dict], grids: list=None, boundar
355355
356356 if grids is None :
357357 grid_var_keys , grids = global_var .grid_cache .get_all (mode = 'torch' )
358+ grids = [grid [global_var .grid_cache .g_func != 0 ] for grid in grids ]
358359 elif grid_var_keys is None :
359360 grid_var_keys , _ = global_var .grid_cache .get_all (mode = 'torch' )
360361
@@ -367,15 +368,23 @@ def solve_epde_system(self, system: Union[SoEq, dict], grids: list=None, boundar
367368
368369 def solve (self , equations : Union [List , SoEq , SolverEquation ], domain : Domain ,
369370 boundary_conditions = None , mode = 'NN' , use_cache : bool = False ,
370- use_fourier : bool = False , fourier_params : dict = None , # epochs = 1e3,
371- use_adaptive_lambdas : bool = False , to_numpy = False , * args , ** kwargs ):
371+ use_fourier : bool = False , fourier_params : dict = None ,
372+ use_adaptive_lambdas : bool = False , to_numpy = False ,
373+ * args , ** kwargs ) -> Tuple [float , Union [torch .Tensor , np .ndarray ]]:
372374
373375 if isinstance (equations , SolverEquation ):
374376 equations_prepared = equations
375377 else :
376378 equations_prepared = SolverEquation ()
377379 for form in equations :
378- equations_prepared .add (form )
380+ print (f'form is solve has a type of { type (form )} : { form } ' )
381+ if isinstance (form , dict ):
382+ equations_prepared .add (form )
383+ elif (isinstance (form , list ) or isinstance (form , tuple )) and len (form ) == 2 :
384+ equations_prepared .add (form [1 ])
385+ else :
386+ raise ValueError ()
387+
379388 if self .net is None :
380389 self .net = self .get_net (equations_prepared , mode , domain , use_fourier ,
381390 fourier_params , device = self ._device )
0 commit comments