Kronos predict vs predict_batch: Single vs Batch Time-Series Inference Explained
The predict() method generates forecasts for a single financial instrument returning one DataFrame, while predict_batch() processes multiple instruments in parallel using batched tensors, returning a list of DataFrames optimized for high-throughput scenarios.
The KronosPredictor class in the Kronos repository provides two distinct inference interfaces for financial time-series forecasting. Understanding the architectural differences between Kronos predict vs predict_batch is essential for optimizing both interactive analysis workflows and production forecasting pipelines that handle large portfolios. Both methods reside in model/kronos.py and leverage the same autoregressive inference core, but differ fundamentally in input handling, tensor dimensionality, and result aggregation.
Key Differences Between predict() and predict_batch()
Input Signatures and Validation Logic
predict() accepts individual arguments for a single series: a pandas.DataFrame, x_timestamp, and y_timestamp. According to the source code in model/kronos.py, it performs validation checks once—ensuring required price columns exist and handling missing volume or amount fields—at lines 200-210【https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L200-L210】.
predict_batch() requires list-based inputs: df_list, x_timestamp_list, and y_timestamp_list, where each list contains one element per financial instrument. The method executes the same validation logic inside a Python for loop (lines 998-1012【https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L998-L1012】), aborting early if any individual series fails the integrity checks.
Tensor Shaping and Batch Dimensions
The internal preprocessing pipeline diverges significantly after validation. For single-series inference, predict() normalizes the data using mean-standard deviation scaling and clips values, then explicitly adds a batch dimension via np.newaxis to create a tensor of shape (1, seq_len, feat) (lines 255-259【https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L255-L259】).
In contrast, predict_batch() normalizes each DataFrame independently using its own statistics, then stacks all series into a single three-dimensional numpy array of shape (B, seq_len, feat) where B represents the batch size (lines 1048-1051【https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L1048-L1051】). This allows the model to process multiple instruments simultaneously in one forward pass.
Core Inference and Return Types
Both methods delegate to self.generate(), which internally calls auto_regressive_inference. However, predict() invokes this at line 261 with the single-series tensor, while predict_batch() passes the batched tensor at line 1060【https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L1060】.
The return structures reflect their input paradigms:
predict()returns a singlepd.DataFramewith columnsopen, high, low, close, volume, amount(lines 267-270【https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L267-L270】)predict_batch()returns aList[pd.DataFrame]preserving the input order, with each DataFrame de-normalized using its specific scaling statistics (lines 1066-1070【https://github.com/shiyu-coder/Kronos/blob/master/model/kronos.py#L1066-L1070】)
Implementation Details in model/kronos.py
Examining the source architecture reveals that both methods share common preprocessing utilities but handle data aggregation differently. The timestamp creation via calc_time_stamps occurs at lines 237-242 for single predictions and lines 1017-1022 for batch operations.
Normalization is consistently applied per-series rather than globally—critical for financial data where different instruments exhibit varying volatility scales. In predict_batch(), this individual normalization occurs inside the preprocessing loop at lines 1030-1035 before the stacking operation at line 1048.
When to Use predict() vs predict_batch()
Use predict() when working with individual assets in interactive Jupyter notebooks, debugging model behavior, or when forecasting portfolios where batch size consistently equals one. The method avoids the list-wrapping overhead and returns results directly without list unpacking.
Use predict_batch() for production forecasting across large universes of stocks or when running historical backtests across multiple instruments. The batched tensor approach eliminates Python loop overhead and maximizes GPU/TPU utilization through vectorized operations.
Code Examples
Single Asset Prediction with predict()
import pandas as pd
from model.kronos import Kronos, KronosTokenizer, KronosPredictor
# Assume `model` and `tokenizer` are already loaded/trained instances
predictor = KronosPredictor(model, tokenizer)
# Load a single DataFrame (must contain open, high, low, close; volume/amount optional)
df = pd.read_parquet("data/stock_A.parquet")
# Historical timestamps (same length as df) and future timestamps for the forecast
x_ts = pd.date_range(start="2023-01-01", periods=len(df), freq="D")
y_ts = pd.date_range(start=x_ts[-1] + pd.Timedelta(days=1), periods=30, freq="D")
# Predict the next 30 days
forecast_df = predictor.predict(df, x_ts, y_ts, pred_len=30)
print(forecast_df.head())
Key method invoked: KronosPredictor.predict() at lines 199-277 in model/kronos.py.
Batch Portfolio Prediction with predict_batch()
import pandas as pd
from model.kronos import Kronos, KronosTokenizer, KronosPredictor
predictor = KronosPredictor(model, tokenizer)
# Prepare several DataFrames (e.g., different stocks)
dfs = [pd.read_parquet(f"data/stock_{sym}.parquet") for sym in ["A", "B", "C"]]
# Corresponding timestamp lists
x_ts_list = [pd.date_range(start="2023-01-01", periods=len(df), freq="D") for df in dfs]
y_ts_list = [pd.date_range(start=xt[-1] + pd.Timedelta(days=1), periods=30, freq="D")
for xt in x_ts_list]
# Batch prediction
forecasts = predictor.predict_batch(dfs, x_ts_list, y_ts_list, pred_len=30)
for sym, df_pred in zip(["A", "B", "C"], forecasts):
print(f"=== Forecast for {sym} ===")
print(df_pred.head())
Key method invoked: KronosPredictor.predict_batch() at lines 990-1070 in model/kronos.py.
Summary
- Input Structure:
predict()takes single DataFrame/timestamp arguments;predict_batch()requires lists of equal length - Tensor Dimensions: Single-series uses shape (1, seq_len, feat); batch processing stacks to (B, seq_len, feat)
- Validation: Both validate per-series, but batch processing validates iteratively with early stopping on errors
- Normalization: Each series uses independent mean/std statistics in both methods
- Return Types:
predict()returnspd.DataFrame;predict_batch()returnsList[pd.DataFrame] - Performance: Batch processing reduces inference overhead for multiple instruments through vectorized forward passes
Frequently Asked Questions
Can predict_batch() be used for a single DataFrame?
Yes, but it requires wrapping the single DataFrame and timestamps in lists (e.g., [df], [x_ts]) and returns a list containing one DataFrame. This adds unnecessary list-wrapping overhead compared to predict(), which is optimized for single-series inference.
Does predict_batch() share model parameters across instruments?
Yes. Both methods call self.generate(), which utilizes the same underlying auto_regressive_inference function. In batch mode (line 1060), this function processes the entire stacked tensor simultaneously, sharing model weights across all instruments in the batch while maintaining separate normalization statistics per series.
What happens if one series in predict_batch() fails validation?
The method aborts early if any series fails validation during the preprocessing loop (lines 998-1012). This ensures data integrity across the entire batch rather than returning partial results for valid series while silently failing on others.
Are normalization statistics shared between batch items?
No. Each DataFrame is normalized independently using its own mean and standard deviation before tensor stacking (lines 1030-1035), and de-normalized individually using those same statistics after inference (lines 1066-1070). This preserves the distributional characteristics of each financial instrument.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →