diff --git a/dinosaur/held_suarez.py b/dinosaur/held_suarez.py index 6281a80..7901f0b 100644 --- a/dinosaur/held_suarez.py +++ b/dinosaur/held_suarez.py @@ -236,7 +236,7 @@ def explicit_terms( nodal_surface_pressure = jnp.exp(nodal_log_surface_pressure) # Pressure at layer centers, with shape (levels, latitude, longitude) - pressure = self.nondim_coords.vertical.pressure_centers( + pressure = self.nondim_coords.vertical.pressure_centers( # pyrefly: ignore[missing-attribute] nodal_surface_pressure ) diff --git a/dinosaur/primitive_equations.py b/dinosaur/primitive_equations.py index 332285a..1c10c55 100644 --- a/dinosaur/primitive_equations.py +++ b/dinosaur/primitive_equations.py @@ -62,12 +62,14 @@ class State: """Records the state of a system described by the primitive equations.""" - vorticity: Array - divergence: Array - temperature_variation: Array - log_surface_pressure: Array - tracers: Mapping[str, Array] = dataclasses.field(default_factory=dict) - sim_time: float | None = None + vorticity: Array | jax.sharding.PartitionSpec + divergence: Array | jax.sharding.PartitionSpec + temperature_variation: Array | jax.sharding.PartitionSpec + log_surface_pressure: Array | jax.sharding.PartitionSpec + tracers: Mapping[str, Array | jax.sharding.PartitionSpec] = dataclasses.field( + default_factory=dict + ) + sim_time: float | Array | jax.sharding.PartitionSpec | None = None def _asdict(state: State) -> dict[str, Any]: @@ -1567,7 +1569,7 @@ def to_nodal_fn(x): # Hybrid vertical velocity / mass flux calculation. nodal_surface_pressure = jnp.exp(to_nodal_fn(state.log_surface_pressure)) delta_p = coords.vertical.layer_thickness(nodal_surface_pressure) # pyrefly: ignore[missing-attribute] - delta_b = coords.vertical.sigma_thickness[:, np.newaxis, np.newaxis] + delta_b = coords.vertical.sigma_thickness[:, np.newaxis, np.newaxis] # pyrefly: ignore[missing-attribute] # D_k = div(v * dp) = dp * div(v) + v . grad(dp) # grad(dp) = grad(da + db * ps) = db * ps * grad(ln ps) @@ -1582,7 +1584,7 @@ def compute_mass_flux(d_k): cumsum_d = jax_numpy_utils.cumsum(d_k, axis=0) # pad top with 0 cumsum_d_padded = jnp.pad(cumsum_d, ((1, 0), (0, 0), (0, 0))) - b_boundaries = coords.vertical.b_boundaries[:, np.newaxis, np.newaxis] + b_boundaries = coords.vertical.b_boundaries[:, np.newaxis, np.newaxis] # pyrefly: ignore[missing-attribute] return -cumsum_d_padded + b_boundaries * sum_d mass_flux_full = compute_mass_flux(d_k_full) @@ -1985,11 +1987,11 @@ def __post_init__(self): raise ValueError('`reference_surface_pressure` must be positive.') self._nondim_reference_surface_pressure = nondim_reference_surface_pressure nondim_a_boundaries = self.physics_specs.nondimensionalize( - self.coords.vertical.a_boundaries * self.hpa_quantity + self.coords.vertical.a_boundaries * self.hpa_quantity # pyrefly: ignore[missing-attribute] ) self.nondim_levels = hybrid_coordinates.HybridCoordinates( a_boundaries=nondim_a_boundaries, # pyrefly: ignore[bad-argument-type] - b_boundaries=self.coords.vertical.b_boundaries, + b_boundaries=self.coords.vertical.b_boundaries, # pyrefly: ignore[missing-attribute] ) nondim_coords = dataclasses.replace( self.coords, diff --git a/dinosaur/typing.py b/dinosaur/typing.py index 9ef93df..49b51eb 100644 --- a/dinosaur/typing.py +++ b/dinosaur/typing.py @@ -68,7 +68,7 @@ class RandomnessState: nodal_value: Pytree | None = None modal_value: Pytree | None = None prng_key: PRNGKeyArray | None = None - prng_step: int | None = None + prng_step: int | Array | None = None @tree_math.struct diff --git a/dinosaur/weatherbench_utils.py b/dinosaur/weatherbench_utils.py index 8c295f1..1880d34 100644 --- a/dinosaur/weatherbench_utils.py +++ b/dinosaur/weatherbench_utils.py @@ -17,6 +17,7 @@ import dataclasses from dinosaur import typing +import jax import tree_math @@ -27,6 +28,6 @@ class State: v: typing.Array t: typing.Array z: typing.Array - sim_time: float + sim_time: float | typing.Array | jax.sharding.PartitionSpec | None tracers: dict[str, typing.Array] = dataclasses.field(default_factory=dict) diagnostics: dict[str, typing.Array] = dataclasses.field(default_factory=dict)