Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
144 changes: 144 additions & 0 deletions download_models.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,144 @@
"""
Kronos Model Download Script
Downloads tokenizer and Kronos-small model from HuggingFace Hub.

Exit codes:
0 -- download + smoke test succeeded
1 -- failure (missing submodule, import error, or download error)
"""
import argparse
import sys
from pathlib import Path

KRONOS_DIR = Path(__file__).resolve().parent
MODEL_DIR = KRONOS_DIR / "model"

DEFAULT_TOKENIZER = "NeoQuasar/Kronos-Tokenizer-base"
DEFAULT_MODEL = "NeoQuasar/Kronos-small"


def _check_model_package() -> None:
if not MODEL_DIR.is_dir():
print("[ERROR] Kronos model package not found at", MODEL_DIR)
print("Hint: Ensure the Kronos submodule is initialized:")
print(" git submodule update --init Kronos")
sys.exit(1)


def _import_model():
"""Import Kronos/KronosTokenizer with a human-readable failure path."""
_check_model_package()
if str(KRONOS_DIR) not in sys.path:
sys.path.insert(0, str(KRONOS_DIR))
try:
from model import Kronos, KronosTokenizer

return Kronos, KronosTokenizer
except ImportError as e:
print(f"[ERROR] Failed to import Kronos model: {e}")
print("Hint: Ensure dependencies are installed "
"(pip install -r Kronos/requirements.txt) and the Kronos")
print(" submodule is initialized (git submodule update --init).")
sys.exit(1)


def _device_label(obj) -> str:
"""Best-effort device string. Works for nn.Module, degrades gracefully."""
if hasattr(obj, "parameters"):
try:
return str(next(obj.parameters()).device)
except (StopIteration, AttributeError):
return "unknown"
return "N/A (non-torch backend)"


def _smoke_test(model, tokenizer) -> bool:
"""
Verify the downloaded artifacts load and produce a forward pass.

Kronos is a time-series foundation model: the tokenizer consumes a numeric
tensor and Kronos consumes token-id tensors, so the smoke test uses
synthetic tensors rather than text.
"""
print("\n[Verify] Running smoke test...")
try:
import torch

# Tokenizer: forward pass on a synthetic time-series batch.
d_in = tokenizer.embed.in_features
x = torch.randn(1, 64, d_in)
with torch.no_grad():
tokenizer(x)
print(" [OK] Tokenizer forward pass succeeded")

# Model: forward pass on random token ids.
s1_vocab = getattr(model, "s1_vocab_size", 4096)
s2_vocab = getattr(model, "s2_vocab_size", 4096)
s1_ids = torch.randint(0, s1_vocab, (1, 64))
s2_ids = torch.randint(0, s2_vocab, (1, 64))
with torch.no_grad():
model(s1_ids, s2_ids)
print(" [OK] Model forward pass succeeded")
return True
except Exception as e:
print(f" [FAIL] Smoke test failed: {e}")
return False


def download_models(
tokenizer_name: str = DEFAULT_TOKENIZER,
model_name: str = DEFAULT_MODEL,
) -> bool:
"""
Download the tokenizer and model, then verify them with a smoke test.

Returns True on success; on failure prints an error and returns False.
"""
print("=" * 50)
print("Kronos Model Downloader")
print("=" * 50)

model_cls, tokenizer_cls = _import_model()

print("\n[1/2] Downloading KronosTokenizer...")
print(f" Model: {tokenizer_name}")
try:
tokenizer = tokenizer_cls.from_pretrained(tokenizer_name)
except Exception as e:
print(f"[ERROR] Tokenizer download failed: {e}")
return False
print(" [OK] Tokenizer downloaded successfully")

print("\n[2/2] Downloading Kronos model...")
print(f" Model: {model_name}")
try:
model = model_cls.from_pretrained(model_name)
except Exception as e:
print(f"[ERROR] Model download failed: {e}")
return False
print(" [OK] Model downloaded successfully")

ok = _smoke_test(model, tokenizer)

print("\n" + "=" * 50)
if ok:
print("Download complete!")
else:
print("Download complete, but the smoke test failed -- the cached")
print("artifacts may be corrupt. Re-run to re-download them.")
print(f"Model device: {_device_label(model)}")
print(f"Tokenizer device: {_device_label(tokenizer)}")
print("=" * 50)
return ok


