Skip to content

Commit 6d237d7

Browse files
sadamovclaude
andcommitted
feat: adopt #652 multi-datastore schema + ForecastBatch return shape
Pre-emptively land the public API shape proposed in #652 on top of the boundary-datastore work in #635, so the model-side adapter (Joel's follow-up) doesn't have to break the schema again later. Config schema (neural_lam/config.py): - Replace `datastore` + `datastore_boundary` top-level keys with a single `datastores: Dict[str, DatastoreSelection]` mapping. The dict key becomes the canonical source name used throughout the pipeline (ForecastBatch field keys, weight / clamping disambiguation, etc). - DatastoreSelection grows optional `inputs:` and `outputs:` per- category variable include-lists, with `None` per-category values meaning "all variables in that category". `outputs:` declares which datastore is the interior (prognostic source); omitted means input- only. - Validate at config load: raise InvalidConfigError when two datastores declare the same variable as an output, pointing the user at mdp's `dim_mapping.name_format` (new builds) or `xr.Dataset.assign_coords` on the existing zarr's small `{category}_feature` coord (milliseconds regardless of zarr size). - `load_config_and_datastore` returns `(config, Dict[str, BaseDatastore])` - the full multi-source mapping - rather than the legacy `(config, interior, boundary)` triple. No transitional adapter. Dataset return shape (neural_lam/weather_dataset.py): - New `ForecastBatch` NamedTuple with per-source dict fields (`init_states`, `target_states`, `forcing`) plus a global `target_times` tensor. PyTorch's default_collate recurses through the NamedTuple and the dicts, stacking per-source tensors along a new batch axis - no custom collate_fn required. - `WeatherDataset.__init__` takes `(datastores, selections, ...)` dicts directly. Internally it still resolves the single interior + optional boundary pair for its slicing/windowing logic - that internal multi-source rewrite is part of #652's model-side follow-up. The PUBLIC return type is already shaped for the multi-source case so the future change is additive on the producer side. - `WeatherDataModule` signature updated to match. - `create_dataarray_from_tensor` refactored to expose a `build_dataarray_from_tensor` staticmethod for callers (e.g. `ForecasterModule._create_dataarray_from_tensor`) that have a datastore but not a full dataset. Production call site (neural_lam/train_model.py): - Unpacks the new `(config, datastores)` return shape. - Uses `_resolve_datastore_roles` to pick out the interior + boundary for the legacy ForecasterModule constructor (which keeps its single-datastore + boundary shape pending Joel's adapter). - WeatherDataModule receives the full multi-source dicts. Example YAMLs (tests/datastore_examples/): - Single-source danra: top-level `datastores: {danra: ...}` wrapping. - danra + era5 boundary: `datastores: {interior: ..., boundary: ...}` with explicit `outputs: {state: }` on interior so the resolver knows which one is the prognostic source. Intentionally NOT in this PR (per #652 follow-up scope): - Model-side adapter. `ForecasterModule.training_step` and friends still unpack the legacy 5-tuple `(init_states, target_states, forcing, boundary, target_times)`, which will fail at runtime when Lightning hands them a `ForecastBatch`. Joel's follow-up replaces those positional unpacks with `batch.init_states["interior"]`-style per-source dict access. - Test updates. Every test that builds a `WeatherDataset` or `NeuralLAMConfig` with the old keyword signature will fail. They need mechanical updates to the new dict shape (and the model- exercising ones should stay skipped until the model adapter lands). - Variable include-lists honoured at runtime. The schema parses `inputs:` / `outputs:` but `WeatherDataset` still fetches all variables per category from each datastore. Filtering by the declared subsets is a small follow-up that depends on whether Joel's adapter wants to do it at dataset level or model level. - Diagnostic outputs. The schema accepts `outputs.diagnostic: [...]` but nothing concatenates them into the target tensor yet; same story (small follow-up after the model adapter decides on the prognostic/diagnostic split semantic). Refs #635, refs #652. Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
1 parent 78a4d52 commit 6d237d7

6 files changed

Lines changed: 399 additions & 171 deletions

File tree

‎neural_lam/config.py‎

Lines changed: 112 additions & 49 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
# Standard library
22
import dataclasses
33
from pathlib import Path
4-
from typing import Dict, Union
4+
from typing import Dict, List, Optional, Union
55

