#!/usr/bin/env python3 """Load an EviSuff adapter stack and optionally run a smoke-test prompt.""" from __future__ import annotations import argparse from pathlib import Path ADAPTERS = ("answer-sft", "no-gate", "full-evisuff") BASE_MODEL = "Qwen/Qwen3-8B" def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--adapter", choices=ADAPTERS, default="full-evisuff") parser.add_argument( "--repo-id", help="Hugging Face model repository. Omit to use the local repository clone.", ) parser.add_argument("--base-model", default=BASE_MODEL) parser.add_argument("--revision", default=None, help="Optional base-model revision.") parser.add_argument("--prompt", help="Optional prompt for a short generation smoke test.") parser.add_argument("--max-new-tokens", type=int, default=128) return parser.parse_args() def adapter_location(repo_id: str | None, adapter: str) -> str: if repo_id: return f"{repo_id}/{adapter}" return str(Path(__file__).resolve().parents[1] / adapter) def load_stack(args: argparse.Namespace): import torch from peft import PeftModel from transformers import AutoModelForCausalLM, AutoTokenizer dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32 base = AutoModelForCausalLM.from_pretrained( args.base_model, revision=args.revision, torch_dtype=dtype, device_map="auto", ) tokenizer = AutoTokenizer.from_pretrained(args.base_model, revision=args.revision) if args.repo_id: answer_model = PeftModel.from_pretrained( base, args.repo_id, subfolder="answer-sft", ) else: answer_model = PeftModel.from_pretrained( base, adapter_location(None, "answer-sft"), ) if args.adapter == "answer-sft": return answer_model, tokenizer merged_answer = answer_model.merge_and_unload() if args.repo_id: model = PeftModel.from_pretrained( merged_answer, args.repo_id, subfolder=args.adapter, ) else: model = PeftModel.from_pretrained( merged_answer, adapter_location(None, args.adapter), ) return model, tokenizer def main() -> None: args = parse_args() model, tokenizer = load_stack(args) print(f"Loaded {args.adapter} with the required adapter stack.") if not args.prompt: return inputs = tokenizer(args.prompt, return_tensors="pt").to(model.device) output = model.generate(**inputs, max_new_tokens=args.max_new_tokens, do_sample=False) print(tokenizer.decode(output[0], skip_special_tokens=True)) if __name__ == "__main__": main()