import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
MODEL_NAME = "ahammad115566/qwen-smeft"
RESPONSE_PREFIX = "\n### Response:\n"
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True,
)
tokenizer = AutoTokenizer.from_pretrained(
MODEL_NAME,
trust_remote_code=True,
)
model = AutoModelForCausalLM.from_pretrained(
MODEL_NAME,
quantization_config=bnb_config,
device_map={"": 0},
trust_remote_code=True,
)
model.eval()
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
def build_prompt(instruction: str) -> str:
"""Construct the prompt using the format used during fine-tuning."""
return f"\n### Instruction:\n{instruction}\n{RESPONSE_PREFIX}"
@torch.inference_mode()
def ask(instruction: str) -> str:
"""Generate a response to an SMEFT-related instruction."""
prompt = build_prompt(instruction)
inputs = tokenizer(
prompt,
return_tensors="pt",
add_special_tokens=True,
).to(model.device)
output_ids = model.generate(
**inputs,
max_new_tokens=2048,
do_sample=False,
repetition_penalty=1.1,
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.pad_token_id,
)
new_tokens = output_ids[0][inputs["input_ids"].shape[-1]:]
return tokenizer.decode(
new_tokens,
skip_special_tokens=True,
).strip()
instruction = """
Which SMEFT operators modify EWPO?
"""
response = ask(instruction)
print(response)