if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Download Kronos models from HuggingFace Hub"
)
parser.add_argument("--tokenizer", default=DEFAULT_TOKENIZER, help="HF tokenizer repo id")
parser.add_argument("--model", default=DEFAULT_MODEL, help="HF model repo id")
args = parser.parse_args()

success = download_models(args.tokenizer, args.model)
sys.exit(0 if success else 1)
221 changes: 141 additions & 80 deletions examples/prediction_example.py
Original file line number Diff line number Diff line change
@@ -1,80 +1,141 @@
import pandas as pd
import matplotlib.pyplot as plt
import sys
sys.path.append("../")
from model import Kronos, KronosTokenizer, KronosPredictor


def plot_prediction(kline_df, pred_df):
pred_df.index = kline_df.index[-pred_df.shape[0]:]
sr_close = kline_df['close']
sr_pred_close = pred_df['close']
sr_close.name = 'Ground Truth'
sr_pred_close.name = "Prediction"

sr_volume = kline_df['volume']
sr_pred_volume = pred_df['volume']
sr_volume.name = 'Ground Truth'
sr_pred_volume.name = "Prediction"

close_df = pd.concat([sr_close, sr_pred_close], axis=1)
volume_df = pd.concat([sr_volume, sr_pred_volume], axis=1)

fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(8, 6), sharex=True)

ax1.plot(close_df['Ground Truth'], label='Ground Truth', color='blue', linewidth=1.5)
ax1.plot(close_df['Prediction'], label='Prediction', color='red', linewidth=1.5)
ax1.set_ylabel('Close Price', fontsize=14)
ax1.legend(loc='lower left', fontsize=12)
ax1.grid(True)

ax2.plot(volume_df['Ground Truth'], label='Ground Truth', color='blue', linewidth=1.5)
ax2.plot(volume_df['Prediction'], label='Prediction', color='red', linewidth=1.5)
ax2.set_ylabel('Volume', fontsize=14)
ax2.legend(loc='upper left', fontsize=12)
ax2.grid(True)

plt.tight_layout()
plt.show()


# 1. Load Model and Tokenizer
tokenizer = KronosTokenizer.from_pretrained("NeoQuasar/Kronos-Tokenizer-base")
model = Kronos.from_pretrained("NeoQuasar/Kronos-small")

# 2. Instantiate Predictor
predictor = KronosPredictor(model, tokenizer, max_context=512)

# 3. Prepare Data
df = pd.read_csv("./data/XSHG_5min_600977.csv")
df['timestamps'] = pd.to_datetime(df['timestamps'])

lookback = 400
pred_len = 120

x_df = df.loc[:lookback-1, ['open', 'high', 'low', 'close', 'volume', 'amount']]
x_timestamp = df.loc[:lookback-1, 'timestamps']
y_timestamp = df.loc[lookback:lookback+pred_len-1, 'timestamps']

# 4. Make Prediction
pred_df = predictor.predict(
df=x_df,
x_timestamp=x_timestamp,
y_timestamp=y_timestamp,
pred_len=pred_len,
T=1.0,
top_p=0.9,
sample_count=1,
verbose=True
)

# 5. Visualize Results
print("Forecasted Data Head:")
print(pred_df.head())

# Combine historical and forecasted data for plotting
kline_df = df.loc[:lookback+pred_len-1]

# visualize
plot_prediction(kline_df, pred_df)

import argparse
import sys
from pathlib import Path

import pandas as pd

# Anchor the project root on this file's location so the example can be run
# from any working directory (e.g. `python examples/prediction_example.py`).
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from model import Kronos, KronosPredictor, KronosTokenizer

REQUIRED_COLUMNS = ['open', 'high', 'low', 'close', 'volume', 'amount', 'timestamps']


def plot_prediction(kline_df, pred_df):
import matplotlib.pyplot as plt

pred_df.index = kline_df.index[-pred_df.shape[0]:]
sr_close = kline_df['close']
sr_pred_close = pred_df['close']
sr_close.name = 'Ground Truth'
sr_pred_close.name = "Prediction"

sr_volume = kline_df['volume']
sr_pred_volume = pred_df['volume']
sr_volume.name = 'Ground Truth'
sr_pred_volume.name = "Prediction"

