diff --git a/benchmark_utils/base_solver.py b/benchmark_utils/base_solver.py index f6f3c66..947c797 100644 --- a/benchmark_utils/base_solver.py +++ b/benchmark_utils/base_solver.py @@ -279,6 +279,10 @@ class BaseTSFMSolver(BaseSolver): dtype The data type of both data and model. Default to bfloat16 on CUDA, float32 elsewhere. + + inference_batch_size + Max number of (series, cutoff) windows sent to the model per + forward call. Subclasses may override this to tune memory/throughput. """ supported_tasks: set[TaskType] @@ -296,6 +300,10 @@ def __init__(self, **kwargs: Any) -> None: super().__init__() self._loaded_model_id: str | None = None self.model: Any = None + + # Possibly overridden by kwargs: + self.inference_batch_size = 128 + for key, value in kwargs.items(): setattr(self, key, value) @@ -579,10 +587,16 @@ def forecast( if not inputs: return ForecastOutput(quantiles=[], quantile_levels=quantile_levels) - # TODO We still do this in batches in case data is very large - - # Get a list of model outputs aligned with inputs - raw = self.forecast_batch(inputs, covariates, prediction_length) + # Run in batches so very large datasets do not go through the model + # as a single oversized batch. Maintains ordering. + raw: list[torch.Tensor] = [] + for start in range(0, len(inputs), self.inference_batch_size): + end = start + self.inference_batch_size + raw.extend( + self.forecast_batch( + inputs[start:end], covariates[start:end], prediction_length + ) + ) per_series_preds = [[None] * n_cutoffs for _, n_cutoffs in per_series_shape] for (series_idx, cutoff_idx), pred in zip(layout, raw):