Skip to content

Stepper

samudra.stepper

Time-stepping primitives for training, validation, and inference.

Provides module-level functions that handle single-step forward passes (train_batch, validate_batch) and multi-step autoregressive rollouts (run_rollout).

validate_rollout(model, dataset, aggregators_by_step, epoch, *, num_model_steps, num_model_steps_forward)

Run a bounded autoregressive validation rollout and record RMSE metrics.

aggregators_by_step maps a rollout horizon, in model steps, to the aggregator that should record metrics for that horizon.

Source code in src/samudra/stepper.py
@torch.no_grad()
def validate_rollout(
    model: BaseModel,
    dataset: InferenceDataset,
    aggregators_by_step: Mapping[int, RolloutValidationAggregator],
    epoch: int,
    *,
    num_model_steps: int,
    num_model_steps_forward: int,
) -> None:
    """Run a bounded autoregressive validation rollout and record RMSE metrics.

    ``aggregators_by_step`` maps a rollout horizon, in model steps, to the
    aggregator that should record metrics for that horizon.
    """
    if not aggregators_by_step:
        raise ValueError("At least one rollout validation aggregator is required")

    dataset.to(get_device())
    num_model_steps = min(num_model_steps, len(dataset))
    horizon_boundaries = _rollout_horizon_boundaries(
        aggregators_by_step,
        num_model_steps,
    )
    chunks = _get_rollout_step_chunks(
        total_steps=num_model_steps,
        num_model_steps_forward=num_model_steps_forward,
        boundaries=horizon_boundaries,
    )

    initial_prognostic = dataset.initial_prognostic
    step = 0
    for loop, num_steps in enumerate(chunks):
        logger.info(
            f"Rollout validation [epoch {epoch}]: loop {loop} of "
            f"{len(chunks) - 1}. Stepping {num_steps} steps forward."
        )
        output = model.inference(
            dataset,
            initial_prognostic=initial_prognostic,
            steps_completed=step,
            num_steps=num_steps,
            epoch=epoch,
        )
        initial_prognostic = output.final_prognostic.clone()
        _record_rollout_horizon_batch(
            aggregators_by_step=aggregators_by_step,
            step=step,
            num_steps=num_steps,
            output=output,
        )
        step += num_steps

run_rollout(model, dataset, inf_aggregator, epoch, output_dir=None, model_path=None, num_model_steps_forward=200, save_zarr=False, data_layout=None, preprocessor=None)

Performs inference, which is an auto-regressive rollout.

Source code in src/samudra/stepper.py
@torch.no_grad()
def run_rollout(
    model: BaseModel,
    dataset: InferenceDataset,
    inf_aggregator: InferenceEvaluatorAggregator,
    epoch: int,
    output_dir: str | PathLike | None = None,
    model_path: str | PathLike | None = None,
    num_model_steps_forward: int = 200,
    save_zarr: bool = False,
    data_layout: DataLayout | None = None,
    preprocessor: BatchPreprocessor | None = None,
) -> None:
    """Performs inference, which is an auto-regressive rollout."""
    if save_zarr:
        if output_dir is None or model_path is None:
            raise ValueError(
                "output_dir and model_path must be provided if save_zarr is True"
            )
        if data_layout is None or preprocessor is None:
            raise ValueError(
                "data_layout and preprocessor must be provided if save_zarr is True"
            )
        coords = dataset.get_coords_dict()
        if num_model_steps_forward > 0:
            chunk_size = num_model_steps_forward
        else:
            chunk_size = 20
        writer = ZarrWriter(
            output_dir,
            coords=coords,
            output_steps=dataset.output_steps,
            model_path=model_path,
            time_chunk_size=chunk_size,
            preprocessor=preprocessor,
            data_layout=data_layout,
        )
    else:
        writer = None
    record_logs = get_record_to_wandb(label="inference")
    logger.info(f"Inference [epoch {epoch}]: processing initial prognostic.")
    logs = inf_aggregator.record_initial_prognostic(
        initial_prognostic=dataset.initial_prognostic.to(get_device()),
    )
    record_logs(logs)
    num_model_steps = len(dataset)
    num_steps_list = _get_rollout_step_chunks(
        total_steps=num_model_steps,
        num_model_steps_forward=num_model_steps_forward,
    )

    num_loops = len(num_steps_list)
    initial_prognostic = dataset.initial_prognostic
    step = 0
    for loop, num_steps in enumerate(num_steps_list):
        logger.info(
            f"Inference [epoch {epoch}]: loop {loop} of {num_loops - 1}. "
            f"Stepping {num_steps} steps forward."
        )
        dataset.to(get_device())
        inference_output: ModelInferenceOutput = model.inference(
            dataset,
            initial_prognostic=initial_prognostic,
            steps_completed=step,
            num_steps=num_steps,
            epoch=epoch,
        )
        # Setting initial prognostic for next loop
        initial_prognostic = inference_output.final_prognostic.clone()
        if writer:
            logger.info("Writing to zarr...")
            writer.record_batch(inference_output)
            writer.write()

        logger.info("Recording logs...")
        logs = inf_aggregator.record_batch(inference_output)
        logger.info("Logging to wandb...")
        record_logs(logs)
        step += num_steps