#!/usr/bin/env python3 """Smoke-test inference for a Hugging Face model repo.""" from __future__ import annotations import argparse import sys import torch from transformers import AutoModelForCausalLM, AutoTokenizer def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser(description="Run a quick generation test against a Hub model.") p.add_argument("--repo_id", type=str, default="AuraWorxAI/weather-llm-initial") p.add_argument( "--prompt", type=str, default="Compare summer weather patterns in Arizona and Washington.", ) p.add_argument("--max_new_tokens", type=int, default=80) p.add_argument("--temperature", type=float, default=0.0) p.add_argument("--top_p", type=float, default=0.9) p.add_argument("--device", type=str, default="auto", choices=["auto", "cpu", "cuda"]) return p.parse_args() def resolve_device(device: str) -> str: if device == "auto": return "cuda" if torch.cuda.is_available() else "cpu" return device def main() -> int: args = parse_args() device = resolve_device(args.device) print(f"Loading repo: {args.repo_id}") print(f"Using device: {device}") try: tokenizer = AutoTokenizer.from_pretrained(args.repo_id) model = AutoModelForCausalLM.from_pretrained(args.repo_id, torch_dtype="auto") model.to(device) model.eval() inputs = tokenizer(args.prompt, return_tensors="pt").to(device) gen_kw: dict = { "max_new_tokens": args.max_new_tokens, "eos_token_id": tokenizer.eos_token_id, "pad_token_id": tokenizer.pad_token_id, } if args.temperature > 0: gen_kw["do_sample"] = True gen_kw["temperature"] = max(args.temperature, 1e-5) gen_kw["top_p"] = args.top_p else: gen_kw["do_sample"] = False with torch.inference_mode(): output_ids = model.generate(**inputs, **gen_kw) text = tokenizer.decode(output_ids[0], skip_special_tokens=True).strip() except Exception as exc: # pragma: no cover print(f"Inference smoke test failed: {exc}", file=sys.stderr) return 1 if not text: print("Inference smoke test failed: empty generation output.", file=sys.stderr) return 2 print("\n=== Prompt ===") print(args.prompt) print("\n=== Output ===") print(text) print("\nHF inference smoke test passed.") return 0 if __name__ == "__main__": raise SystemExit(main())