koreallmdev's picture
Publish deepseek70b-qlora-7816-conditional-stable-adapter
378b4ec verified
Raw
History Blame Contribute Delete
4.92 kB
#!/usr/bin/env python3
from __future__ import annotations
import argparse, json, re, time
from collections import Counter
from pathlib import Path
import torch
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
BLOCKING={"EMPTY","REFUSAL","PROMPT_COPY","REPETITION","HAN_OUTSIDE_CODE"}
def norm(x): return re.sub(r"\s+"," ",x).strip()
def get_flags(text,prompt):
if not text.strip(): return ["EMPTY"]
out=[]; low=text.lower()
if any(x in low for x in ("λ„μ™€λ“œλ¦΄ 수 μ—†","μ œκ³΅ν•  수 μ—†","λ‹΅λ³€ν•  수 μ—†","i cannot","i can't")): out.append("REFUSAL")
p=norm(prompt).lower(); g=norm(text).lower()
if len(p)>=80 and g.startswith(p[:120]): out.append("PROMPT_COPY")
lines=[norm(x) for x in text.splitlines() if len(norm(x))>=12]
if any(v>=3 for v in Counter(lines).values()) or re.search(r"(.)\1{19,}",text,re.S): out.append("REPETITION")
outside=re.sub(r"```.*?```","",text,flags=re.S)
if re.search(r"[\u3400-\u4DBF\u4E00-\u9FFF]",outside): out.append("HAN_OUTSIDE_CODE")
return sorted(set(out))
def self_test():
assert "REPETITION" in get_flags("same sentence\nsame sentence\nsame sentence\n","x")
assert "REPETITION" not in get_flags("one sentence\ntwo sentence\nthree sentence\n","x")
print("SELF_TEST=PASS"); return 0
def main():
ap=argparse.ArgumentParser()
ap.add_argument("--prompt",default=""); ap.add_argument("--prompt-file",default="")
ap.add_argument("--release-dir",default="/home/saul9523/dgx_ai_factory/releases/deepseek70b_qlora_conditional_stable_current"); ap.add_argument("--model-dir",default="/home/saul9523/dgx_ai_factory/models/deepseek_r1_distill_llama_70b_hf")
ap.add_argument("--adapter-dir",default=""); ap.add_argument("--json-output",action="store_true")
ap.add_argument("--self-test",action="store_true"); ap.add_argument("--max-input-tokens",type=int,default=896)
a=ap.parse_args()
if a.self_test: return self_test()
release=Path(a.release_dir).resolve(); model_dir=Path(a.model_dir).resolve()
adapter=Path(a.adapter_dir).resolve() if a.adapter_dir else release/"adapter"
prompt=a.prompt or (Path(a.prompt_file).read_text(encoding="utf-8") if a.prompt_file else "")
if not prompt.strip(): raise SystemExit("[FATAL] --prompt λ˜λŠ” --prompt-file ν•„μš”")
policy=json.loads((release/"runtime_policy.json").read_text(encoding="utf-8"))
tok=AutoTokenizer.from_pretrained(model_dir,use_fast=True,trust_remote_code=True)
if tok.eos_token_id is None: tok.eos_token_id=128001
if tok.pad_token_id is None: tok.pad_token_id=tok.eos_token_id
tok.padding_side="left"
q=BitsAndBytesConfig(load_in_4bit=True,bnb_4bit_quant_type="nf4",bnb_4bit_use_double_quant=True,bnb_4bit_compute_dtype=torch.bfloat16)
base=AutoModelForCausalLM.from_pretrained(model_dir,quantization_config=q,dtype=torch.bfloat16,device_map={"":0},low_cpu_mem_usage=True,trust_remote_code=True,attn_implementation="sdpa")
base.config.use_cache=True
model=PeftModel.from_pretrained(base,adapter,is_trainable=False); model.eval()
def generate(settings):
try: rendered=tok.apply_chat_template([{"role":"user","content":prompt}],tokenize=False,add_generation_prompt=True)
except Exception: rendered=prompt
enc=tok(rendered,return_tensors="pt",truncation=True,max_length=a.max_input_tokens,add_special_tokens=True)
dev=next(model.parameters()).device; enc={k:v.to(dev) for k,v in enc.items()}
started=time.perf_counter()
with torch.inference_mode():
out=model.generate(**enc,do_sample=bool(settings.get("do_sample",False)),max_new_tokens=int(settings.get("max_new_tokens",192)),repetition_penalty=float(settings.get("repetition_penalty",1.08)),no_repeat_ngram_size=int(settings.get("no_repeat_ngram_size",8)),use_cache=True,pad_token_id=tok.pad_token_id,eos_token_id=tok.eos_token_id)
ids=out[0,enc["input_ids"].shape[1]:]
text=tok.decode(ids,skip_special_tokens=True).strip()
return {"text":text,"tokens":int(ids.numel()),"seconds":round(time.perf_counter()-started,4),"flags":get_flags(text,prompt)}
attempts=[generate(policy["primary"])]
if set(attempts[-1]["flags"]) & BLOCKING and policy.get("retry",{}).get("enabled",True):
retry=dict(policy["retry"]); retry.pop("enabled",None); retry.pop("max_retries",None); attempts.append(generate(retry))
chosen=attempts[-1]; status="PASS" if not(set(chosen["flags"])&BLOCKING) else "FAIL"
result={"status":status,"attempt_count":len(attempts),"selected_flags":chosen["flags"],"generation":chosen["text"],"attempts":attempts}
print(json.dumps(result,ensure_ascii=False,indent=2) if a.json_output else chosen["text"]+f"\n\n[status={status} attempts={len(attempts)} flags={chosen['flags']}]")
return 0 if status=="PASS" else 20
if __name__=="__main__": raise SystemExit(main())