27 lines
1.2 KiB
Python
27 lines
1.2 KiB
Python
|
|
# -*- coding: utf-8 -*-
|
||
|
|
import torch, json
|
||
|
|
from transformers import AutoTokenizer, AutoModelForCausalLM
|
||
|
|
|
||
|
|
MODEL_DIR = r"/root/AI-Model-Training-Test/runs/qwen4b_sft_merged3"
|
||
|
|
tok = AutoTokenizer.from_pretrained(MODEL_DIR, trust_remote_code=True)
|
||
|
|
model = AutoModelForCausalLM.from_pretrained(MODEL_DIR, device_map="auto", torch_dtype=torch.bfloat16 if True else (torch.float16 if False else None), trust_remote_code=True)
|
||
|
|
if tok.pad_token is None: tok.pad_token = tok.eos_token
|
||
|
|
|
||
|
|
def chat_once(system_text: str, user_text: str, max_new_tokens=256):
|
||
|
|
msgs = [{"role":"system","content":system_text},
|
||
|
|
{"role":"user","content":user_text}]
|
||
|
|
x = tok.apply_chat_template(msgs, return_tensors="pt", add_generation_prompt=True).to(model.device)
|
||
|
|
with torch.no_grad():
|
||
|
|
y = model.generate(x, max_new_tokens=max_new_tokens, do_sample=False, eos_token_id=tok.eos_token_id)
|
||
|
|
print(tok.decode(y[0], skip_special_tokens=True))
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
sys = "You are a strict detector for sensitive entities. Output ONLY one JSON object."
|
||
|
|
while True:
|
||
|
|
try:
|
||
|
|
q = input("text> ").strip()
|
||
|
|
if not q: continue
|
||
|
|
chat_once(sys, q)
|
||
|
|
except (EOFError, KeyboardInterrupt):
|
||
|
|
break
|