66
# Third-party
77
import dataclass_wizard
@@ -18,26 +18,47 @@
1818
@dataclasses.dataclass
1919
class DatastoreSelection:
2020
"""
21-
Configuration for selecting a datastore to use with neural-lam.
21+
Configuration for selecting a datastore and declaring how its variables
22+
are consumed by the model.
2223
2324
Attributes
2425
----------
2526
kind : str
2627
The kind of datastore to use, currently `mdp` or `npyfilesmeps` are
2728
implemented.
2829
config_path : str
29-
The path to the configuration file for the selected datastore, this is
30-
assumed to be relative to the configuration file for neural-lam.
30+
The path to the configuration file for the selected datastore, this
31+
is assumed to be relative to the configuration file for neural-lam.
32+
inputs : Dict[str, List[str] or None] or None, optional
33+
Per-category lists of variable names this datastore contributes as
34+
model inputs. Categories are typically ``state``, ``forcing``,
35+
``static``. If the whole field is ``None`` (the default), every
36+
variable in every category that the datastore exposes is treated
37+
as an input. If a category key is present with a ``null`` / ``None``
38+
value, all variables in that category are used. An explicit empty
39+
list excludes the category.
40+
outputs : Dict[str, List[str] or None] or None, optional
41+
Per-category lists of variable names this datastore contributes as
42+
model outputs (prediction targets). Categories may include ``state``
43+
for prognostic outputs (those that are fed back as input in the
44+
next autoregressive step) and ``diagnostic`` for predict-only
45+
outputs. Prognostic outputs are the intersection of
46+
``inputs["state"]`` and ``outputs["state"]``; everything in
47+
``outputs`` that is not also in ``inputs["state"]`` is diagnostic.
48+
If ``None`` (the default), this datastore is treated as input-only
49+
(no contribution to predictions). ``null`` per-category values
50+
follow the same "all available" convention as ``inputs``.
3151
"""
3252

3353
kind: str
54+
config_path: str
55+
inputs: Optional[Dict[str, Optional[List[str]]]] = None
56+
outputs: Optional[Dict[str, Optional[List[str]]]] = None
3457

3558
def __post_init__(self):
3659
if self.kind not in DATASTORES:
3760
raise ValueError(f"Datastore kind {self.kind} is not implemented")
3861

39-
config_path: str
40-
4162

4263
@dataclasses.dataclass
4364
class ManualStateFeatureWeighting:
@@ -89,10 +110,10 @@ class TrainingConfig:
89110
Attributes
90111
----------
91112
state_feature_weighting : Union[ManualStateFeatureWeighting,
92-
UnformFeatureWeighting]
113+
UniformFeatureWeighting]
93114
The method to use for weighting the state features in the loss
94-
function. Defaults to uniform weighting (`UnformFeatureWeighting`, i.e.
95-
all features are weighted equally).
115+
function. Defaults to uniform weighting (`UniformFeatureWeighting`,
116+
i.e. all features are weighted equally).
96117
"""
97118

98119
state_feature_weighting: Union[
@@ -112,31 +133,35 @@ class NeuralLAMConfig(dataclass_wizard.JSONWizard, dataclass_wizard.YAMLWizard):
112133
113134
Attributes
114135
----------
115-
datastore : DatastoreSelection
116-
The configuration for the datastore to use.
136+
datastores : Dict[str, DatastoreSelection]
137+
Mapping from user-chosen datastore name to its selection and role
138+
declaration. The dict key becomes the canonical source name used
139+
throughout the pipeline (in ``ForecastBatch`` field dict keys, in
140+
weight / clamping config keys when collisions need to be
141+
disambiguated, etc).
117142
training : TrainingConfig
118143
The configuration for training the model.
119144
"""
120145

121-
datastore: DatastoreSelection
122-
datastore_boundary: Union[DatastoreSelection, None] = None
146+
datastores: Dict[str, DatastoreSelection]
123147
training: TrainingConfig = dataclasses.field(default_factory=TrainingConfig)
124148

