weather-llm-initial / scripts /test_hf_model.py
NagacharanVemula
Improve model card and add HF inference smoke test.
22e1f58
Raw
History Blame Contribute Delete
2.55 kB
#!/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())