close_df = pd.concat([sr_close, sr_pred_close], axis=1)
volume_df = pd.concat([sr_volume, sr_pred_volume], axis=1)

fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(8, 6), sharex=True)

ax1.plot(close_df['Ground Truth'], label='Ground Truth', color='blue', linewidth=1.5)
ax1.plot(close_df['Prediction'], label='Prediction', color='red', linewidth=1.5)
ax1.set_ylabel('Close Price', fontsize=14)
ax1.legend(loc='lower left', fontsize=12)
ax1.grid(True)

ax2.plot(volume_df['Ground Truth'], label='Ground Truth', color='blue', linewidth=1.5)
ax2.plot(volume_df['Prediction'], label='Prediction', color='red', linewidth=1.5)
ax2.set_ylabel('Volume', fontsize=14)
ax2.legend(loc='upper left', fontsize=12)
ax2.grid(True)

plt.tight_layout()
plt.show()


def parse_args():
parser = argparse.ArgumentParser(
description="Run Kronos inference on an OHLCV CSV file.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument("csv_path", type=Path,
help="Path to a CSV with columns: " + ", ".join(REQUIRED_COLUMNS))
parser.add_argument("--lookback", type=int, default=400,
help="Number of historical bars to condition on.")
parser.add_argument("--pred_len", type=int, default=120,
help="Number of bars to forecast.")
parser.add_argument("--max_context", type=int, default=512,
help="Maximum context length fed to the model.")
parser.add_argument("--device", type=str, default=None,
help="Inference device (defaults to KRONOS_DEVICE, then auto-detect).")
parser.add_argument("--top_k", type=int, default=0,
help="Top-k sampling threshold (0 disables).")
parser.add_argument("--top_p", type=float, default=0.9,
help="Nucleus sampling threshold.")
parser.add_argument("--T", type=float, default=1.0,
help="Sampling temperature.")
parser.add_argument("--sample_count", type=int, default=1,
help="Parallel samples per series (averaged).")
parser.add_argument("--no-show", action="store_true",
help="Skip the matplotlib plot.")
return parser.parse_args()


def main():
args = parse_args()

if not args.csv_path.is_file():
raise FileNotFoundError(f"CSV file not found: {args.csv_path}")

df = pd.read_csv(args.csv_path)
missing = [c for c in REQUIRED_COLUMNS if c not in df.columns]
if missing:
raise ValueError(
f"CSV at {args.csv_path} is missing required columns: {missing}. "
f"Expected columns: {REQUIRED_COLUMNS}"
)

df['timestamps'] = pd.to_datetime(df['timestamps'])

if args.lookback < 1 or args.pred_len < 1:
raise ValueError("lookback and pred_len must both be >= 1.")
if args.lookback + args.pred_len > len(df):
raise ValueError(
f"lookback + pred_len ({args.lookback + args.pred_len}) exceeds the "
f"number of rows in the CSV ({len(df)})."
)

# 1. Load Model and Tokenizer
tokenizer = KronosTokenizer.from_pretrained("NeoQuasar/Kronos-Tokenizer-base")
model = Kronos.from_pretrained("NeoQuasar/Kronos-small")

# 2. Instantiate Predictor
predictor = KronosPredictor(model, tokenizer, device=args.device, max_context=args.max_context)

# 3. Prepare Data
lookback = args.lookback
pred_len = args.pred_len

x_df = df.loc[:lookback - 1, ['open', 'high', 'low', 'close', 'volume', 'amount']]
x_timestamp = df.loc[:lookback - 1, 'timestamps']
y_timestamp = df.loc[lookback:lookback + pred_len - 1, 'timestamps']

# 4. Make Prediction
pred_df = predictor.predict(
df=x_df,
x_timestamp=x_timestamp,
y_timestamp=y_timestamp,
pred_len=pred_len,
T=args.T,
top_k=args.top_k,
top_p=args.top_p,
sample_count=args.sample_count,
verbose=True
)

# 5. Visualize Results
print("Forecasted Data Head:")
print(pred_df.head())

if not args.no_show:
# Combine historical and forecasted data for plotting
kline_df = df.loc[:lookback + pred_len - 1]
plot_prediction(kline_df, pred_df)


if __name__ == '__main__':
main()
Loading