LiveHouse-TS / scripts /validate_model.py
ziyuzhou02's picture
Deploy GitHub main 3feb6cda1511
e317359 verified
Raw
History Blame Contribute Delete
5.75 kB
#!/usr/bin/env python3
"""Validate an external user model wrapper for LiveHouse-TS."""
from __future__ import annotations
import argparse
import importlib
import sys
from pathlib import Path
import numpy as np
import pandas as pd
from gluonts.dataset.common import ListDataset
from gluonts.model.forecast import Forecast, QuantileForecast
DEFAULT_QUANTILES = [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9]
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--model-class",
required=True,
help="Full import path to the model class, e.g. user_models.my_user_model.MyUserModel",
)
parser.add_argument(
"--checkpoint",
default=None,
help="Path to the model checkpoint / weights file (optional)",
)
parser.add_argument(
"--prediction-length",
type=int,
default=24,
help="Prediction length to test with (default: 24)",
)
parser.add_argument(
"--freq",
default="1H",
help="Time series frequency to test with (default: 1H)",
)
args = parser.parse_args()
print("=== TSFM External Model Validator ===")
print(f"Model Class: {args.model_class}")
print(f"Checkpoint: {args.checkpoint}")
print(f"Pred Length: {args.prediction_length}")
print(f"Frequency: {args.freq}\n")
# Resolve import path
# Add root and space to sys.path
repo_root = Path(__file__).resolve().parents[1]
space_path = repo_root / "space"
if str(space_path) not in sys.path:
sys.path.insert(0, str(space_path))
if str(repo_root / "src") not in sys.path:
sys.path.insert(0, str(repo_root / "src"))
print("[Step 1] Loading model module...")
try:
module_name, class_name = args.model_class.rsplit(".", 1)
module = importlib.import_module(module_name)
model_class = getattr(module, class_name)
print(f" Successfully loaded {class_name} from {module_name}")
except Exception as e:
print(f" [ERROR] Failed to import model class: {e}")
sys.exit(1)
# Instantiate model
print("[Step 2] Instantiating model...")
try:
predictor = model_class(
prediction_length=args.prediction_length,
checkpoint_path=args.checkpoint,
quantile_levels=DEFAULT_QUANTILES,
)
print(" Successfully instantiated the model predictor.")
except TypeError as te:
print(f" [ERROR] Constructor signature mismatch: {te}")
print(" Note: Your constructor MUST accept (prediction_length: int, checkpoint_path: str | None, quantile_levels: list[float] | None)")
sys.exit(1)
except Exception as e:
print(f" [ERROR] Failed to instantiate model: {e}")
sys.exit(1)
# Create dummy data
print("[Step 3] Preparing dummy evaluation dataset...")
# 100 timesteps of dummy values
history_len = 100
dummy_target = np.sin(np.arange(history_len) * 0.1) + np.random.normal(0, 0.1, history_len)
# ListDataset expects target to be float32
entry = {
"item_id": "dummy_ts_0",
"start": pd.Period("2026-06-01 00:00", freq=args.freq),
"target": dummy_target.astype(np.float32),
}
dataset = ListDataset([entry], freq=args.freq)
# Run prediction
print("[Step 4] Running model predictions...")
try:
forecast_iter = predictor.predict(dataset)
forecasts = list(forecast_iter)
except Exception as e:
print(f" [ERROR] predict() call failed: {e}")
sys.exit(1)
if not forecasts:
print(" [ERROR] Predictor returned an empty forecast list/iterator.")
sys.exit(1)
print(f" Successfully generated {len(forecasts)} forecasts.")
# Validate output format
print("[Step 5] Checking output forecast structure...")
fc = forecasts[0]
# Check if subclass of Forecast
if not isinstance(fc, Forecast):
print(f" [WARNING] Output item is type {type(fc)}, which does not inherit from gluonts.model.forecast.Forecast.")
else:
print(" Forecast inherits from gluonts.model.forecast.Forecast. [OK]")
# Check prediction length shape
try:
p50 = fc.quantile(0.5) if hasattr(fc, "quantile") else fc.mean
actual_len = len(p50)
if actual_len != args.prediction_length:
print(f" [ERROR] Prediction length mismatch: expected {args.prediction_length}, got {actual_len}.")
sys.exit(1)
print(f" Forecast length matches prediction length {args.prediction_length}. [OK]")
except Exception as e:
print(f" [ERROR] Failed to extract p50 / mean forecast: {e}")
sys.exit(1)
# Check quantiles if it's a QuantileForecast
if isinstance(fc, QuantileForecast) or hasattr(fc, "quantile"):
print(" Checking forecast quantiles...")
try:
for q in [0.1, 0.5, 0.9]:
q_vals = fc.quantile(q)
if np.isnan(q_vals).any() or np.isinf(q_vals).any():
print(f" [ERROR] Quantile {q} contains NaN or Inf values.")
sys.exit(1)
print(f" - Quantile {q} is valid (no NaN/Inf).")
print(" Quantile checks passed. [OK]")
except Exception as e:
print(f" [ERROR] Failed to query quantiles: {e}")
sys.exit(1)
else:
print(" [WARNING] Forecast object does not support quantiles (p10/p50/p90 visual bands will fallback to mean).")
print("\n=========================================")
print("🎉 SUCCESS: Model wrapper validation PASSED!")
print("=========================================")
if __name__ == "__main__":
main()