Download model/legacy/features.py from OneScience-Group/NeuralGCM: direct link, hf CLI and curl.
- Browser
- Download file 32 kB
-
https://huggingface.co/OneScience-Group/NeuralGCM/resolve/main/model/legacy/features.py
- Command line
-
hf download hf://OneScience-Group/NeuralGCM/model/legacy/features.py
-
curl -L -o features.py https://huggingface.co/OneScience-Group/NeuralGCM/resolve/main/model/legacy/features.py
32 kB
| # Copyright 2024 Google LLC | |
| # | |
| # Licensed under the Apache License, Version 2.0 (the "License"); | |
| # you may not use this file except in compliance with the License. | |
| # You may obtain a copy of the License at | |
| # | |
| # https://www.apache.org/licenses/LICENSE-2.0 | |
| # | |
| # Unless required by applicable law or agreed to in writing, software | |
| # distributed under the License is distributed on an "AS IS" BASIS, | |
| # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | |
| # See the License for the specific language governing permissions and | |
| # limitations under the License. | |
| """Modules that computes relevant state features to be used by ML components.""" | |
| from typing import Any, Callable, Mapping, Optional, Protocol, Sequence | |
| from dinosaur import coordinate_systems | |
| from dinosaur import primitive_equations | |
| from dinosaur import pytree_utils | |
| from dinosaur import radiation | |
| from dinosaur import scales | |
| from dinosaur import spherical_harmonic | |
| from dinosaur import typing | |
| from dinosaur import xarray_utils | |
| import gin | |
| import haiku as hk | |
| import jax | |
| import jax.numpy as jnp | |
| from model.legacy import transforms | |
| import numpy as np | |
| Array = typing.Array | |
| Pytree = typing.Pytree | |
| TransformModule = typing.TransformModule | |
| KeyWithCosLatFactor = typing.KeyWithCosLatFactor | |
| class FeaturesFn(Protocol): | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| ... | |
| FeaturesModule = Callable[..., FeaturesFn] | |
| class PrimitiveEquationsDiagnosticState(hk.Module): | |
| """Features modules that returns processed DiagnosticState for PE.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| self.coords = coords | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> primitive_equations.DiagnosticState: | |
| del memory, diagnostics, randomness, forcing # unused | |
| if not isinstance(inputs, primitive_equations.State): | |
| inputs = primitive_equations.State(**inputs) | |
| d_state = primitive_equations.compute_diagnostic_state(inputs, self.coords) | |
| return self.features_transform_fn(d_state.asdict()) | |
| class VelocityAndPrognostics(hk.Module): | |
| """Features module that returns prognostics + u,v and optionally gradients.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| fields_to_include: Optional[Sequence[str]] = None, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| compute_gradients_module: TransformModule = transforms.EmptyTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| self.coords = coords | |
| self.fields_to_include = fields_to_include | |
| self.compute_gradients_fn = compute_gradients_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| def _extract_features( | |
| self, | |
| inputs: typing.Pytree, | |
| prefix: str = '', | |
| ) -> typing.Pytree: | |
| """Returns a nodal velocity and prognostic features.""" | |
| # Note: all intermediate features have an explicit cos-lat factors in key. | |
| # These factors are removed in the `__call__` method before returning. | |
| # compute `u, v` if div/curl is available and `u, v` not in prognosics. | |
| if set(['vorticity', 'divergence']).issubset(inputs.keys()) and not set( | |
| ['u', 'v'] | |
| ).intersection(inputs.keys()): | |
| cos_lat_u, cos_lat_v = spherical_harmonic.get_cos_lat_vector( | |
| inputs['vorticity'], inputs['divergence'], self.coords.horizontal | |
| ) | |
| modal_features = { | |
| KeyWithCosLatFactor(prefix + 'u', 1): cos_lat_u, | |
| KeyWithCosLatFactor(prefix + 'v', 1): cos_lat_v, | |
| } | |
| else: | |
| modal_features = {} | |
| prognostics_keys = list(inputs.keys()) | |
| prognostics_keys.remove('tracers') | |
| prognostics_keys.remove('sim_time') | |
| for k in prognostics_keys: | |
| if self.fields_to_include is None or k in self.fields_to_include: | |
| modal_features[KeyWithCosLatFactor(prefix + k, 0)] = inputs[k] | |
| for k, v in inputs['tracers'].items(): | |
| if self.fields_to_include is None or k in self.fields_to_include: | |
| modal_features[KeyWithCosLatFactor(prefix + k, 0)] = v | |
| # Computing gradient features and adjusting cos_lat factors. | |
| modal_features = self.coords.with_dycore_sharding(modal_features) | |
| diff_operator_features = self.compute_gradients_fn(modal_features) | |
| sec_lat = 1 / self.coords.horizontal.cos_lat | |
| sec2_lat = self.coords.horizontal.sec2_lat | |
| sec_lat_scales = {0: 1, 1: sec_lat, 2: sec2_lat} | |
| # Computing all features in nodal space. | |
| features = {} | |
| for k, v in (diff_operator_features | modal_features).items(): | |
| sec_lat_scale = sec_lat_scales[k.factor_order] | |
| features[k.name] = self.coords.horizontal.to_nodal(v) * sec_lat_scale | |
| features = self.coords.with_dycore_sharding(features) | |
| return features | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del memory, diagnostics, randomness, forcing # unused. | |
| nodal_features = self._extract_features(inputs) | |
| return self.features_transform_fn(nodal_features) | |
| class MemoryVelocityAndValues(VelocityAndPrognostics): | |
| """Similar to `VelocityAndPrognostics`, but operates on memory.""" | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del inputs, diagnostics, randomness, forcing # unused. | |
| nodal_features = self._extract_features(memory, 'memory_') | |
| return self.features_transform_fn(nodal_features) | |
| class NodalInputVelocityAndPrognostics(VelocityAndPrognostics): | |
| """Features modules that returns velocities, temperature, and optionally gradients.""" | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| to_modal_fn = self.coords.horizontal.to_modal | |
| inputs = to_modal_fn(inputs) | |
| memory = to_modal_fn(memory) | |
| return super().__call__(inputs, memory, randomness, forcing) | |
| class RadiationFeatures(hk.Module): | |
| """Feature module that computes incident radiation flux.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| ref_datetime_str = str(aux_features[xarray_utils.REFERENCE_DATETIME_KEY]) | |
| self.solar_radiation = radiation.SolarRadiation.normalized( | |
| coords=coords, | |
| physics_specs=physics_specs, | |
| reference_datetime=np.datetime64(ref_datetime_str), | |
| ) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del memory, diagnostics, randomness, forcing # unused. | |
| features = {} | |
| features['radiation'] = self.solar_radiation.radiation_flux( | |
| inputs['sim_time'] | |
| ) | |
| # TODO(janniyuval) add a flag that allow to get radiation of next time step | |
| # insert a feature axis. | |
| features = jax.tree_util.tree_map(lambda x: jnp.expand_dims(x, 0), features) | |
| return self.features_transform_fn(features) | |
| class OrbitalTimeFeatures(hk.Module): | |
| """Feature module that computes orbital time features.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| ref_datetime_str = str(aux_features[xarray_utils.REFERENCE_DATETIME_KEY]) | |
| self.solar_radiation = radiation.SolarRadiation.normalized( | |
| coords=coords, | |
| physics_specs=physics_specs, | |
| reference_datetime=np.datetime64(ref_datetime_str), | |
| ) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del memory, diagnostics, randomness, forcing # unused. | |
| features = {} | |
| # Cosine and sine of Earth's orbital phase around the Sun | |
| orbital_time = self.solar_radiation.time_to_orbital_time(inputs['sim_time']) | |
| # Convert from orbital_phase=0 on January 1st UTC to orbital_phase=0 at the | |
| # approximate perihelion (when earth is closest to the sun). | |
| orbital_phase = orbital_time.orbital_phase - radiation.PERIHELION | |
| # All longitude, latitude locations share the same orbital phase | |
| ones = jnp.ones(self.solar_radiation.coords.surface_nodal_shape) | |
| features['cos_orbital_phase'] = jnp.cos(orbital_phase) * ones | |
| features['sin_orbital_phase'] = jnp.sin(orbital_phase) * ones | |
| # Cosine and sine of local hour angle (angle from solar noon) | |
| solar_hour_angle = self.solar_radiation.solar_hour_angle(inputs['sim_time']) | |
| solar_hour_angle = jnp.expand_dims(solar_hour_angle, 0) | |
| features['cos_solar_hour'] = jnp.cos(solar_hour_angle) | |
| features['sin_solar_hour'] = jnp.sin(solar_hour_angle) | |
| # TODO(janniyuval) add a flag that allow to get radiation of next time step | |
| return self.features_transform_fn(features) | |
| class ForcingFeatures(hk.Module): | |
| """Feature module that provides forcing values as features.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| forcing_to_include: Sequence[str] = tuple(), | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.forcing_to_include = forcing_to_include | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| def __call__( | |
| self, | |
| inputs: Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Forcing] = None, | |
| ) -> Pytree: | |
| del inputs, memory, diagnostics, randomness | |
| features = {} | |
| for key in self.forcing_to_include: | |
| value = forcing[key] # pyrefly: ignore[unsupported-operation] | |
| # Expect singleton "level" dimension for surface forcings | |
| if value.ndim > 3: | |
| raise ValueError( | |
| f'Expected forcing "{key}" to have ndim <= 3, got {value.ndim}' | |
| ) | |
| if value.ndim == 2: | |
| value = jnp.expand_dims(value, axis=0) | |
| if value.shape[0] != 1: | |
| raise ValueError( | |
| f'Expected forcing "{key}" to have leading dimension 1' | |
| f'for level, got {value.shape}' | |
| ) | |
| features[key] = value | |
| return self.features_transform_fn(features) | |
| class LatitudeFeatures(hk.Module): | |
| """Feature module that creates cos and sin of latitude as features.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| self.coords = coords | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del inputs, memory, diagnostics, randomness, forcing # unused. | |
| _, sin_lat = self.coords.horizontal.nodal_mesh | |
| sin_features = sin_lat[np.newaxis, ...] | |
| cos_features = jnp.cos(jnp.arcsin(sin_features)) | |
| features = { | |
| 'cos_latitude': cos_features, | |
| 'sin_latitude': sin_features, | |
| } | |
| return self.features_transform_fn(features) | |
| class RandomnessFeatures(hk.Module): | |
| """Feature module that returns fields from `randomness` as features.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del inputs, memory, diagnostics, forcing # unused. | |
| if randomness is None: | |
| random_features = {} | |
| elif isinstance(randomness, dict): | |
| random_features, _ = pytree_utils.flatten_dict(randomness) | |
| elif isinstance(randomness, jax.Array): | |
| random_features = {'randomness': randomness} | |
| else: | |
| raise ValueError(f'randomness has unsupported {type(randomness)=}.') | |
| # random fields are 2D by construction, adding a feature/level dimension. | |
| if randomness is not None: | |
| ndims = set(x.ndim for x in jax.tree_util.tree_leaves(random_features)) | |
| if not ndims.issubset({2, 3}): | |
| raise ValueError( | |
| f'Random fields expected to be 2D and/or 3D. Found {ndims=}' | |
| ) | |
| def make_3d(x): | |
| if x.ndim == 3: | |
| return x | |
| if x.ndim == 2: | |
| return x[np.newaxis, ...] | |
| random_features = jax.tree_util.tree_map(make_3d, random_features) | |
| return self.features_transform_fn(random_features) | |
| class OrographyFeatures(hk.Module): | |
| """Feature module that computes orographic features.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| if xarray_utils.OROGRAPHY not in aux_features: | |
| raise ValueError('OrographyFeatures requires orography in aux_features.') | |
| self.nodal_orography = aux_features[xarray_utils.OROGRAPHY] | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del inputs, memory, diagnostics, randomness, forcing # unused. | |
| features = { | |
| xarray_utils.OROGRAPHY: jnp.expand_dims(self.nodal_orography, 0), | |
| } | |
| return self.features_transform_fn(features) | |
| class OneHotAuxFeatures(hk.Module): | |
| """Feature module that produces one-hot encodings from binary covariates.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| covariate_keys: Sequence[str] = gin.REQUIRED, # pyrefly: ignore[bad-function-definition] | |
| convert_float_to_int: bool = False, | |
| name: Optional[str] = None, | |
| ): | |
| del coords, dt, physics_specs # unused. | |
| super().__init__(name=name) | |
| covariates = {} | |
| num_classes = {} | |
| for key in covariate_keys: | |
| if key not in aux_features: | |
| raise ValueError(f'Covariate {key} not found in aux_features.') | |
| if not np.issubdtype(aux_features[key].dtype, np.integer): | |
| if convert_float_to_int: | |
| aux_features[key] = np.round(aux_features[key]).astype(int) | |
| else: | |
| raise ValueError( | |
| f'Covariate {key} is expected to be integer dtype, ' | |
| f'but is: {aux_features[key].dtype}' | |
| ) | |
| covariates[key] = aux_features[key] | |
| num_classes[key] = np.unique(aux_features[key]).size | |
| self.covariates = covariates | |
| self.num_classes = num_classes | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> dict[str, jnp.ndarray]: | |
| del inputs, memory, diagnostics, randomness, forcing # unused. | |
| features = { | |
| k: jax.nn.one_hot(v, self.num_classes[k], axis=0) | |
| for k, v in self.covariates.items() | |
| } | |
| return features | |
| class LearnedPositionalFeatures(hk.Module): | |
| """Feature module with learned params at surface nodal locations.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| latent_size: int, | |
| scale: float = 1.0, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.scale = scale | |
| self.padding = coords.horizontal.nodal_padding | |
| unpadded_nodal_shape = tuple( | |
| x - y for x, y in zip(coords.horizontal.nodal_shape, self.padding) | |
| ) | |
| self.positional_features = hk.get_parameter( | |
| 'learned_positional_features', | |
| (latent_size,) + unpadded_nodal_shape, | |
| jnp.float32, | |
| init=hk.initializers.Constant(0.0), | |
| ) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> dict[str, jnp.ndarray]: | |
| """Returns scaled parameter values at surface nodal locations.""" | |
| del inputs, memory, diagnostics, randomness, forcing # unused. | |
| pad_x, pad_y = self.padding | |
| positional_features = self.scale * jnp.pad( | |
| self.positional_features, [(0, 0), (0, pad_x), (0, pad_y)] | |
| ) | |
| return {'learned_positional_features': positional_features} | |
| class EmbeddingSurfaceFeatures(hk.Module): | |
| """Feature module that specifies embedding surface outputs as features. | |
| Returns {feature_name: nn_output} | |
| where nn_output.shape = (output_size, lon, lat). | |
| """ | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| feature_name: str, | |
| output_size: int, | |
| embedding_module: typing.EmbeddingModule, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| # output shapes are arrays to be pytree leaves for tree_map | |
| output_shapes = { | |
| feature_name: np.asarray((output_size,) + coords.horizontal.nodal_shape) | |
| } | |
| self.embedding_fn = embedding_module( | |
| coords, dt, physics_specs, aux_features, output_shapes=output_shapes | |
| ) | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| features = self.embedding_fn( | |
| inputs, memory, diagnostics, randomness, forcing # pyrefly: ignore[bad-argument-type] | |
| ) | |
| return self.features_transform_fn(features) | |
| class EmbeddingVolumeFeatures(hk.Module): | |
| """Feature module that specifies embedding volume outputs as features. | |
| Returns {feature_name_0: nn_output_0, | |
| feature_name_1: nn_output_1, | |
| ... | |
| } | |
| where the NN output array has shape (output_size, level, lon, lat), which is | |
| unpacked over output_size such that nn_output_{i}.shape = (level, lon, lat) | |
| for each i in range(output_size). | |
| """ | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| feature_name: str, | |
| output_size: int, | |
| embedding_module: typing.EmbeddingModule, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| # output shapes are arrays to be pytree leaves for tree_map | |
| output_shapes = { | |
| f'{feature_name}_{i}': np.asarray(coords.nodal_shape) | |
| for i in range(output_size) | |
| } | |
| self.embedding_fn = embedding_module( | |
| coords, dt, physics_specs, aux_features, output_shapes=output_shapes | |
| ) | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| features = self.embedding_fn( | |
| inputs, memory, diagnostics, randomness, forcing # pyrefly: ignore[bad-argument-type] | |
| ) | |
| return self.features_transform_fn(features) | |
| class FloatDataFeatures(hk.Module): | |
| """Feature module that supplies floating point covariates from data.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| covariate_data_path: str = gin.REQUIRED, # pyrefly: ignore[bad-function-definition] | |
| covariate_keys: Sequence[str] = gin.REQUIRED, # pyrefly: ignore[bad-function-definition] | |
| renaming_dict: Optional[Mapping[str, str]] = None, | |
| compute_gradients_module: TransformModule = transforms.EmptyTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.covariates = {} | |
| self.compute_gradients_fn = compute_gradients_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| self.coords = coords | |
| ds = xarray_utils.ds_from_path_or_aux(covariate_data_path, aux_features) | |
| if renaming_dict is not None: | |
| ds = ds.rename(renaming_dict) | |
| lon, lat = (ds[xarray_utils.XR_LON_NAME], ds[xarray_utils.XR_LAT_NAME]) | |
| xarray_utils.verify_grid_consistency(lon, lat, coords.horizontal) | |
| lon_lat_order = (xarray_utils.XR_LON_NAME, xarray_utils.XR_LAT_NAME) | |
| for key in covariate_keys: | |
| data = ds[key].transpose(*lon_lat_order) | |
| data_units = scales.parse_units(data.attrs['units']) | |
| data = physics_specs.nondimensionalize(data.values * data_units) | |
| if data.ndim != 3: | |
| data = data[np.newaxis, ...] | |
| self.covariates[key] = data | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> dict[str, jnp.ndarray]: | |
| del inputs, memory, diagnostics, forcing, randomness # unused. | |
| features = {k: v for k, v in self.covariates.items()} | |
| modal_features = self.coords.horizontal.to_modal(features) | |
| modal_features = { # jit should eliminate to_modal if it is not used. | |
| KeyWithCosLatFactor(k, 0): v for k, v in modal_features.items() | |
| } | |
| modal_gradient_features = self.compute_gradients_fn(modal_features) | |
| sec_lat = 1 / self.coords.horizontal.cos_lat | |
| sec2_lat = self.coords.horizontal.sec2_lat | |
| sec_lat_scales = {0: 1, 1: sec_lat, 2: sec2_lat} | |
| for k, v in modal_gradient_features.items(): | |
| sec_lat_scale = sec_lat_scales[k.factor_order] | |
| features[k.name] = self.coords.horizontal.to_nodal(v) * sec_lat_scale | |
| return features | |
| class CombinedFeatures(hk.Module): | |
| """Feature module that combines multiple feature modules together.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| feature_modules: Sequence[FeaturesModule] = gin.REQUIRED, # pyrefly: ignore[bad-function-definition] | |
| feature_module_names_to_exclude: Sequence[str] = tuple(), | |
| features_to_exclude: Sequence[str] = tuple(), | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.feature_fns = [ | |
| module(coords, dt, physics_specs, aux_features) | |
| for module in feature_modules | |
| ] | |
| self.feature_module_names_to_exclude = feature_module_names_to_exclude | |
| self.features_to_exclude = features_to_exclude | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Forcing] = None, | |
| ) -> dict[str, jnp.ndarray]: | |
| all_features = {} | |
| for feature_fn in self.feature_fns: | |
| if type(feature_fn).__name__ not in self.feature_module_names_to_exclude: | |
| features = feature_fn(inputs, memory, diagnostics, randomness, forcing) | |
| for k, v in features.items(): | |
| if k in all_features: | |
| raise ValueError(f'Encountered duplicate feature {k}') | |
| all_features[k] = v | |
| all_features = self.features_transform_fn(all_features) | |
| for k in self.features_to_exclude: | |
| all_features.pop(k, None) | |
| return all_features | |
| class NullFeatures(hk.Module): | |
| """Placeholder features module that returns an empty dict.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| name: Optional[str] = None, | |
| ): | |
| del coords, dt, physics_specs, aux_features # unused | |
| super().__init__(name=name) | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> dict[str, jnp.ndarray]: | |
| del inputs, memory, diagnostics, randomness, forcing # unused | |
| return {} | |
| class PressureFeatures(hk.Module): | |
| """Feature module that computes pressure.""" | |
| def __init__( | |
| self, | |
| coords: coordinate_systems.CoordinateSystem, | |
| dt: float, | |
| physics_specs: Any, | |
| aux_features: typing.AuxFeatures, | |
| features_transform_module: TransformModule = transforms.IdentityTransform, | |
| name: Optional[str] = None, | |
| ): | |
| super().__init__(name=name) | |
| self.coords = coords | |
| self.features_transform_fn = features_transform_module( | |
| coords, dt, physics_specs, aux_features | |
| ) | |
| def _nodal_pressure( | |
| self, | |
| inputs: typing.Pytree, | |
| prefix: str = '', | |
| ) -> Mapping[str, Array]: | |
| """Computes nodal pressure from model inputs.""" | |
| # Compute nodal, dimensionalized quantities | |
| to_nodal_fn = self.coords.horizontal.to_nodal | |
| sigma = self.coords.vertical.centers | |
| surface_pressure = jnp.exp(to_nodal_fn(inputs['log_surface_pressure'])) | |
| pressure = surface_pressure * sigma[:, jnp.newaxis, jnp.newaxis] | |
| nodal_features = {prefix + 'pressure': pressure} | |
| return nodal_features | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del memory, diagnostics, randomness, forcing # unused. | |
| nodal_features = self._nodal_pressure(inputs) | |
| return self.features_transform_fn(nodal_features) | |
| class MemoryPressureFeatures(PressureFeatures): | |
| """Feature module that computes pressure from memory values.""" | |
| def __call__( | |
| self, | |
| inputs: typing.Pytree, | |
| memory: Optional[typing.PyTreeState] = None, | |
| diagnostics: Optional[typing.Pytree] = None, | |
| randomness: Optional[typing.PyTreeState] = None, | |
| forcing: Optional[typing.Pytree] = None, | |
| ) -> typing.Pytree: | |
| del inputs, diagnostics, randomness, forcing # unused. | |
| nodal_features = self._nodal_pressure(memory, 'memory_') | |
| return self.features_transform_fn(nodal_features) | |