125149
class _(dataclass_wizard.JSONWizard.Meta):
126150
"""
127151
Define the configuration class as a JSON wizard class.
128152
129-
Together `tag_key` and `auto_assign_tags` enable that when a `Union` of
130-
types are used for an attribute, the specific type to deserialize to
131-
can be specified in the serialised data using the `tag_key` value. In
132-
our case we call the tag key `__config_class__` to indicate to the
133-
user that they should pick a dataclass describing configuration in
134-
neural-lam. This Union-based selection allows us to support different
135-
configuration attributes for different choices of methods for example
136-
and is used when picking between different feature weighting methods in
137-
the `TrainingConfig` class. `auto_assign_tags` is set to True to
138-
automatically set that tag key (i.e. `__config_class__` in the config
139-
file) should just be the class name of the dataclass to deserialize to.
153+
Together `tag_key` and `auto_assign_tags` enable that when a `Union`
154+
of types are used for an attribute, the specific type to deserialize
155+
to can be specified in the serialised data using the `tag_key`
156+
value. In our case we call the tag key `__config_class__` to
157+
indicate to the user that they should pick a dataclass describing
158+
configuration in neural-lam. This Union-based selection allows us
159+
to support different configuration attributes for different choices
160+
of methods for example and is used when picking between different
161+
feature weighting methods in the `TrainingConfig` class.
162+
`auto_assign_tags` is set to True to automatically set that tag key
163+
(i.e. `__config_class__` in the config file) should just be the
164+
class name of the dataclass to deserialize to.
140165
"""
141166

142167
tag_key = "__config_class__"
@@ -145,25 +170,62 @@ class _(dataclass_wizard.JSONWizard.Meta):
145170
# dataclasses used
146171
# TODO: this should be enabled once
147172
# https://github.com/rnag/dataclass-wizard/issues/137 is fixed, but
148-
# currently cannot be used together with `auto_assign_tags` due to a
149-
# bug it seems
173+
# currently cannot be used together with `auto_assign_tags` due to
174+
# a bug it seems
150175
# raise_on_unknown_json_key = True
151176

152177

153178
class InvalidConfigError(Exception):
154179
pass
155180

156181

