import os, json, argparse HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) os.environ["HF_HOME"] = os.path.join(HERE, "hf_cache") import torch from transformers import AutoModelForCausalLM, AutoTokenizer def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", default="Qwen/Qwen3-0.6B") ap.add_argument("--adapter", default=None) ap.add_argument("--eval", default=os.path.join(HERE, "eval", "chat_eval.jsonl")) ap.add_argument("--out", default=os.path.join(HERE, "eval", "baseline_chat.json")) a = ap.parse_args() device = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu") tok = AutoTokenizer.from_pretrained(a.model) model = AutoModelForCausalLM.from_pretrained(a.model, dtype=torch.float32).to(device) if a.adapter: from peft import PeftModel model = PeftModel.from_pretrained(model, a.adapter).to(device) model.eval() cases = [json.loads(l) for l in open(a.eval) if l.strip()] results = [] for c in cases: msgs = [{"role": "user", "content": c["query"]}] text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True, enable_thinking=False) ids = tok(text, return_tensors="pt").to(device) out = model.generate(**ids, max_new_tokens=256, do_sample=False) reply = tok.decode(out[0][ids.input_ids.shape[1]:], skip_special_tokens=True) results.append({"id": c["id"], "kind": c["kind"], "query": c["query"], "reply": reply}) print(f"[{c['id']}] {c['kind']}: {reply[:80].strip()}...") json.dump({"model": a.model, "adapter": a.adapter, "n": len(cases), "results": results}, open(a.out, "w"), indent=2) print(f"\nsaved {len(cases)} chat outputs -> {a.out}") if __name__ == "__main__": main()