"""Projection of a 2D function using kernel-based approximation spaces.""" import matplotlib.pyplot as plt import torch from scimba_torch.approximation_space.kernelx_space import ( GaussianKernel, KernelxSpace, MultiquadraticKernel, ) from scimba_torch.domain.meshless_domain.domain_2d import Square2D from scimba_torch.integration.monte_carlo import DomainSampler, TensorizedSampler from scimba_torch.integration.monte_carlo_parameters import UniformParametricSampler from scimba_torch.numerical_solvers.collocation_projector import ( CollocationProjector, LinearProjector, ) from scimba_torch.plots.plots_nd import plot_abstract_approx_spaces from scimba_torch.utils.scimba_tensors import LabelTensor def func_test(x: LabelTensor, mu: LabelTensor): x1, x2 = x.get_components() sigma = 0.4 return torch.sin(torch.cos(x1) * torch.sin(x2) / (2 * sigma**2)) torch.manual_seed(0) domain_x = Square2D([(-1.0, 1.0), (-1.0, 1.0)], is_main_domain=True) sampler = TensorizedSampler([DomainSampler(domain_x), UniformParametricSampler([])]) space1 = KernelxSpace( 1, 0, kernel_type=MultiquadraticKernel, nb_centers=1000, spatial_domain=domain_x, beta=-2, integrator=sampler, anisotropic=False, ) p = LinearProjector(space1, func_test) p.solve(n_collocation=2000) # opt_1 = { # "name": "adam", # "optimizer_args": {"lr": 4.0e-2}, # } # p = CollocationProjector(space1, func_test, optimizers=opt_1) # p.solve(epochs=2000, n_collocation=1000, verbose=True) opt_2 = { "name": "adam", "optimizer_args": {"lr": 6.0e-2}, } space2 = KernelxSpace( 1, 0, kernel_type=GaussianKernel, nb_centers=1000, spatial_domain=domain_x, beta=-2, integrator=sampler, anisotropic=False, ) p2 = CollocationProjector(space2, func_test, optimizer=opt_2) p2.solve(epochs=2000, n_collocation=2000) print("done") plot_abstract_approx_spaces( (p.space, p2.space), # the approximation spaces (domain_x), # the spatial domain ((),), # the parametric domain loss=( None, p2.losses, ), # for plot of the loss: the losses solution=(func_test), # for plot of the exact sol: sol error=(func_test), # for plot of the error with respect to a func: the func draw_contours=True, n_drawn_contours=20, n_visu=256, ) plt.show()