182+
def _validate_output_name_collisions(
183+
datastores: Dict[str, Union[MDPDatastore, NpyFilesDatastoreMEPS]],
184+
selections: Dict[str, DatastoreSelection],
185+
) -> None:
186+
"""Raise :class:`InvalidConfigError` if two datastores would contribute
187+
a variable with the same name to the model's output set, since
188+
downstream sites (weight dicts, metric keys, saved zarr coords) cannot
189+
disambiguate.
190+
191+
Fix is to give the colliding variable a unique name in one of the
192+
contributing zarrs (mdp's ``dim_mapping.name_format`` for new builds,
193+
or ``xr.Dataset.assign_coords`` on the small ``{category}_feature``
194+
coord array of an existing zarr - a milliseconds operation regardless
195+
of zarr size).
196+
"""
197+
seen: Dict[str, str] = {}
198+
for ds_name, sel in selections.items():
199+
if sel.outputs is None:
200+
continue
201+
for category, var_list in sel.outputs.items():
202+
if var_list is None:
203+
var_list = datastores[ds_name].get_vars_names(category)
204+
for var in var_list:
205+
if var in seen:
206+
raise InvalidConfigError(
207+
f"Variable '{var}' is declared as an output in "
208+
f"both datastores '{seen[var]}' and '{ds_name}'. "
209+
"Rename the variable in one of the source zarrs "
210+
"(via mdp's `dim_mapping.name_format` for new "
211+
"builds, or via `xr.Dataset.assign_coords` on the "
212+
"existing zarr's `{category}_feature` coord - a "
213+
"milliseconds operation regardless of zarr size). "
214+
"See mllam/neural-lam#652."
215+
)
216+
seen[var] = ds_name
217+
218+
157219
def load_config_and_datastore(
158220
config_path: str,
159221
) -> tuple[
160222
NeuralLAMConfig,
161-
Union[MDPDatastore, NpyFilesDatastoreMEPS],
162-
Union[MDPDatastore, NpyFilesDatastoreMEPS, None],
223+
Dict[str, Union[MDPDatastore, NpyFilesDatastoreMEPS]],
163224
]:
164-
"""
165-
Load the neural-lam configuration and the datastores specified in the
166-
configuration.
225+
"""Load the neural-lam configuration and instantiate each datastore.
226+
227+
The configuration uses the multi-datastore schema introduced for #652:
228+
a top-level ``datastores`` mapping with one entry per source.
167229
168230
Parameters
169231
----------
@@ -172,9 +234,11 @@ def load_config_and_datastore(
172234
173235
Returns
174236
-------
175-
tuple[NeuralLAMConfig, datastore, datastore_boundary]
176-
The Neural-LAM configuration, the loaded (interior) datastore,
177-
and the boundary datastore (or None if not configured).
237+
config : NeuralLAMConfig
238+
The parsed configuration.
239+
datastores : Dict[str, BaseDatastore]
240+
Mapping from each user-chosen datastore name to the loaded
241+
datastore object, in the same order as declared in the config.
178242
"""
179243
try:
180244
config = NeuralLAMConfig.from_yaml_file(config_path)
@@ -183,22 +247,21 @@ def load_config_and_datastore(
183247
"There was an error loading the configuration file at "
184248
f"{config_path}. "
185249
) from ex
186-
# datastore config is assumed to be relative to the config file
187-
datastore_config_path = (
188-
Path(config_path).parent / config.datastore.config_path
189-
)
190-
datastore = init_datastore(
191-
datastore_kind=config.datastore.kind, config_path=datastore_config_path
192-
)
193250

194-
datastore_boundary = None
195-
if config.datastore_boundary is not None:
196-
datastore_boundary_config_path = (
197-
Path(config_path).parent / config.datastore_boundary.config_path
251+
if not config.datastores:
252+
raise InvalidConfigError(
253+
f"Configuration at {config_path} declares no datastores. "
254+
"Add at least one entry under the top-level `datastores:` key."
198255
)
199-
datastore_boundary = init_datastore(
200-
datastore_kind=config.datastore_boundary.kind,
201-
config_path=datastore_boundary_config_path,
256+
257+
config_dir = Path(config_path).parent
258+
loaded: Dict[str, Union[MDPDatastore, NpyFilesDatastoreMEPS]] = {}
259+
for name, selection in config.datastores.items():
260+
datastore_config_path = config_dir / selection.config_path
261+
loaded[name] = init_datastore(
262+
datastore_kind=selection.kind,
263+
config_path=datastore_config_path,
202264
)
203265

204-
return config, datastore, datastore_boundary
266+
_validate_output_name_collisions(loaded, config.datastores)
267+
return config, loaded

‎neural_lam/models/module.py‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -243,10 +243,17 @@ def _create_dataarray_from_tensor(
243243
split: str,
244244
category: str,
245245
) -> xr.DataArray:
246-
weather_dataset = WeatherDataset(datastore=self.datastore, split=split)
246+
# Use the staticmethod variant so we don't instantiate a full
247+
# WeatherDataset (which now requires the multi-datastore dict from
248+
# #652) just to build a single DataArray. The reference dataarray
249+
# only needs the per-grid coords from the datastore.
250+
reference = self.datastore.get_dataarray(category=category, split=split)
247251
time = np.array(time.cpu(), dtype="datetime64[ns]")
248-
da = weather_dataset.create_dataarray_from_tensor(
249-
tensor=tensor, time=time, category=category
252+
da = WeatherDataset.build_dataarray_from_tensor(
253+
reference_dataarray=reference,
254+
tensor=tensor,
255+
time=time,
256+
category=category,
250257
)
251258
return da
252259

‎neural_lam/train_model.py‎

Lines changed: 15 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -391,8 +391,18 @@ def main(input_args=None):
391391
seed.seed_everything(args.seed)
392392

393393
# Load neural-lam configuration and datastores to use
394-
config, datastore, datastore_boundary = load_config_and_datastore(
395-
config_path=args.config_path
394+
config, datastores = load_config_and_datastore(config_path=args.config_path)
395+
396+
# Resolve interior + boundary roles for the legacy single-source
397+
# model side (ForecasterModule, predictor). Multi-source consumption
398+
# on the model side is tracked in #652.
399+
# Local
400+
from .weather_dataset import _resolve_datastore_roles
401+
402+
interior_name, boundary_name = _resolve_datastore_roles(config.datastores)
403+
datastore = datastores[interior_name]
404+
datastore_boundary = (
405+
datastores[boundary_name] if boundary_name is not None else None
396406
)
397407

398408
# Check --var_leads_metrics_watch variable indices against the datastore
@@ -408,16 +418,16 @@ def main(input_args=None):
408418
f"{len(state_var_names)} state variables)."
409419
)
410420

411-
# Create datamodule
421+
# Create datamodule - takes the full multi-source dicts
412422
data_module = WeatherDataModule(
413-
datastore=datastore,
423+
datastores=datastores,
424+
selections=config.datastores,
414425
ar_steps_train=args.ar_steps_train,
415426
ar_steps_eval=args.ar_steps_eval,
416427
num_past_forcing_steps=args.num_past_forcing_steps,
417428
num_future_forcing_steps=args.num_future_forcing_steps,
418429
num_past_boundary_steps=args.num_past_boundary_steps,
419430
num_future_boundary_steps=args.num_future_boundary_steps,
420-
datastore_boundary=datastore_boundary,
421431
load_single_member=args.load_single_member,
422432
batch_size=args.batch_size,
423433
num_workers=args.num_workers,

0 commit comments

Comments
 (0)