| 12 | ) |
| 13 | |
| 14 | def get_positional_encoding(name, hparams=None): |
| 15 | if name == "direct": |
| 16 | return PE.Direct() |
| 17 | elif name == "cartesian3d": |
| 18 | return PE.Cartesian3D() |
| 19 | elif name == "sphericalharmonics": |
| 20 | |
| 21 | # default to analytic |
| 22 | if "harmonics_calculation" not in hparams.keys(): |
| 23 | hparams["harmonics_calculation"] = "analytic" |
| 24 | |
| 25 | if "harmonics_calculation" in hparams.keys() and hparams['harmonics_calculation'] == "discretized": |
| 26 | return PE.DiscretizedSphericalHarmonics(legendre_polys=hparams['legendre_polys']) |
| 27 | else: |
| 28 | return PE.SphericalHarmonics(legendre_polys=hparams['legendre_polys'], |
| 29 | harmonics_calculation=hparams['harmonics_calculation']) |
| 30 | elif name == "theory": |
| 31 | return PE.Theory(min_radius=hparams['min_radius'], |
| 32 | max_radius=hparams['max_radius'], |
| 33 | frequency_num=hparams['frequency_num']) |
| 34 | elif name == "wrap": |
| 35 | return PE.Wrap() |
| 36 | elif name in ["grid", "spherec", "spherecplus", "spherem", "spheremplus"]: |
| 37 | return PE.GridAndSphere(min_radius=hparams['min_radius'], |
| 38 | max_radius=hparams['max_radius'], |
| 39 | frequency_num=hparams['frequency_num'], |
| 40 | name=name) |
| 41 | else: |
| 42 | raise ValueError(f"{name} not a known positional encoding.") |
| 43 | |
| 44 | def get_neural_network(name, input_dim, hparams=None): |
| 45 | if name == "linear": |