From 5ea5508fe438e6fc3ffab50f4a566312e44a59f1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Wed, 23 Jul 2025 16:14:14 +0200 Subject: [PATCH 01/27] fix unused argument --- .../training/diagnostics/callbacks/lora.py | 55 +++++++++++++++++++ 1 file changed, 55 insertions(+) create mode 100644 training/src/anemoi/training/diagnostics/callbacks/lora.py diff --git a/training/src/anemoi/training/diagnostics/callbacks/lora.py b/training/src/anemoi/training/diagnostics/callbacks/lora.py new file mode 100644 index 0000000000..4d63f3e225 --- /dev/null +++ b/training/src/anemoi/training/diagnostics/callbacks/lora.py @@ -0,0 +1,55 @@ +# (C) Copyright 2024 Anemoi contributors. +# (C) Copyright 2025 BULL SAS. +# +# This software is licensed under the terms of the Apache Licence Version 2.0 +# which can be obtained at http://www.apache.org/licenses/LICENSE-2.0. +# +# In applying this licence, ECMWF does not waive the privileges and immunities +# granted to it by virtue of its status as an intergovernmental organisation +# nor does it submit to any jurisdiction. +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +from peft import LoraConfig +from peft import TaskType +from peft import get_peft_model +from pytorch_lightning.callbacks import Callback + +if TYPE_CHECKING: + import pytorch_lightning as pl + from omegaconf import OmegaConf + +LOGGER = logging.getLogger(__name__) + + +class LoRAAdapters(Callback): + """Inject LoRA adapters in a pre-trained model for fine-tuning.""" + + def __init__(self, config: OmegaConf, **kwargs) -> None: + """Initialize LoRA injection callback. + + Parameters + ---------- + config : dict + Dictionary with configuration settings + kwargs : dict + Keyword arguments for LoRA configuration, such as `r`, `lora_alpha`, `lora_dropout`, etc. + See `peft.LoraConfig` for details. + """ + super().__init__() + self.config = config + self.lora_config = LoraConfig(**kwargs, task_type=TaskType.FEATURE_EXTRACTION) + + def on_fit_start(self, _: pl.Trainer, pl_module: pl.LightningModule) -> None: + """Check the order of the variables in the model from checkpoint and the training data. + + Parameters + ---------- + _ : pl.Trainer + Not used + pl_module : pl.LightningModule + Pytorch Lightning module + """ + get_peft_model(pl_module.model, self.lora_config) From 67b0d8a1ef99495b4b9a8b6c9de43b12f7707455 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Fri, 25 Jul 2025 14:48:37 +0200 Subject: [PATCH 02/27] Add lora as a lightning_module inheriting from GraphForecaster, activate tensorboard logging --- .../training/diagnostics/callbacks/lora.py | 55 ----------- .../training/train/forecaster/__init__.py | 3 +- .../training/train/forecaster/forecaster.py | 6 +- .../train/forecaster/loraforecaster.py | 98 +++++++++++++++++++ training/src/anemoi/training/train/train.py | 1 + 5 files changed, 106 insertions(+), 57 deletions(-) delete mode 100644 training/src/anemoi/training/diagnostics/callbacks/lora.py create mode 100644 training/src/anemoi/training/train/forecaster/loraforecaster.py diff --git a/training/src/anemoi/training/diagnostics/callbacks/lora.py b/training/src/anemoi/training/diagnostics/callbacks/lora.py deleted file mode 100644 index 4d63f3e225..0000000000 --- a/training/src/anemoi/training/diagnostics/callbacks/lora.py +++ /dev/null @@ -1,55 +0,0 @@ -# (C) Copyright 2024 Anemoi contributors. -# (C) Copyright 2025 BULL SAS. -# -# This software is licensed under the terms of the Apache Licence Version 2.0 -# which can be obtained at http://www.apache.org/licenses/LICENSE-2.0. -# -# In applying this licence, ECMWF does not waive the privileges and immunities -# granted to it by virtue of its status as an intergovernmental organisation -# nor does it submit to any jurisdiction. -from __future__ import annotations - -import logging -from typing import TYPE_CHECKING - -from peft import LoraConfig -from peft import TaskType -from peft import get_peft_model -from pytorch_lightning.callbacks import Callback - -if TYPE_CHECKING: - import pytorch_lightning as pl - from omegaconf import OmegaConf - -LOGGER = logging.getLogger(__name__) - - -class LoRAAdapters(Callback): - """Inject LoRA adapters in a pre-trained model for fine-tuning.""" - - def __init__(self, config: OmegaConf, **kwargs) -> None: - """Initialize LoRA injection callback. - - Parameters - ---------- - config : dict - Dictionary with configuration settings - kwargs : dict - Keyword arguments for LoRA configuration, such as `r`, `lora_alpha`, `lora_dropout`, etc. - See `peft.LoraConfig` for details. - """ - super().__init__() - self.config = config - self.lora_config = LoraConfig(**kwargs, task_type=TaskType.FEATURE_EXTRACTION) - - def on_fit_start(self, _: pl.Trainer, pl_module: pl.LightningModule) -> None: - """Check the order of the variables in the model from checkpoint and the training data. - - Parameters - ---------- - _ : pl.Trainer - Not used - pl_module : pl.LightningModule - Pytorch Lightning module - """ - get_peft_model(pl_module.model, self.lora_config) diff --git a/training/src/anemoi/training/train/forecaster/__init__.py b/training/src/anemoi/training/train/forecaster/__init__.py index a545beabf5..2ae1aa3dee 100644 --- a/training/src/anemoi/training/train/forecaster/__init__.py +++ b/training/src/anemoi/training/train/forecaster/__init__.py @@ -10,5 +10,6 @@ from .ensforecaster import GraphEnsForecaster from .forecaster import GraphForecaster from .interpolator import GraphInterpolator +from .loraforecaster import LoRAGraphForecaster -__all__ = ["GraphEnsForecaster", "GraphForecaster", "GraphInterpolator"] +__all__ = ["GraphEnsForecaster", "GraphForecaster", "GraphInterpolator", "LoRAGraphForecaster"] diff --git a/training/src/anemoi/training/train/forecaster/forecaster.py b/training/src/anemoi/training/train/forecaster/forecaster.py index f118252f5f..8231e5a757 100644 --- a/training/src/anemoi/training/train/forecaster/forecaster.py +++ b/training/src/anemoi/training/train/forecaster/forecaster.py @@ -102,7 +102,11 @@ def __init__( self.latlons_data = graph_data[config.graph.data].x self.statistics_tendencies = statistics_tendencies - self.logger_enabled = config.diagnostics.log.wandb.enabled or config.diagnostics.log.mlflow.enabled + self.logger_enabled = ( + config.diagnostics.log.wandb.enabled + or config.diagnostics.log.mlflow.enabled + or config.diagnostics.log.tensorboard.enabled + ) metadata_extractor = ExtractVariableGroupAndLevel( variable_groups=config.model_dump(by_alias=True).training.variable_groups, diff --git a/training/src/anemoi/training/train/forecaster/loraforecaster.py b/training/src/anemoi/training/train/forecaster/loraforecaster.py new file mode 100644 index 0000000000..b9851cd33b --- /dev/null +++ b/training/src/anemoi/training/train/forecaster/loraforecaster.py @@ -0,0 +1,98 @@ +# (C) Copyright 2024 Anemoi contributors. +# +# This software is licensed under the terms of the Apache Licence Version 2.0 +# which can be obtained at http://www.apache.org/licenses/LICENSE-2.0. +# +# In applying this licence, ECMWF does not waive the privileges and immunities +# granted to it by virtue of its status as an intergovernmental organisation +# nor does it submit to any jurisdiction. +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +from peft import LoraConfig +from peft import get_peft_model + +from anemoi.models.data_indices.collection import IndexCollection +from anemoi.training.train.forecaster import GraphForecaster + +if TYPE_CHECKING: + import torch + from torch_geometric.data import HeteroData + + from anemoi.models.data_indices.collection import IndexCollection + from anemoi.training.schemas.base_schema import BaseSchema + + +LOGGER = logging.getLogger(__name__) + + +class LoRAGraphForecaster(GraphForecaster): + """Graph neural network forecaster for PyTorch Lightning.""" + + def __init__( + self, + *, + config: BaseSchema, + graph_data: HeteroData, + truncation_data: dict, + statistics: dict, + statistics_tendencies: dict, + data_indices: IndexCollection, + metadata: dict, + supporting_arrays: dict, + ) -> None: + """Initialize graph neural network forecaster. + + Parameters + ---------- + config : DictConfig + Job configuration + graph_data : HeteroData + Graph object + statistics : dict + Statistics of the training data + data_indices : IndexCollection + Indices of the training data, + metadata : dict + Provenance information + supporting_arrays : dict + Supporting NumPy arrays to store in the checkpoint + + """ + super().__init__( + config=config, + graph_data=graph_data, + truncation_data=truncation_data, + statistics=statistics, + statistics_tendencies=statistics_tendencies, + data_indices=data_indices, + metadata=metadata, + supporting_arrays=supporting_arrays, + ) + + self.lora_config = LoraConfig( + r=8, + lora_alpha=32, + target_modules=[ + "lin_key", + "lin_query", + "lin_value", + "lin_self", + "lin_v", + "lin_k", + "lin_q", + "projection", + "mlp.0", + "mlp.2", + "node_dst_mlp.0", + "node_dst_mlp.2", + ], + modules_to_save=["node_data_extractor.1"], + task_type=None, + ) + + def on_load_checkpoint(self, checkpoint: torch.nn.module) -> None: + self._ckpt_model_name_to_index = checkpoint["hyper_parameters"]["data_indices"].name_to_index + get_peft_model(self.model, self.lora_config) diff --git a/training/src/anemoi/training/train/train.py b/training/src/anemoi/training/train/train.py index e8860a5e21..7ebece7c3f 100644 --- a/training/src/anemoi/training/train/train.py +++ b/training/src/anemoi/training/train/train.py @@ -219,6 +219,7 @@ def model(self) -> pl.LightningModule: if self.config.training.transfer_learning: LOGGER.info("Loading weights with Transfer Learning from %s", self.last_checkpoint) model = transfer_learning_loading(model, self.last_checkpoint) + model.on_load_checkpoint(self.last_checkpoint) else: LOGGER.info("Restoring only model weights from %s", self.last_checkpoint) # pop data_indices so that the data indices on the checkpoint do not get overwritten From 86cf84db16eefadc5efb556e7b4ebaa650537946 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Wed, 5 Nov 2025 14:53:03 +0100 Subject: [PATCH 03/27] Fix meging with main --- training/src/anemoi/training/train/tasks/loraforecaster.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/training/src/anemoi/training/train/tasks/loraforecaster.py b/training/src/anemoi/training/train/tasks/loraforecaster.py index b9851cd33b..d7b3bb6101 100644 --- a/training/src/anemoi/training/train/tasks/loraforecaster.py +++ b/training/src/anemoi/training/train/tasks/loraforecaster.py @@ -1,4 +1,5 @@ # (C) Copyright 2024 Anemoi contributors. +# Copyright (C) Bull S.A.S - 2025 # # This software is licensed under the terms of the Apache Licence Version 2.0 # which can be obtained at http://www.apache.org/licenses/LICENSE-2.0. @@ -14,11 +15,11 @@ from peft import LoraConfig from peft import get_peft_model -from anemoi.models.data_indices.collection import IndexCollection -from anemoi.training.train.forecaster import GraphForecaster +import torch + +from anemoi.training.train.tasks import GraphForecaster if TYPE_CHECKING: - import torch from torch_geometric.data import HeteroData from anemoi.models.data_indices.collection import IndexCollection From 79fb33a0e190bbb9f76eab731f04191a9517e200 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Wed, 5 Nov 2025 14:53:19 +0100 Subject: [PATCH 04/27] Add comment to identify change related to lora --- training/src/anemoi/training/train/train.py | 1 + 1 file changed, 1 insertion(+) diff --git a/training/src/anemoi/training/train/train.py b/training/src/anemoi/training/train/train.py index cef73cd70c..5f1039d93e 100644 --- a/training/src/anemoi/training/train/train.py +++ b/training/src/anemoi/training/train/train.py @@ -229,6 +229,7 @@ def model(self) -> pl.LightningModule: if self.config.training.transfer_learning: LOGGER.info("Loading weights with Transfer Learning from %s", self.last_checkpoint) model = transfer_learning_loading(model, self.last_checkpoint) + # Added for LoRA #TODO remove when better strategy is implemented model.on_load_checkpoint(self.last_checkpoint) else: LOGGER.info("Restoring only model weights from %s", self.last_checkpoint) From e06f366a24424e3ef616bb306d8f727788c6f4e4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Thu, 6 Nov 2025 11:14:58 +0100 Subject: [PATCH 05/27] update dependencies --- training/pyproject.toml | 1 + 1 file changed, 1 insertion(+) diff --git a/training/pyproject.toml b/training/pyproject.toml index 446115a814..52b4795b50 100644 --- a/training/pyproject.toml +++ b/training/pyproject.toml @@ -52,6 +52,7 @@ dependencies = [ "mlflow-skinny>=2.11.1", "numpy<2", # Pinned until we can confirm it works with anemoi graphs "nvidia-ml-py>=13.580.82", + "peft", "pydantic>=2.9", "pyshtools>=4.13", "pytorch-lightning>=2.1", From 95311cf582e97b19215a04931a28af0695368ca2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Thu, 6 Nov 2025 14:21:36 +0100 Subject: [PATCH 06/27] Add LoRAForecaster to training schemas --- training/src/anemoi/training/schemas/training.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/training/src/anemoi/training/schemas/training.py b/training/src/anemoi/training/schemas/training.py index c4b63401e3..759525379d 100644 --- a/training/src/anemoi/training/schemas/training.py +++ b/training/src/anemoi/training/schemas/training.py @@ -386,6 +386,11 @@ class DiffusionTendForecasterSchema(ForecasterSchema): "Training objective." +class LoRAForecasterSchema(ForecasterSchema): + model_task: Literal["anemoi.training.train.tasks.LoRAGraphForecaster",] = Field(..., alias="model_task") + "Training objective." + + class InterpolationSchema(BaseTrainingSchema): model_task: Literal["anemoi.training.train.tasks.GraphInterpolator"] = Field(..., alias="model_task") "Training objective." From 7e6159a62cc9e19209e4c53a5a24edddd3019806 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Thu, 6 Nov 2025 14:28:19 +0100 Subject: [PATCH 07/27] fix lora schema --- training/src/anemoi/training/schemas/training.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/training/src/anemoi/training/schemas/training.py b/training/src/anemoi/training/schemas/training.py index 759525379d..1f250a3613 100644 --- a/training/src/anemoi/training/schemas/training.py +++ b/training/src/anemoi/training/schemas/training.py @@ -405,6 +405,7 @@ class InterpolationSchema(BaseTrainingSchema): | ForecasterEnsSchema | InterpolationSchema | DiffusionForecasterSchema - | DiffusionTendForecasterSchema, + | DiffusionTendForecasterSchema + | LoRAForecasterSchema, Discriminator("model_task"), ] From 7862505dd68ee7fd38aecfbb2894cc26d5db532d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Thu, 6 Nov 2025 14:45:10 +0100 Subject: [PATCH 08/27] fix lora schema --- training/src/anemoi/training/schemas/training.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/training/src/anemoi/training/schemas/training.py b/training/src/anemoi/training/schemas/training.py index 1f250a3613..de8d8d6415 100644 --- a/training/src/anemoi/training/schemas/training.py +++ b/training/src/anemoi/training/schemas/training.py @@ -387,7 +387,7 @@ class DiffusionTendForecasterSchema(ForecasterSchema): class LoRAForecasterSchema(ForecasterSchema): - model_task: Literal["anemoi.training.train.tasks.LoRAGraphForecaster",] = Field(..., alias="model_task") + model_task: Literal["anemoi.training.train.tasks.LoRAGraphForecaster"] = Field(..., alias="model_task") "Training objective." From 4c19dd992454e00fade04b8201bc31493cf5dcb4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Wed, 12 Nov 2025 13:59:35 +0100 Subject: [PATCH 09/27] update peft version --- training/pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/training/pyproject.toml b/training/pyproject.toml index 52b4795b50..81e3688f0a 100644 --- a/training/pyproject.toml +++ b/training/pyproject.toml @@ -52,7 +52,7 @@ dependencies = [ "mlflow-skinny>=2.11.1", "numpy<2", # Pinned until we can confirm it works with anemoi graphs "nvidia-ml-py>=13.580.82", - "peft", + "peft>=0.17.1", "pydantic>=2.9", "pyshtools>=4.13", "pytorch-lightning>=2.1", From 152ae04a1a5ca90e346cbf85ea9b82ee67ee5c89 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Wed, 12 Nov 2025 14:00:17 +0100 Subject: [PATCH 10/27] add lora config --- training/src/anemoi/training/config/lora.yaml | 14 ++ .../anemoi/training/config/training/lora.yaml | 147 ++++++++++++++++++ 2 files changed, 161 insertions(+) create mode 100644 training/src/anemoi/training/config/lora.yaml create mode 100644 training/src/anemoi/training/config/training/lora.yaml diff --git a/training/src/anemoi/training/config/lora.yaml b/training/src/anemoi/training/config/lora.yaml new file mode 100644 index 0000000000..d630ad68f5 --- /dev/null +++ b/training/src/anemoi/training/config/lora.yaml @@ -0,0 +1,14 @@ +defaults: +- data: zarr +- dataloader: native_grid +- diagnostics: evaluation +- datamodule: single +- hardware: example +- graph: multi_scale +- model: gnn +- training: lora +- _self_ + + +# set to true to switch on config validation +config_validation: True diff --git a/training/src/anemoi/training/config/training/lora.yaml b/training/src/anemoi/training/config/training/lora.yaml new file mode 100644 index 0000000000..dba45ae031 --- /dev/null +++ b/training/src/anemoi/training/config/training/lora.yaml @@ -0,0 +1,147 @@ +--- +defaults: + - scalers: global + +# resume or fork a training from a checkpoint last.ckpt or specified in hardware.files.warm_start +run_id: null +fork_run_id: ??? +transfer_learning: False # activate to perform transfer learning +load_weights_only: True # only load model weights, do not restore optimiser states etc. + +# run in deterministic mode ; slows down +deterministic: False + +# miscellaneous +precision: 16-mixed + +# multistep input +# 1 = single step scheme, X(t-1) used to predict X(t) +# k > 1: multistep scheme, uses [X(t-k), X(t-k+1), ... X(t-1)] to predict X(t) +# Deepmind use k = 2 in their model +multistep_input: 2 + +# gradient accumulation across K batches, K >= 1 (if K == 1 then no accumulation) +# the effective batch size becomes num-devices * batch_size * k +accum_grad_batches: 1 + +num_sanity_val_steps: 6 + +# clipp gradients, 0 : don't clip, default algorithm: norm, alternative: value +gradient_clip: + val: 32. + algorithm: value + +# stochastic weight averaging +# https://pytorch.org/blog/stochastic-weight-averaging-in-pytorch/ +swa: + enabled: False + lr: 1.e-4 + +# Optimizer settings +optimizer: + zero: False # use ZeroRedundancyOptimizer ; saves memory for larger models + kwargs: + betas: [0.9, 0.95] + +# select model +model_task: anemoi.training.train.tasks.LoRAGraphForecaster + +# select strategy +strategy: + _target_: anemoi.training.distributed.strategy.DDPGroupStrategy + num_gpus_per_model: ${hardware.num_gpus_per_model} + read_group_size: ${dataloader.read_group_size} + +# loss functions + +# dynamic rescaling of the loss gradient +# see https://arxiv.org/pdf/2306.06079.pdf, section 4.3.2 +# don't enable this by default until it's been tested and proven beneficial +loss_gradient_scaling: False + +# loss function for the model +training_loss: + # loss class to initialise + _target_: anemoi.training.losses.MSELoss + # Scalers to include in loss calculation + # A selection of available scalers are listed in training/scalers. + # '*' is a valid entry to use all `scalers` given, if a scaler is to be excluded + # add `!scaler_name`, i.e. ['*', '!scaler_1'], and `scaler_1` will not be added. + scalers: ['pressure_level', 'general_variable', 'node_weights'] + ignore_nans: False + +# Validation metrics calculation, +# This may be a list, in which case all metrics will be calculated +# and logged according to their name. +# These metrics are calculated in the output model space, and thus +# have undergone postprocessing. +validation_metrics: + # loss class to initialise + mse: + _target_: anemoi.training.losses.MSELoss + # Scalers to include in loss calculation + # Cannot scale over the variable dimension due to possible remappings. + # Available scalers include: + # - 'loss_weights_mask': Giving imputed NaNs a zero weight in the loss function + # Use the `scale_validation_metrics` section to variable scale. + scalers: ['node_weights'] + # other kwargs + ignore_nans: True + +# Variable groups definition for scaling +# The variable level scaling methods are defined under training/scalers +# A default group is required and is appended as prefix to the metric of all variables not assigned to a group. +# Variables are assigned to a group by their param if contained in the metadata, else by their name. + +# If more complex grouping is required, groups can be defined as a dictionary, such that all +# keys must be evaluate to True. +# .e.g. to set the variable group based on if the metadata specifies the variable is a pressure level +# you can write the following: +# variable_groups: +# default: sfc +# pl: +# is_pressure_level: True +# See `anemoi.transform.variables.Variable` for the available metadata. +# Note that the former formulation of +# : +# variable_groups: +# default: sfc +# pl: [q, t, u, v, w, z] +# +# still works +# param is an alias for the variable name in the case of no metadata. + +variable_groups: + default: sfc + pl: + param: [q, t, u, v, w, z] + +metrics: +- z_500 +- t_850 +- u_850 +- v_850 + +# length of the "rollout" window (see Keisler's paper) +rollout: + start: 1 + # increase rollout every n epochs + epoch_increment: 0 + # maximum rollout to use + max: 1 + +# Set max_epochs or max_steps. Training stops at the first limit reached. +max_epochs: null +max_steps: 150000 + +lr: + warmup: 1000 # number of warmup iterations + rate: 0.625e-4 #local_lr + iterations: ${training.max_steps} # NOTE: When max_epochs < max_steps, scheduler will run for max_steps + min: 3e-7 #Not scaled by #GPU + +# Changes in per-gpu batch_size should come with a rescaling of the local_lr +# in order to keep a constant global_lr +# global_lr = local_lr * num_gpus_per_node * num_nodes / gpus_per_model + +submodules_to_freeze: [] From e12b38801181eedc1d00397d577e0f66892c0e37 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Wed, 12 Nov 2025 15:02:05 +0100 Subject: [PATCH 11/27] Fix lora training test definition --- .../tests/integration/config/test_lora.yaml | 15 +++++++++ training/tests/integration/conftest.py | 33 +++++++++++++++++++ .../tests/integration/test_training_cycle.py | 13 ++++++++ 3 files changed, 61 insertions(+) create mode 100644 training/tests/integration/config/test_lora.yaml diff --git a/training/tests/integration/config/test_lora.yaml b/training/tests/integration/config/test_lora.yaml new file mode 100644 index 0000000000..dedd810281 --- /dev/null +++ b/training/tests/integration/config/test_lora.yaml @@ -0,0 +1,15 @@ +# Modifications for the basic template "config.yaml" +hardware: + files: + dataset: anemoi-integration-tests/training/datasets/aifs-ea-an-oper-0001-mars-o48-1979-19-6h-v6-testset.zarr + paths: + data: https://object-store.os-api.cci1.ecmwf.int/ml-tests/test-data/samples + +training: + fork_run_id: "dummy_id" + +dataloader: + training: + end: 1979-01-08 12:00:00 + validation: + start: 1979-01-08 18:00:00 diff --git a/training/tests/integration/conftest.py b/training/tests/integration/conftest.py index d8ad77c00f..1254879428 100644 --- a/training/tests/integration/conftest.py +++ b/training/tests/integration/conftest.py @@ -383,3 +383,36 @@ def diffusion_config( cfg = OmegaConf.merge(template, testing_modifications_with_temp_dir, use_case_modifications) OmegaConf.resolve(cfg) return cfg, dataset_urls[0] + + +@pytest.fixture +def lora_config( + testing_modifications_with_temp_dir: DictConfig, + get_tmp_paths: GetTmpPaths, + get_test_data: GetTestData, + migrator: Migrator, +) -> tuple[DictConfig, str]: + with initialize(version_base=None, config_path="../../src/anemoi/training/config", job_name="test_lora"): + template = compose(config_name="lora") + + use_case_modifications = OmegaConf.load(Path.cwd() / "training/tests/integration/config/test_lora.yaml") + assert isinstance(use_case_modifications, DictConfig) + + tmp_dir, rel_paths, dataset_urls = get_tmp_paths(use_case_modifications, ["dataset"]) + use_case_modifications.hardware.paths.data = tmp_dir + use_case_modifications.hardware.files.dataset = rel_paths[0] + + cfg = OmegaConf.merge(template, testing_modifications_with_temp_dir, use_case_modifications) + OmegaConf.resolve(cfg) + assert isinstance(cfg, DictConfig) + + existing_ckpt = get_test_data( + "anemoi-integration-tests/training/checkpoints/testing-checkpoint-gnn-global-2025-07-31.ckpt", + ) + _, new_ckpt, _ = migrator.sync(existing_ckpt) + + checkpoint_dir = Path(cfg.hardware.paths.output + "checkpoint/dummy_id") + checkpoint_dir.mkdir(parents=True, exist_ok=True) + torch.save(new_ckpt, checkpoint_dir / "last.ckpt") + + return cfg, dataset_urls[0] \ No newline at end of file diff --git a/training/tests/integration/test_training_cycle.py b/training/tests/integration/test_training_cycle.py index 66afc239ce..cf0eb8a290 100644 --- a/training/tests/integration/test_training_cycle.py +++ b/training/tests/integration/test_training_cycle.py @@ -206,3 +206,16 @@ def test_training_cycle_diffusion(diffusion_config: tuple[DictConfig, str], get_ def test_config_validation_diffusion(diffusion_config: tuple[DictConfig, str]) -> None: cfg, _ = diffusion_config BaseSchema(**cfg) + + +@skip_if_offline +@pytest.mark.slow +def test_training_cycle_lora(lora_config: tuple[DictConfig, list[str]], get_test_archive: GetTestArchive) -> None: + cfg, url = lora_config + get_test_archive(url) + AnemoiTrainer(cfg).train() + + +def test_config_validation_lora(lora_config: tuple[DictConfig, str]) -> None: + cfg, _ = lora_config + BaseSchema(**cfg) From a7cdc3dbf42fe7b203d498d9f126d5a808cb583e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Wed, 12 Nov 2025 17:26:55 +0100 Subject: [PATCH 12/27] Make lora parameters configurable --- training/src/anemoi/training/config/lora.yaml | 7 +++++++ .../src/anemoi/training/schemas/training.py | 20 ++++++++++++++++++ .../training/train/tasks/loraforecaster.py | 21 +------------------ 3 files changed, 28 insertions(+), 20 deletions(-) diff --git a/training/src/anemoi/training/config/lora.yaml b/training/src/anemoi/training/config/lora.yaml index d630ad68f5..419db4ffd9 100644 --- a/training/src/anemoi/training/config/lora.yaml +++ b/training/src/anemoi/training/config/lora.yaml @@ -12,3 +12,10 @@ defaults: # set to true to switch on config validation config_validation: True + +training: + lora_config: + r: 8 + lora_alpha: 32 + target_modules: ['mlp.0', 'mlp.2', 'dummy_layer'] + modules_to_save: ['node_data_extractor.1'] diff --git a/training/src/anemoi/training/schemas/training.py b/training/src/anemoi/training/schemas/training.py index de8d8d6415..2714f644fa 100644 --- a/training/src/anemoi/training/schemas/training.py +++ b/training/src/anemoi/training/schemas/training.py @@ -13,6 +13,7 @@ from typing import Any from typing import Literal +from peft import LoraConfig from pydantic import AfterValidator from pydantic import Discriminator from pydantic import Field @@ -111,6 +112,23 @@ class TargetForcing(BaseModel): "Use target time as a fraction between input boundary times as input." +class LoRAConfig(BaseModel): + """LoRA parameters. + + See https://huggingface.co/docs/peft/package_reference/lora#peft.LoraConfig + for more information. + """ + + r: int = 8 + "Lora attention dimension (the “rank”)." + lora_alpha: int = 32 + "The alpha parameter for Lora scaling." + target_modules: list[str] = Field(examples=["mlp.0", "mlp.2"]) + "The names of the modules to apply the adapter to." + modules_to_save: list[str] = Field(examples=["node_data_extractor.1"]) + "List of modules apart from adapter layers to be set as trainable." + + class LossScalingSchema(BaseModel): default: int = 1 "Default scaling value applied to the variables loss. Default to 1." @@ -389,6 +407,8 @@ class DiffusionTendForecasterSchema(ForecasterSchema): class LoRAForecasterSchema(ForecasterSchema): model_task: Literal["anemoi.training.train.tasks.LoRAGraphForecaster"] = Field(..., alias="model_task") "Training objective." + lora_config: LoRAConfig + "Configuration for the LoRA adapter." class InterpolationSchema(BaseTrainingSchema): diff --git a/training/src/anemoi/training/train/tasks/loraforecaster.py b/training/src/anemoi/training/train/tasks/loraforecaster.py index d7b3bb6101..6a37821183 100644 --- a/training/src/anemoi/training/train/tasks/loraforecaster.py +++ b/training/src/anemoi/training/train/tasks/loraforecaster.py @@ -73,26 +73,7 @@ def __init__( supporting_arrays=supporting_arrays, ) - self.lora_config = LoraConfig( - r=8, - lora_alpha=32, - target_modules=[ - "lin_key", - "lin_query", - "lin_value", - "lin_self", - "lin_v", - "lin_k", - "lin_q", - "projection", - "mlp.0", - "mlp.2", - "node_dst_mlp.0", - "node_dst_mlp.2", - ], - modules_to_save=["node_data_extractor.1"], - task_type=None, - ) + self.lora_config = LoraConfig(**config.model_dump(by_alias=True).training.lora_config) def on_load_checkpoint(self, checkpoint: torch.nn.module) -> None: self._ckpt_model_name_to_index = checkpoint["hyper_parameters"]["data_indices"].name_to_index From 2d8ca6ec64a772cf87643255ee06fb014d06a608 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Fri, 14 Nov 2025 15:12:38 +0100 Subject: [PATCH 13/27] fix lora injection by adding a hook in BaseGraphModule --- training/src/anemoi/training/train/tasks/base.py | 3 +++ training/src/anemoi/training/train/tasks/loraforecaster.py | 4 ++-- training/src/anemoi/training/train/train.py | 4 +++- 3 files changed, 8 insertions(+), 3 deletions(-) diff --git a/training/src/anemoi/training/train/tasks/base.py b/training/src/anemoi/training/train/tasks/base.py index 09d585a69d..06f41288e3 100644 --- a/training/src/anemoi/training/train/tasks/base.py +++ b/training/src/anemoi/training/train/tasks/base.py @@ -311,6 +311,9 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: def on_load_checkpoint(self, checkpoint: torch.nn.Module) -> None: self._ckpt_model_name_to_index = checkpoint["hyper_parameters"]["data_indices"].name_to_index + + def on_checkpoint_loaded(self) -> None: + pass def update_scalers(self, callback: AvailableCallbacks) -> None: """Update scalers, calling the defined function on them, updating if not None.""" diff --git a/training/src/anemoi/training/train/tasks/loraforecaster.py b/training/src/anemoi/training/train/tasks/loraforecaster.py index 6a37821183..955ff0bd6a 100644 --- a/training/src/anemoi/training/train/tasks/loraforecaster.py +++ b/training/src/anemoi/training/train/tasks/loraforecaster.py @@ -75,6 +75,6 @@ def __init__( self.lora_config = LoraConfig(**config.model_dump(by_alias=True).training.lora_config) - def on_load_checkpoint(self, checkpoint: torch.nn.module) -> None: - self._ckpt_model_name_to_index = checkpoint["hyper_parameters"]["data_indices"].name_to_index + def on_checkpoint_loaded(self) -> None: get_peft_model(self.model, self.lora_config) + LOGGER.info("LoRA adapters injected into the model") diff --git a/training/src/anemoi/training/train/train.py b/training/src/anemoi/training/train/train.py index 5f1039d93e..58ecbcc89b 100644 --- a/training/src/anemoi/training/train/train.py +++ b/training/src/anemoi/training/train/train.py @@ -230,13 +230,15 @@ def model(self) -> pl.LightningModule: LOGGER.info("Loading weights with Transfer Learning from %s", self.last_checkpoint) model = transfer_learning_loading(model, self.last_checkpoint) # Added for LoRA #TODO remove when better strategy is implemented - model.on_load_checkpoint(self.last_checkpoint) + model.on_checkpoint_loaded() else: LOGGER.info("Restoring only model weights from %s", self.last_checkpoint) # pop data_indices so that the data indices on the checkpoint do not get overwritten # by the data indices from the new config kwargs.pop("data_indices") model = model_task.load_from_checkpoint(self.last_checkpoint, **kwargs, strict=False) + # Added for LoRA #TODO remove when better strategy is implemented + model.on_checkpoint_loaded() model.data_indices = self.data_indices # check data indices in original checkpoint and current data indices are the same From 29b3ff801be83c70684e676c8bdf2219dcea8a81 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Fri, 14 Nov 2025 15:57:41 +0100 Subject: [PATCH 14/27] remove unused import --- training/src/anemoi/training/schemas/training.py | 1 - 1 file changed, 1 deletion(-) diff --git a/training/src/anemoi/training/schemas/training.py b/training/src/anemoi/training/schemas/training.py index 2714f644fa..a33bf02348 100644 --- a/training/src/anemoi/training/schemas/training.py +++ b/training/src/anemoi/training/schemas/training.py @@ -13,7 +13,6 @@ from typing import Any from typing import Literal -from peft import LoraConfig from pydantic import AfterValidator from pydantic import Discriminator from pydantic import Field From 5bb9b3007c5f9e78d57aa1dc8dc76f58d4544573 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Mon, 17 Nov 2025 16:56:34 +0100 Subject: [PATCH 15/27] Improve formatting --- training/src/anemoi/training/schemas/training.py | 4 ++-- training/src/anemoi/training/train/tasks/__init__.py | 2 +- training/src/anemoi/training/train/tasks/base.py | 2 +- training/src/anemoi/training/train/tasks/loraforecaster.py | 4 +--- training/src/anemoi/training/train/train.py | 6 ++++-- 5 files changed, 9 insertions(+), 9 deletions(-) diff --git a/training/src/anemoi/training/schemas/training.py b/training/src/anemoi/training/schemas/training.py index a33bf02348..d34d3797e8 100644 --- a/training/src/anemoi/training/schemas/training.py +++ b/training/src/anemoi/training/schemas/training.py @@ -113,8 +113,8 @@ class TargetForcing(BaseModel): class LoRAConfig(BaseModel): """LoRA parameters. - - See https://huggingface.co/docs/peft/package_reference/lora#peft.LoraConfig + + See https://huggingface.co/docs/peft/package_reference/lora#peft.LoraConfig for more information. """ diff --git a/training/src/anemoi/training/train/tasks/__init__.py b/training/src/anemoi/training/train/tasks/__init__.py index 49bca447a1..c35a6cfd22 100644 --- a/training/src/anemoi/training/train/tasks/__init__.py +++ b/training/src/anemoi/training/train/tasks/__init__.py @@ -21,5 +21,5 @@ "GraphEnsForecaster", "GraphForecaster", "GraphInterpolator", - "LoRAGraphForecaster" + "LoRAGraphForecaster", ] diff --git a/training/src/anemoi/training/train/tasks/base.py b/training/src/anemoi/training/train/tasks/base.py index 06f41288e3..f90532ff20 100644 --- a/training/src/anemoi/training/train/tasks/base.py +++ b/training/src/anemoi/training/train/tasks/base.py @@ -311,7 +311,7 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: def on_load_checkpoint(self, checkpoint: torch.nn.Module) -> None: self._ckpt_model_name_to_index = checkpoint["hyper_parameters"]["data_indices"].name_to_index - + def on_checkpoint_loaded(self) -> None: pass diff --git a/training/src/anemoi/training/train/tasks/loraforecaster.py b/training/src/anemoi/training/train/tasks/loraforecaster.py index 955ff0bd6a..d40e2a1b71 100644 --- a/training/src/anemoi/training/train/tasks/loraforecaster.py +++ b/training/src/anemoi/training/train/tasks/loraforecaster.py @@ -1,5 +1,5 @@ # (C) Copyright 2024 Anemoi contributors. -# Copyright (C) Bull S.A.S - 2025 +# Copyright (C) Bull S.A.S - 2025 # # This software is licensed under the terms of the Apache Licence Version 2.0 # which can be obtained at http://www.apache.org/licenses/LICENSE-2.0. @@ -15,8 +15,6 @@ from peft import LoraConfig from peft import get_peft_model -import torch - from anemoi.training.train.tasks import GraphForecaster if TYPE_CHECKING: diff --git a/training/src/anemoi/training/train/train.py b/training/src/anemoi/training/train/train.py index 58ecbcc89b..3df620268f 100644 --- a/training/src/anemoi/training/train/train.py +++ b/training/src/anemoi/training/train/train.py @@ -229,7 +229,8 @@ def model(self) -> pl.LightningModule: if self.config.training.transfer_learning: LOGGER.info("Loading weights with Transfer Learning from %s", self.last_checkpoint) model = transfer_learning_loading(model, self.last_checkpoint) - # Added for LoRA #TODO remove when better strategy is implemented + # Added for LoRA + # TODO(Mikael): remove when better strategy is implemented model.on_checkpoint_loaded() else: LOGGER.info("Restoring only model weights from %s", self.last_checkpoint) @@ -237,7 +238,8 @@ def model(self) -> pl.LightningModule: # by the data indices from the new config kwargs.pop("data_indices") model = model_task.load_from_checkpoint(self.last_checkpoint, **kwargs, strict=False) - # Added for LoRA #TODO remove when better strategy is implemented + # Added for LoRA + # TODO(Mikael): remove when better strategy is implemented model.on_checkpoint_loaded() model.data_indices = self.data_indices From 620358aef341416113c0140a1c5dcacebdc6b23b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Wed, 19 Nov 2025 15:23:29 +0100 Subject: [PATCH 16/27] Add lora config default parameters in training:lora config --- training/src/anemoi/training/config/training/lora.yaml | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/training/src/anemoi/training/config/training/lora.yaml b/training/src/anemoi/training/config/training/lora.yaml index dba45ae031..817d83cbab 100644 --- a/training/src/anemoi/training/config/training/lora.yaml +++ b/training/src/anemoi/training/config/training/lora.yaml @@ -46,6 +46,12 @@ optimizer: # select model model_task: anemoi.training.train.tasks.LoRAGraphForecaster +lora_config: + r: 8 + lora_alpha: 32 + target_modules: ['mlp.0', 'mlp.2', 'dummy_layer'] + modules_to_save: ['node_data_extractor.1'] + # select strategy strategy: _target_: anemoi.training.distributed.strategy.DDPGroupStrategy From 1fe72051c83fb8bec11f02b5143aa060843d8c5d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mika=C3=ABl=20Jacquemont?= Date: Mon, 24 Nov 2025 17:33:12 +0100 Subject: [PATCH 17/27] add docstring to on_checkpoint_loaded --- training/src/anemoi/training/train/tasks/base.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/training/src/anemoi/training/train/tasks/base.py b/training/src/anemoi/training/train/tasks/base.py index f90532ff20..8c654418d3 100644 --- a/training/src/anemoi/training/train/tasks/base.py +++ b/training/src/anemoi/training/train/tasks/base.py @@ -313,6 +313,10 @@ def on_load_checkpoint(self, checkpoint: torch.nn.Module) -> None: self._ckpt_model_name_to_index = checkpoint["hyper_parameters"]["data_indices"].name_to_index def on_checkpoint_loaded(self) -> None: + """Called by anemoi-training once a model has been loaded from a checkpoint, either via + transfer learning or weights only setting. Override to perform actions on the model + once the checkpoint has been loaded. + """ pass def update_scalers(self, callback: AvailableCallbacks) -> None: From d95acc3bc3b29022b9628358e733a57d887bbe47 Mon Sep 17 00:00:00 2001 From: mikael10j Date: Fri, 3 Apr 2026 17:27:54 +0200 Subject: [PATCH 18/27] update lora forecaster to load lora ckpt --- .../src/anemoi/training/train/tasks/loraforecaster.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/training/src/anemoi/training/train/tasks/loraforecaster.py b/training/src/anemoi/training/train/tasks/loraforecaster.py index d40e2a1b71..d5010c3be6 100644 --- a/training/src/anemoi/training/train/tasks/loraforecaster.py +++ b/training/src/anemoi/training/train/tasks/loraforecaster.py @@ -13,6 +13,7 @@ from typing import TYPE_CHECKING from peft import LoraConfig +from peft import PeftModel from peft import get_peft_model from anemoi.training.train.tasks import GraphForecaster @@ -73,6 +74,14 @@ def __init__( self.lora_config = LoraConfig(**config.model_dump(by_alias=True).training.lora_config) + def on_load_checkpoint(self, checkpoint) -> None: + if 'LoRAGraphForecaster' in checkpoint['hyper_parameters']['config'].training.model_task: + self._inject_lora_adapters() + def on_checkpoint_loaded(self) -> None: + if not isinstance(self.model, PeftModel): + self._inject_lora_adapters() + + def _inject_lora_adapters(self) -> None: get_peft_model(self.model, self.lora_config) LOGGER.info("LoRA adapters injected into the model") From 8ed69040c0040acf5edaff862578bb414ff5bac4 Mon Sep 17 00:00:00 2001 From: mikael10j Date: Fri, 3 Apr 2026 17:44:56 +0200 Subject: [PATCH 19/27] add test to restart from a lora ckpt --- .../tests/integration/test_training_cycle.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/training/tests/integration/test_training_cycle.py b/training/tests/integration/test_training_cycle.py index cf0eb8a290..5bd4724ef5 100644 --- a/training/tests/integration/test_training_cycle.py +++ b/training/tests/integration/test_training_cycle.py @@ -215,6 +215,22 @@ def test_training_cycle_lora(lora_config: tuple[DictConfig, list[str]], get_test get_test_archive(url) AnemoiTrainer(cfg).train() + output_dir = Path(cfg.hardware.paths.output + "checkpoint") + + assert output_dir.exists(), f"Checkpoint directory not found at: {output_dir}" + + run_dirs = [item for item in output_dir.iterdir() if item.is_dir()] + assert ( + len(run_dirs) == 1 + ), f"Expected exactly one run_id directory, found {len(run_dirs)}: {[d.name for d in run_dirs]}" + + checkpoint_dir = run_dirs[0] + assert len(list(checkpoint_dir.glob("anemoi-by_epoch-*.ckpt"))) == 2, "Expected 2 checkpoints after first run" + + cfg.training.run_id = checkpoint_dir.name + cfg.training.max_epochs = cfg.training.max_epochs + 3 + AnemoiTrainer(cfg).train() + def test_config_validation_lora(lora_config: tuple[DictConfig, str]) -> None: cfg, _ = lora_config From 135dbf501b827d425704747e321aadad300a6fd1 Mon Sep 17 00:00:00 2001 From: mikael10j Date: Wed, 8 Apr 2026 17:46:48 +0200 Subject: [PATCH 20/27] fix lora integration test --- training/tests/integration/test_training_cycle.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/training/tests/integration/test_training_cycle.py b/training/tests/integration/test_training_cycle.py index 5bd4724ef5..c659c480b5 100644 --- a/training/tests/integration/test_training_cycle.py +++ b/training/tests/integration/test_training_cycle.py @@ -219,11 +219,11 @@ def test_training_cycle_lora(lora_config: tuple[DictConfig, list[str]], get_test assert output_dir.exists(), f"Checkpoint directory not found at: {output_dir}" - run_dirs = [item for item in output_dir.iterdir() if item.is_dir()] + run_dirs = [item for item in output_dir.iterdir() if item.is_dir() and item.name != "dummy_id"] assert ( len(run_dirs) == 1 ), f"Expected exactly one run_id directory, found {len(run_dirs)}: {[d.name for d in run_dirs]}" - + checkpoint_dir = run_dirs[0] assert len(list(checkpoint_dir.glob("anemoi-by_epoch-*.ckpt"))) == 2, "Expected 2 checkpoints after first run" From 805659f89d48564220f46963942780d440286cd2 Mon Sep 17 00:00:00 2001 From: mikael10j Date: Wed, 8 Apr 2026 17:47:04 +0200 Subject: [PATCH 21/27] fix ckpt name to index --- training/src/anemoi/training/train/tasks/loraforecaster.py | 1 + 1 file changed, 1 insertion(+) diff --git a/training/src/anemoi/training/train/tasks/loraforecaster.py b/training/src/anemoi/training/train/tasks/loraforecaster.py index d5010c3be6..5ff6bdd564 100644 --- a/training/src/anemoi/training/train/tasks/loraforecaster.py +++ b/training/src/anemoi/training/train/tasks/loraforecaster.py @@ -75,6 +75,7 @@ def __init__( self.lora_config = LoraConfig(**config.model_dump(by_alias=True).training.lora_config) def on_load_checkpoint(self, checkpoint) -> None: + self._ckpt_model_name_to_index = checkpoint["hyper_parameters"]["data_indices"].name_to_index if 'LoRAGraphForecaster' in checkpoint['hyper_parameters']['config'].training.model_task: self._inject_lora_adapters() From 0d8834bf3ed839698d3b7a3529beb29c3b0782d5 Mon Sep 17 00:00:00 2001 From: mikael10j Date: Fri, 24 Apr 2026 16:54:33 +0200 Subject: [PATCH 22/27] update lora training config Co-authored-by: Copilot --- training/src/anemoi/training/config/training/lora.yaml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/training/src/anemoi/training/config/training/lora.yaml b/training/src/anemoi/training/config/training/lora.yaml index 817d83cbab..8680d07fb7 100644 --- a/training/src/anemoi/training/config/training/lora.yaml +++ b/training/src/anemoi/training/config/training/lora.yaml @@ -44,7 +44,7 @@ optimizer: betas: [0.9, 0.95] # select model -model_task: anemoi.training.train.tasks.LoRAGraphForecaster +training_method: anemoi.training.train.methods.LoRASingleTraining lora_config: r: 8 From a97e870972e3d94c9ef51df9b7971e50e4e2ffc4 Mon Sep 17 00:00:00 2001 From: mikael10j Date: Fri, 24 Apr 2026 17:22:06 +0200 Subject: [PATCH 23/27] fix format, fix lora integration test config, fix lora type checking import, fix lora type annotation, move torch into type-checking block Co-authored-by: Copilot --- training/src/anemoi/training/schemas/training.py | 4 +--- training/src/anemoi/training/train/methods/__init__.py | 4 ++-- training/src/anemoi/training/train/methods/lora_single.py | 8 +++++--- training/tests/integration/conftest.py | 6 +++--- training/tests/integration/test_training_cycle.py | 2 +- 5 files changed, 12 insertions(+), 12 deletions(-) diff --git a/training/src/anemoi/training/schemas/training.py b/training/src/anemoi/training/schemas/training.py index 5a818db661..230b91231f 100644 --- a/training/src/anemoi/training/schemas/training.py +++ b/training/src/anemoi/training/schemas/training.py @@ -492,9 +492,7 @@ class DiffusionTendencyTrainingSchema(BaseTrainingSchema): class LoRASingleTrainingSchema(BaseTrainingSchema): - training_method: Literal["anemoi.training.train.methods.LoRASingleTraining"] = Field( - ..., - alias="training_method") + training_method: Literal["anemoi.training.train.methods.LoRASingleTraining"] = Field(..., alias="training_method") "Training objective." lora_config: LoRAConfig "Configuration for the LoRA adapter." diff --git a/training/src/anemoi/training/train/methods/__init__.py b/training/src/anemoi/training/train/methods/__init__.py index 75d8d85c47..84a6166ee3 100644 --- a/training/src/anemoi/training/train/methods/__init__.py +++ b/training/src/anemoi/training/train/methods/__init__.py @@ -10,13 +10,13 @@ from .diffusion import DiffusionTendencyTraining from .diffusion import DiffusionTraining from .ensemble import EnsembleTraining -from .single import SingleTraining from .lora_single import LoRASingleTraining +from .single import SingleTraining __all__ = [ "DiffusionTendencyTraining", "DiffusionTraining", "EnsembleTraining", - "SingleTraining", "LoRASingleTraining", + "SingleTraining", ] diff --git a/training/src/anemoi/training/train/methods/lora_single.py b/training/src/anemoi/training/train/methods/lora_single.py index 152df8d053..8737fbbdd7 100644 --- a/training/src/anemoi/training/train/methods/lora_single.py +++ b/training/src/anemoi/training/train/methods/lora_single.py @@ -19,10 +19,12 @@ from anemoi.training.train.methods.single import SingleTraining if TYPE_CHECKING: + import torch from torch_geometric.data import HeteroData from anemoi.models.data_indices.collection import IndexCollection from anemoi.training.schemas.base_schema import BaseSchema + from training.src.anemoi.training.tasks.base import BaseTask LOGGER = logging.getLogger(__name__) @@ -57,16 +59,16 @@ def __init__( self.lora_config = LoraConfig(**config.model_dump(by_alias=True).training.lora_config) - def on_load_checkpoint(self, checkpoint) -> None: + def on_load_checkpoint(self, checkpoint: torch.nn.Module) -> None: self._update_checkpoint_state_dict_for_load(checkpoint) self._ckpt_model_name_to_index = { dataset_name: data_indices.name_to_index for dataset_name, data_indices in checkpoint["hyper_parameters"]["data_indices"].items() } - if 'LoRASingleTraining' in checkpoint['hyper_parameters']['config'].training.training_method: + if "LoRASingleTraining" in checkpoint["hyper_parameters"]["config"].training.training_method: self._inject_lora_adapters() - + def on_checkpoint_loaded(self) -> None: if not isinstance(self.model, PeftModel): self._inject_lora_adapters() diff --git a/training/tests/integration/conftest.py b/training/tests/integration/conftest.py index ec20d98a62..eb92f9c7a0 100644 --- a/training/tests/integration/conftest.py +++ b/training/tests/integration/conftest.py @@ -552,7 +552,7 @@ def diffusion_config( @pytest.fixture def lora_config( testing_modifications_with_temp_dir: DictConfig, - get_tmp_paths: GetTmpPaths, + get_tmp_path: GetTmpPath, get_test_data: GetTestData, migrator: Migrator, ) -> tuple[DictConfig, str]: @@ -567,8 +567,8 @@ def lora_config( OmegaConf.resolve(cfg) existing_ckpt = get_test_data( - "anemoi-integration-tests/training/checkpoints/testing-checkpoint-gnn-global-2025-07-31.ckpt", - ) + "anemoi-integration-tests/training/checkpoints/testing-checkpoint-gnn-global-2025-07-31.ckpt", + ) _, new_ckpt, _ = migrator.sync(existing_ckpt) checkpoint_dir = Path(cfg.system.output.root + "/" + cfg.system.output.checkpoints.root + "/dummy_id") diff --git a/training/tests/integration/test_training_cycle.py b/training/tests/integration/test_training_cycle.py index eb18ebce2b..dbda216c98 100644 --- a/training/tests/integration/test_training_cycle.py +++ b/training/tests/integration/test_training_cycle.py @@ -360,7 +360,7 @@ def test_training_cycle_lora(lora_config: tuple[DictConfig, list[str]], get_test assert ( len(run_dirs) == 1 ), f"Expected exactly one run_id directory, found {len(run_dirs)}: {[d.name for d in run_dirs]}" - + checkpoint_dir = run_dirs[0] assert len(list(checkpoint_dir.glob("anemoi-by_epoch-*.ckpt"))) == 2, "Expected 2 checkpoints after first run" From 75fb6fd602792393adb132f3fa0cebf71daca51f Mon Sep 17 00:00:00 2001 From: mikael10j Date: Fri, 24 Apr 2026 17:28:40 +0200 Subject: [PATCH 24/27] fix default lora config Co-authored-by: Copilot --- training/src/anemoi/training/config/lora.yaml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/training/src/anemoi/training/config/lora.yaml b/training/src/anemoi/training/config/lora.yaml index 419db4ffd9..03c0f98a7a 100644 --- a/training/src/anemoi/training/config/lora.yaml +++ b/training/src/anemoi/training/config/lora.yaml @@ -2,10 +2,10 @@ defaults: - data: zarr - dataloader: native_grid - diagnostics: evaluation -- datamodule: single -- hardware: example +- system: example - graph: multi_scale - model: gnn +- task: forecaster - training: lora - _self_ From db380709fe12462ccc9adcfe5da7f299a199b9e8 Mon Sep 17 00:00:00 2001 From: mikael10j Date: Mon, 27 Apr 2026 09:45:09 +0200 Subject: [PATCH 25/27] update lora config, lint BaseTrainingModule, update ckpt version Co-authored-by: Copilot --- .../anemoi/training/config/training/lora.yaml | 93 +++++++------------ .../src/anemoi/training/train/methods/base.py | 10 +- training/tests/integration/conftest.py | 2 +- 3 files changed, 42 insertions(+), 63 deletions(-) diff --git a/training/src/anemoi/training/config/training/lora.yaml b/training/src/anemoi/training/config/training/lora.yaml index 8680d07fb7..ae6fa37ad7 100644 --- a/training/src/anemoi/training/config/training/lora.yaml +++ b/training/src/anemoi/training/config/training/lora.yaml @@ -1,6 +1,7 @@ --- defaults: - scalers: global + - optimization: default # resume or fork a training from a checkpoint last.ckpt or specified in hardware.files.warm_start run_id: null @@ -14,12 +15,6 @@ deterministic: False # miscellaneous precision: 16-mixed -# multistep input -# 1 = single step scheme, X(t-1) used to predict X(t) -# k > 1: multistep scheme, uses [X(t-k), X(t-k+1), ... X(t-1)] to predict X(t) -# Deepmind use k = 2 in their model -multistep_input: 2 - # gradient accumulation across K batches, K >= 1 (if K == 1 then no accumulation) # the effective batch size becomes num-devices * batch_size * k accum_grad_batches: 1 @@ -37,12 +32,6 @@ swa: enabled: False lr: 1.e-4 -# Optimizer settings -optimizer: - zero: False # use ZeroRedundancyOptimizer ; saves memory for larger models - kwargs: - betas: [0.9, 0.95] - # select model training_method: anemoi.training.train.methods.LoRASingleTraining @@ -55,7 +44,7 @@ lora_config: # select strategy strategy: _target_: anemoi.training.distributed.strategy.DDPGroupStrategy - num_gpus_per_model: ${hardware.num_gpus_per_model} + num_gpus_per_model: ${system.hardware.num_gpus_per_model} read_group_size: ${dataloader.read_group_size} # loss functions @@ -67,14 +56,16 @@ loss_gradient_scaling: False # loss function for the model training_loss: - # loss class to initialise - _target_: anemoi.training.losses.MSELoss - # Scalers to include in loss calculation - # A selection of available scalers are listed in training/scalers. - # '*' is a valid entry to use all `scalers` given, if a scaler is to be excluded - # add `!scaler_name`, i.e. ['*', '!scaler_1'], and `scaler_1` will not be added. - scalers: ['pressure_level', 'general_variable', 'node_weights'] - ignore_nans: False + datasets: + data: # user-defined key in data + # loss class to initialise + _target_: anemoi.training.losses.MSELoss + # Scalers to include in loss calculation + # A selection of available scalers are listed in training/scalers. + # '*' is a valid entry to use all `scalers` given, if a scaler is to be excluded + # add `!scaler_name`, i.e. ['*', '!scaler_1'], and `scaler_1` will not be added. + scalers: ['pressure_level', 'general_variable', 'node_weights'] + ignore_nans: False # Validation metrics calculation, # This may be a list, in which case all metrics will be calculated @@ -82,17 +73,19 @@ training_loss: # These metrics are calculated in the output model space, and thus # have undergone postprocessing. validation_metrics: - # loss class to initialise - mse: - _target_: anemoi.training.losses.MSELoss - # Scalers to include in loss calculation - # Cannot scale over the variable dimension due to possible remappings. - # Available scalers include: - # - 'loss_weights_mask': Giving imputed NaNs a zero weight in the loss function - # Use the `scale_validation_metrics` section to variable scale. - scalers: ['node_weights'] - # other kwargs - ignore_nans: True + datasets: + data: # user-defined key in data + # loss class to initialise + mse: + _target_: anemoi.training.losses.MSELoss + # Scalers to include in loss calculation + # Cannot scale over the variable dimension due to possible remappings. + # Available scalers include: + # - 'loss_weights_mask': Giving imputed NaNs a zero weight in the loss function + # Use the `scale_validation_metrics` section to variable scale. + scalers: ['node_weights'] + # other kwargs + ignore_nans: True # Variable groups definition for scaling # The variable level scaling methods are defined under training/scalers @@ -118,36 +111,22 @@ validation_metrics: # param is an alias for the variable name in the case of no metadata. variable_groups: - default: sfc - pl: - param: [q, t, u, v, w, z] + datasets: + data: # user-defined key in data + default: sfc + pl: + param: [q, t, u, v, w, z] metrics: -- z_500 -- t_850 -- u_850 -- v_850 - -# length of the "rollout" window (see Keisler's paper) -rollout: - start: 1 - # increase rollout every n epochs - epoch_increment: 0 - # maximum rollout to use - max: 1 + datasets: + data: # user-defined key in data + - z_500 + - t_850 + - u_850 + - v_850 # Set max_epochs or max_steps. Training stops at the first limit reached. max_epochs: null max_steps: 150000 -lr: - warmup: 1000 # number of warmup iterations - rate: 0.625e-4 #local_lr - iterations: ${training.max_steps} # NOTE: When max_epochs < max_steps, scheduler will run for max_steps - min: 3e-7 #Not scaled by #GPU - -# Changes in per-gpu batch_size should come with a rescaling of the local_lr -# in order to keep a constant global_lr -# global_lr = local_lr * num_gpus_per_node * num_nodes / gpus_per_model - submodules_to_freeze: [] diff --git a/training/src/anemoi/training/train/methods/base.py b/training/src/anemoi/training/train/methods/base.py index 2954923a1e..826d784cdb 100644 --- a/training/src/anemoi/training/train/methods/base.py +++ b/training/src/anemoi/training/train/methods/base.py @@ -443,12 +443,12 @@ def on_load_checkpoint(self, checkpoint: torch.nn.Module) -> None: } def on_checkpoint_loaded(self) -> None: - """Called by anemoi-training once a model has been loaded from a checkpoint, either via - transfer learning or weights only setting. Override to perform actions on the model - once the checkpoint has been loaded. + """Called once a model has been loaded from a checkpoint. + + Override to perform actions on the model. This applies to both transfer + learning and weights-only settings. """ - pass - + def _update_scaler_for_dataset( self, name: str, diff --git a/training/tests/integration/conftest.py b/training/tests/integration/conftest.py index eb92f9c7a0..0529c4f15a 100644 --- a/training/tests/integration/conftest.py +++ b/training/tests/integration/conftest.py @@ -567,7 +567,7 @@ def lora_config( OmegaConf.resolve(cfg) existing_ckpt = get_test_data( - "anemoi-integration-tests/training/checkpoints/testing-checkpoint-gnn-global-2025-07-31.ckpt", + "anemoi-integration-tests/training/checkpoints/testing-checkpoint-gnn-global-2026-03-06.ckpt", ) _, new_ckpt, _ = migrator.sync(existing_ckpt) From 1e746b553863c6b338dd51709c65ec3512ddf5ca Mon Sep 17 00:00:00 2001 From: mikael10j Date: Mon, 27 Apr 2026 14:23:21 +0200 Subject: [PATCH 26/27] update weight averaging in lora config Co-authored-by: Copilot --- training/src/anemoi/training/config/training/lora.yaml | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/training/src/anemoi/training/config/training/lora.yaml b/training/src/anemoi/training/config/training/lora.yaml index ae6fa37ad7..6cf0594922 100644 --- a/training/src/anemoi/training/config/training/lora.yaml +++ b/training/src/anemoi/training/config/training/lora.yaml @@ -2,6 +2,7 @@ defaults: - scalers: global - optimization: default + - weight_averaging: null # resume or fork a training from a checkpoint last.ckpt or specified in hardware.files.warm_start run_id: null @@ -26,12 +27,6 @@ gradient_clip: val: 32. algorithm: value -# stochastic weight averaging -# https://pytorch.org/blog/stochastic-weight-averaging-in-pytorch/ -swa: - enabled: False - lr: 1.e-4 - # select model training_method: anemoi.training.train.methods.LoRASingleTraining From 8e930edb512fe6d58f99e280a8c9d82bbd3dda4e Mon Sep 17 00:00:00 2001 From: mikael10j Date: Tue, 28 Apr 2026 11:51:24 +0200 Subject: [PATCH 27/27] update lora training config, improve lora integration test, fix lora config loading in training method Co-authored-by: Copilot --- .../anemoi/training/config/training/lora.yaml | 9 ++++++--- .../training/train/methods/lora_single.py | 2 +- .../tests/integration/config/test_lora.yaml | 6 +++--- training/tests/integration/conftest.py | 3 +++ .../tests/integration/test_training_cycle.py | 18 ++++++++++++++++-- 5 files changed, 29 insertions(+), 9 deletions(-) diff --git a/training/src/anemoi/training/config/training/lora.yaml b/training/src/anemoi/training/config/training/lora.yaml index 6cf0594922..d1707a8ff8 100644 --- a/training/src/anemoi/training/config/training/lora.yaml +++ b/training/src/anemoi/training/config/training/lora.yaml @@ -7,8 +7,11 @@ defaults: # resume or fork a training from a checkpoint last.ckpt or specified in hardware.files.warm_start run_id: null fork_run_id: ??? -transfer_learning: False # activate to perform transfer learning +transfer_learning: True # activate to perform transfer learning load_weights_only: True # only load model weights, do not restore optimiser states etc. +update_ds_stats_on_ckpt_load: + states: False # rebuild state processors from current dataset when loading a checkpoint + tendencies: True # rebuild tendency processors from current dataset when loading a checkpoint # run in deterministic mode ; slows down deterministic: False @@ -59,7 +62,7 @@ training_loss: # A selection of available scalers are listed in training/scalers. # '*' is a valid entry to use all `scalers` given, if a scaler is to be excluded # add `!scaler_name`, i.e. ['*', '!scaler_1'], and `scaler_1` will not be added. - scalers: ['pressure_level', 'general_variable', 'node_weights'] + scalers: ['pressure_level', 'general_variable', 'node_weights', 'time_steps'] ignore_nans: False # Validation metrics calculation, @@ -78,7 +81,7 @@ validation_metrics: # Available scalers include: # - 'loss_weights_mask': Giving imputed NaNs a zero weight in the loss function # Use the `scale_validation_metrics` section to variable scale. - scalers: ['node_weights'] + scalers: ['node_weights', 'time_steps'] # other kwargs ignore_nans: True diff --git a/training/src/anemoi/training/train/methods/lora_single.py b/training/src/anemoi/training/train/methods/lora_single.py index 8737fbbdd7..5952ecc757 100644 --- a/training/src/anemoi/training/train/methods/lora_single.py +++ b/training/src/anemoi/training/train/methods/lora_single.py @@ -57,7 +57,7 @@ def __init__( supporting_arrays=supporting_arrays, ) - self.lora_config = LoraConfig(**config.model_dump(by_alias=True).training.lora_config) + self.lora_config = LoraConfig(**config.training.lora_config) def on_load_checkpoint(self, checkpoint: torch.nn.Module) -> None: self._update_checkpoint_state_dict_for_load(checkpoint) diff --git a/training/tests/integration/config/test_lora.yaml b/training/tests/integration/config/test_lora.yaml index 2b4c247462..9b44cac73a 100644 --- a/training/tests/integration/config/test_lora.yaml +++ b/training/tests/integration/config/test_lora.yaml @@ -1,7 +1,7 @@ # Modifications for the basic template "config.yaml" system: input: - dataset: anemoi-integration-tests/training/datasets/aifs-ea-an-oper-0001-mars-o48-1979-19-6h-v6-testset.zarr + dataset: anemoi-integration-tests/training/datasets/aifs-ea-an-oper-0001-mars-o96-2017-2017-6h-v8-testing.zarr training: fork_run_id: "dummy_id" @@ -10,8 +10,8 @@ dataloader: training: datasets: data: - end: 1979-01-07 12:00:00 + end: 2017-01-08 12:00:00 validation: datasets: data: - start: 1979-01-08 18:00:00 + start: 2017-01-08 18:00:00 diff --git a/training/tests/integration/conftest.py b/training/tests/integration/conftest.py index 0529c4f15a..137caa7b49 100644 --- a/training/tests/integration/conftest.py +++ b/training/tests/integration/conftest.py @@ -575,6 +575,9 @@ def lora_config( checkpoint_dir.mkdir(parents=True, exist_ok=True) torch.save(new_ckpt, checkpoint_dir / "last.ckpt") + cfg.diagnostics.plot.callbacks = [] # remove plotting callbacks as they are tested in global training cycle test + cfg.diagnostics.callbacks = [] # remove RolloutEval callback as it is tested in global training cycle test + return cfg, url_dataset diff --git a/training/tests/integration/test_training_cycle.py b/training/tests/integration/test_training_cycle.py index dbda216c98..dc5b85d0c5 100644 --- a/training/tests/integration/test_training_cycle.py +++ b/training/tests/integration/test_training_cycle.py @@ -351,7 +351,14 @@ def test_config_validation_diffusion(diffusion_config: tuple[DictConfig, str]) - def test_training_cycle_lora(lora_config: tuple[DictConfig, list[str]], get_test_archive: GetTestArchive) -> None: cfg, url = lora_config get_test_archive(url) - AnemoiTrainer(cfg).train() + trainer = AnemoiTrainer(cfg) + has_lora = False + for name, _param in trainer.model.named_parameters(): + if "lora_A" in name: + has_lora = True + break + assert has_lora, "Expected model to have LoRA embeddings for LoRA training" + trainer.train() output_dir = Path(cfg.system.output.root + "/" + cfg.system.output.checkpoints.root) assert output_dir.exists(), f"Checkpoint directory not found at: {output_dir}" @@ -366,7 +373,14 @@ def test_training_cycle_lora(lora_config: tuple[DictConfig, list[str]], get_test cfg.training.run_id = checkpoint_dir.name cfg.training.max_epochs = cfg.training.max_epochs + 3 - AnemoiTrainer(cfg).train() + trainer = AnemoiTrainer(cfg) + has_lora = False + for name, _param in trainer.model.named_parameters(): + if "lora_A" in name: + has_lora = True + break + assert has_lora, "Expected model to have LoRA embeddings for LoRA training" + trainer.train() def test_config_validation_lora(lora_config: tuple[DictConfig, str]) -> None: