Adonis3039's picture
Upload 19 files
f0630f0 verified
Raw History Blame Contribute Delete
2.82 kB
#!/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()