#!/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()