初始化项目,由ModelHub XC社区提供模型
Model: karthik-2905/AL1-model-B Source: Original Platform
This commit is contained in:
47
mlb/build_sft.py
Normal file
47
mlb/build_sft.py
Normal file
@@ -0,0 +1,47 @@
|
||||
import os, json, random
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
TOOLS_F = os.path.join(HERE, "data", "sft_tools.jsonl")
|
||||
CHAT_F = os.path.join(HERE, "data", "sft_chat.jsonl")
|
||||
EVAL_TOOLS = os.path.join(HERE, "eval", "tools_eval.jsonl")
|
||||
EVAL_CHAT = os.path.join(HERE, "eval", "chat_eval.jsonl")
|
||||
TRAIN_OUT = os.path.join(HERE, "data", "sft_train.jsonl")
|
||||
VAL_OUT = os.path.join(HERE, "data", "sft_val.jsonl")
|
||||
VAL_FRAC = 0.05
|
||||
random.seed(0)
|
||||
|
||||
def load(path):
|
||||
return [json.loads(l) for l in open(path) if l.strip()]
|
||||
|
||||
def main():
|
||||
eval_qs = set()
|
||||
for f in (EVAL_TOOLS, EVAL_CHAT):
|
||||
for c in load(f):
|
||||
eval_qs.add(c["query"].strip())
|
||||
|
||||
rows = load(TOOLS_F) + load(CHAT_F)
|
||||
kept, dropped = [], 0
|
||||
for r in rows:
|
||||
if r["messages"][0]["content"].strip() in eval_qs:
|
||||
dropped += 1
|
||||
continue
|
||||
kept.append(r)
|
||||
|
||||
random.shuffle(kept)
|
||||
n_val = int(len(kept) * VAL_FRAC)
|
||||
val, train = kept[:n_val], kept[n_val:]
|
||||
|
||||
for path, split in ((TRAIN_OUT, train), (VAL_OUT, val)):
|
||||
with open(path, "w") as f:
|
||||
for r in split:
|
||||
f.write(json.dumps(r) + "\n")
|
||||
|
||||
def has_tools(r): return "tools" in r
|
||||
print(f"leakage-dropped: {dropped}")
|
||||
print(f"train: {len(train)} (tool={sum(has_tools(r) for r in train)}, "
|
||||
f"chat={sum(not has_tools(r) for r in train)})")
|
||||
print(f"val: {len(val)} (tool={sum(has_tools(r) for r in val)}, "
|
||||
f"chat={sum(not has_tools(r) for r in val)})")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
38
mlb/check_sft_format.py
Normal file
38
mlb/check_sft_format.py
Normal file
@@ -0,0 +1,38 @@
|
||||
import os, json
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
os.environ["HF_HOME"] = os.path.join(HERE, "hf_cache")
|
||||
import sys; sys.path.insert(0, os.path.join(HERE, "mlb"))
|
||||
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
MODEL = "Qwen/Qwen3-0.6B"
|
||||
DATA = os.path.join(HERE, "data", "sft_tools.jsonl")
|
||||
|
||||
def main():
|
||||
tok = AutoTokenizer.from_pretrained(MODEL)
|
||||
rows = [json.loads(l) for l in open(DATA) if l.strip()]
|
||||
print(f"loaded {len(rows)} rows")
|
||||
|
||||
lengths, bad = [], 0
|
||||
for r in rows:
|
||||
text = tok.apply_chat_template(
|
||||
r["messages"], tools=r["tools"], tokenize=False,
|
||||
add_generation_prompt=False, enable_thinking=False)
|
||||
if "<tool_call>" not in text:
|
||||
bad += 1
|
||||
lengths.append(len(tok(text).input_ids))
|
||||
|
||||
lengths.sort()
|
||||
print("tool_call missing in rendered text:", bad)
|
||||
print("token length min/mean/max:",
|
||||
lengths[0], round(sum(lengths) / len(lengths), 1), lengths[-1])
|
||||
print("p95 length:", lengths[int(0.95 * len(lengths))])
|
||||
|
||||
print("\n----- SAMPLE RENDERED TRAINING STRING -----\n")
|
||||
s = tok.apply_chat_template(rows[0]["messages"], tools=rows[0]["tools"],
|
||||
tokenize=False, add_generation_prompt=False,
|
||||
enable_thinking=False)
|
||||
print(s)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
43
mlb/eval_chat.py
Normal file
43
mlb/eval_chat.py
Normal file
@@ -0,0 +1,43 @@
|
||||
import os, json, argparse
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
os.environ["HF_HOME"] = os.path.join(HERE, "hf_cache")
|
||||
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--model", default="Qwen/Qwen3-0.6B")
|
||||
ap.add_argument("--adapter", default=None)
|
||||
ap.add_argument("--eval", default=os.path.join(HERE, "eval", "chat_eval.jsonl"))
|
||||
ap.add_argument("--out", default=os.path.join(HERE, "eval", "baseline_chat.json"))
|
||||
a = ap.parse_args()
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu")
|
||||
tok = AutoTokenizer.from_pretrained(a.model)
|
||||
model = AutoModelForCausalLM.from_pretrained(a.model, dtype=torch.float32).to(device)
|
||||
if a.adapter:
|
||||
from peft import PeftModel
|
||||
model = PeftModel.from_pretrained(model, a.adapter).to(device)
|
||||
model.eval()
|
||||
|
||||
cases = [json.loads(l) for l in open(a.eval) if l.strip()]
|
||||
results = []
|
||||
|
||||
for c in cases:
|
||||
msgs = [{"role": "user", "content": c["query"]}]
|
||||
text = tok.apply_chat_template(msgs, tokenize=False,
|
||||
add_generation_prompt=True, enable_thinking=False)
|
||||
ids = tok(text, return_tensors="pt").to(device)
|
||||
out = model.generate(**ids, max_new_tokens=256, do_sample=False)
|
||||
reply = tok.decode(out[0][ids.input_ids.shape[1]:], skip_special_tokens=True)
|
||||
results.append({"id": c["id"], "kind": c["kind"],
|
||||
"query": c["query"], "reply": reply})
|
||||
print(f"[{c['id']}] {c['kind']}: {reply[:80].strip()}...")
|
||||
|
||||
json.dump({"model": a.model, "adapter": a.adapter, "n": len(cases), "results": results},
|
||||
open(a.out, "w"), indent=2)
|
||||
print(f"\nsaved {len(cases)} chat outputs -> {a.out}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
58
mlb/eval_tools.py
Normal file
58
mlb/eval_tools.py
Normal file
@@ -0,0 +1,58 @@
|
||||
import os, json, argparse
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
os.environ["HF_HOME"] = os.path.join(HERE, "hf_cache")
|
||||
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
from tools_common import TOOLS, parse_tool_calls
|
||||
|
||||
def norm_args(d):
|
||||
return {k: (v.strip() if isinstance(v, str) else v) for k, v in (d or {}).items()}
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--model", default="Qwen/Qwen3-0.6B")
|
||||
ap.add_argument("--adapter", default=None)
|
||||
ap.add_argument("--eval", default=os.path.join(HERE, "eval", "tools_eval.py"))
|
||||
ap.add_argument("--out", default=os.path.join(HERE, "eval", "baseline_tools.json"))
|
||||
a = ap.parse_args()
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu")
|
||||
tok = AutoTokenizer.from_pretrained(a.model)
|
||||
model = AutoModelForCausalLM.from_pretrained(a.model, dtype=torch.float32).to(device)
|
||||
if a.adapter:
|
||||
from peft import PeftModel
|
||||
model = PeftModel.from_pretrained(model, a.adapter).to(device)
|
||||
model.eval()
|
||||
|
||||
cases = [json.loads(l) for l in open(a.eval) if l.strip()]
|
||||
results, name_hits, full_hits = [], 0, 0
|
||||
|
||||
for c in cases:
|
||||
msgs = [{"role": "user", "content": c["query"]}]
|
||||
text = tok.apply_chat_template(msgs, tools=TOOLS, tokenize=False,
|
||||
add_generation_prompt=True, enable_thinking=False)
|
||||
ids = tok(text, return_tensors="pt").to(device)
|
||||
out = model.generate(**ids, max_new_tokens=128, do_sample=False)
|
||||
reply = tok.decode(out[0][ids.input_ids.shape[1]:], skip_special_tokens=True)
|
||||
|
||||
calls = parse_tool_calls(reply)
|
||||
got = calls[0] if calls else None
|
||||
exp = c["expected"]
|
||||
name_ok = bool(got) and got.get("name") == exp["name"]
|
||||
full_ok = name_ok and norm_args(got.get("arguments")) == norm_args(exp["arguments"])
|
||||
name_hits += name_ok
|
||||
full_hits += full_ok
|
||||
results.append({"query": c["query"], "expected": exp,
|
||||
"got": got, "raw": reply,
|
||||
"name_ok": name_ok, "full_ok": full_ok})
|
||||
|
||||
n = len(cases)
|
||||
summary = {"model": a.model, "adapter": a.adapter, "n": n,
|
||||
"name_acc": round(name_hits / n, 3),
|
||||
"full_acc": round(full_hits / n, 3)}
|
||||
json.dump({"summary": summary, "results": results}, open(a.out, "w"), indent=2)
|
||||
print(json.dumps(summary, indent=2))
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
37
mlb/gen_sft_chat.py
Normal file
37
mlb/gen_sft_chat.py
Normal file
@@ -0,0 +1,37 @@
|
||||
import os, json, random
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
os.environ["HF_HOME"] = os.path.join(HERE, "hf_cache")
|
||||
|
||||
from datasets import load_dataset
|
||||
|
||||
OUT = os.path.join(HERE, "data", "sft_chat.jsonl")
|
||||
N = 500
|
||||
random.seed(0)
|
||||
|
||||
def main():
|
||||
ds = load_dataset("databricks/databricks-dolly-15k", split="train")
|
||||
idx = list(range(len(ds)))
|
||||
random.shuffle(idx)
|
||||
|
||||
rows = []
|
||||
for i in idx:
|
||||
r = ds[i]
|
||||
instr, ctx, resp = r["instruction"].strip(), r["context"].strip(), r["response"].strip()
|
||||
if not instr or not resp:
|
||||
continue
|
||||
if len(resp) > 1200: # keep seq length sane
|
||||
continue
|
||||
user = f"{instr}\n\n{ctx}" if ctx else instr
|
||||
rows.append({"messages": [{"role": "user", "content": user},
|
||||
{"role": "assistant", "content": resp}]})
|
||||
if len(rows) >= N:
|
||||
break
|
||||
|
||||
with open(OUT, "w") as f:
|
||||
for r in rows:
|
||||
f.write(json.dumps(r) + "\n")
|
||||
print(f"wrote {len(rows)} chat examples -> {OUT}")
|
||||
print("sample user:", rows[0]["messages"][0]["content"][:80])
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
90
mlb/gen_sft_tools.py
Normal file
90
mlb/gen_sft_tools.py
Normal file
@@ -0,0 +1,90 @@
|
||||
import os, json, random
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
import sys; sys.path.insert(0, os.path.join(HERE, "mlb"))
|
||||
from tools_common import TOOLS
|
||||
|
||||
OUT = os.path.join(HERE, "data", "sft_tools.jsonl")
|
||||
N = 1000
|
||||
random.seed(0)
|
||||
|
||||
FILES = ["config.py", "app.py", "main.py", "README.md", "utils.py", "src/utils.py",
|
||||
"tests/test_api.py", "server.js", "index.html", "data/train.csv", "Makefile",
|
||||
"requirements.txt", "model.py", "train.py", ".env.example", "docker-compose.yml"]
|
||||
DIRS = ["src", "tests", "data", ".", "configs", "scripts", "app", "lib", "models"]
|
||||
CMDS = ["pytest", "git status", "ls -la", "npm test", "python train.py",
|
||||
"pip install -r requirements.txt", "git log --oneline", "make build",
|
||||
"docker build .", "ruff check .", "python -m pytest tests/", "git diff"]
|
||||
EDITS = [("5000", "8080"), ("foo", "bar"), ("debug=False", "debug=True"),
|
||||
("localhost", "0.0.0.0"), ("v1", "v2"), ("GET", "POST"),
|
||||
("255", "128"), ("lr=1e-3", "lr=3e-4"), ("True", "False")]
|
||||
WRITES = [("hello.py", "print('hi')"), ("notes.txt", "hello world"),
|
||||
(".gitignore", "__pycache__/"), ("VERSION", "1.0.0"),
|
||||
("run.sh", "#!/bin/bash\npython main.py"), ("todo.md", "- ship it")]
|
||||
|
||||
READ_T = ["Show me the contents of {p}", "Read {p}", "Open {p}", "What's in {p}?",
|
||||
"Display {p}", "Print out {p} for me", "Can you show me {p}?"]
|
||||
LIST_T = ["List the files in {p}", "What's in the {p} directory?", "Show the entries in {p}",
|
||||
"List everything under {p}", "ls {p}", "Show me what's inside {p}"]
|
||||
WRITE_T = ["Create a file {p} containing {c}", "Make a new file {p} with the text {c}",
|
||||
"Write {c} to {p}", "Save {c} into a file called {p}"]
|
||||
EDIT_T = ["In {p} replace {a} with {b}", "Change {a} to {b} in {p}",
|
||||
"Update {p}: swap {a} for {b}", "In the file {p}, replace {a} with {b}",
|
||||
"Edit {p} and change {a} into {b}", "Modify {p} so {a} becomes {b}"]
|
||||
BASH_T = ["Run {c}", "Execute {c}", "Run the command {c}", "Please run {c}",
|
||||
"Can you run {c}?", "Kick off {c}", "Go ahead and run {c}"]
|
||||
|
||||
WEIGHTS = [("edit_file", 0.28), ("run_bash", 0.28), ("read_file", 0.16),
|
||||
("list_dir", 0.14), ("write_file", 0.14)]
|
||||
|
||||
def pick_tool():
|
||||
r, acc = random.random(), 0.0
|
||||
for name, w in WEIGHTS:
|
||||
acc += w
|
||||
if r <= acc:
|
||||
return name
|
||||
return WEIGHTS[-1][0]
|
||||
|
||||
def make_example():
|
||||
t = pick_tool()
|
||||
if t == "read_file":
|
||||
p = random.choice(FILES)
|
||||
q = random.choice(READ_T).format(p=p); args = {"path": p}
|
||||
elif t == "list_dir":
|
||||
p = random.choice(DIRS)
|
||||
q = random.choice(LIST_T).format(p=p); args = {"path": p}
|
||||
elif t == "write_file":
|
||||
p, c = random.choice(WRITES)
|
||||
q = random.choice(WRITE_T).format(p=p, c=c); args = {"path": p, "content": c}
|
||||
elif t == "edit_file":
|
||||
p = random.choice(FILES); a, b = random.choice(EDITS)
|
||||
q = random.choice(EDIT_T).format(p=p, a=a, b=b)
|
||||
args = {"path": p, "old_string": a, "new_string": b}
|
||||
else: # run_bash
|
||||
c = random.choice(CMDS)
|
||||
q = random.choice(BASH_T).format(c=c); args = {"command": c}
|
||||
call = json.dumps({"name": t, "arguments": args})
|
||||
assistant = f"<tool_call>\n{call}\n</tool_call>"
|
||||
return {"messages": [{"role": "user", "content": q},
|
||||
{"role": "assistant", "content": assistant}],
|
||||
"tools": TOOLS}
|
||||
|
||||
def main():
|
||||
seen, rows = set(), []
|
||||
while len(rows) < N:
|
||||
ex = make_example()
|
||||
key = ex["messages"][0]["content"] + ex["messages"][1]["content"]
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key); rows.append(ex)
|
||||
with open(OUT, "w") as f:
|
||||
for r in rows:
|
||||
f.write(json.dumps(r) + "\n")
|
||||
dist = {}
|
||||
for r in rows:
|
||||
name = json.loads(r["messages"][1]["content"].split("<tool_call>\n")[1].split("\n</tool_call>")[0])["name"]
|
||||
dist[name] = dist.get(name, 0) + 1
|
||||
print(f"wrote {len(rows)} -> {OUT}")
|
||||
print("distribution:", dist)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
47
mlb/tools_common.py
Normal file
47
mlb/tools_common.py
Normal file
@@ -0,0 +1,47 @@
|
||||
import re, json
|
||||
|
||||
TOOLS = [
|
||||
{"type": "function", "function": {
|
||||
"name": "read_file",
|
||||
"description": "Read and return the contents of a file.",
|
||||
"parameters": {"type": "object",
|
||||
"properties": {"path": {"type": "string", "description": "File path to read."}},
|
||||
"required": ["path"]}}},
|
||||
{"type": "function", "function": {
|
||||
"name": "write_file",
|
||||
"description": "Create or overwrite a file with the given content.",
|
||||
"parameters": {"type": "object",
|
||||
"properties": {"path": {"type": "string"}, "content": {"type": "string"}},
|
||||
"required": ["path", "content"]}}},
|
||||
{"type": "function", "function": {
|
||||
"name": "edit_file",
|
||||
"description": "Replace an exact substring in a file with new text.",
|
||||
"parameters": {"type": "object",
|
||||
"properties": {"path": {"type": "string"},
|
||||
"old_string": {"type": "string"},
|
||||
"new_string": {"type": "string"}},
|
||||
"required": ["path", "old_string", "new_string"]}}},
|
||||
{"type": "function", "function": {
|
||||
"name": "list_dir",
|
||||
"description": "List the entries in a directory.",
|
||||
"parameters": {"type": "object",
|
||||
"properties": {"path": {"type": "string"}},
|
||||
"required": ["path"]}}},
|
||||
{"type": "function", "function": {
|
||||
"name": "run_bash",
|
||||
"description": "Run a shell command and return its output.",
|
||||
"parameters": {"type": "object",
|
||||
"properties": {"command": {"type": "string"}},
|
||||
"required": ["command"]}}},
|
||||
]
|
||||
|
||||
_TOOLCALL_RE = re.compile(r"<tool_call>\s*(\{.*?\})\s*</tool_call>", re.DOTALL)
|
||||
|
||||
def parse_tool_calls(text):
|
||||
calls = []
|
||||
for m in _TOOLCALL_RE.findall(text):
|
||||
try:
|
||||
calls.append(json.loads(m))
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
return calls
|
||||
71
mlb/train_sft.py
Normal file
71
mlb/train_sft.py
Normal file
@@ -0,0 +1,71 @@
|
||||
import os, argparse
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
os.environ.setdefault("HF_HOME", os.path.join(HERE, "hf_cache"))
|
||||
|
||||
from datasets import load_dataset
|
||||
from peft import LoraConfig
|
||||
from trl import SFTConfig, SFTTrainer
|
||||
|
||||
MODEL = "Qwen/Qwen3-0.6B"
|
||||
TRAIN = os.path.join(HERE, "data", "sft_train.jsonl")
|
||||
VAL = os.path.join(HERE, "data", "sft_val.jsonl")
|
||||
|
||||
QWEN_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj",
|
||||
"gate_proj", "up_proj", "down_proj"]
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--lr", type=float, default=2e-4)
|
||||
ap.add_argument("--epochs", type=float, default=3)
|
||||
ap.add_argument("--frac", type=float, default=1.0) # 0.1 for LR sweep
|
||||
ap.add_argument("--dora", action="store_true") # run 1 = off
|
||||
ap.add_argument("--out", default=os.path.join(HERE, "adapters", "sft-lora"))
|
||||
ap.add_argument("--eval_steps", type=int, default=50)
|
||||
ap.add_argument("--logging_steps", type=int, default=10)
|
||||
ap.add_argument("--batch", type=int, default=8)
|
||||
ap.add_argument("--grad_accum", type=int, default=4)
|
||||
ap.add_argument("--no_grad_ckpt", action="store_true")
|
||||
a = ap.parse_args()
|
||||
|
||||
train = load_dataset("json", data_files=TRAIN, split="train")
|
||||
val = load_dataset("json", data_files=VAL, split="train")
|
||||
if a.frac < 1.0:
|
||||
train = train.select(range(int(len(train) * a.frac)))
|
||||
|
||||
peft = LoraConfig(r=16, lora_alpha=32, lora_dropout=0.05,
|
||||
bias="none", task_type="CAUSAL_LM",
|
||||
target_modules=QWEN_TARGETS, use_dora=a.dora)
|
||||
|
||||
cfg = SFTConfig(
|
||||
output_dir=a.out,
|
||||
model_init_kwargs={"dtype": "bfloat16"},
|
||||
max_length=512,
|
||||
packing=False,
|
||||
assistant_only_loss=True,
|
||||
use_liger_kernel=False,
|
||||
per_device_train_batch_size=a.batch,
|
||||
per_device_eval_batch_size=a.batch,
|
||||
gradient_accumulation_steps=a.grad_accum,
|
||||
num_train_epochs=a.epochs,
|
||||
learning_rate=a.lr,
|
||||
lr_scheduler_type="cosine",
|
||||
warmup_ratio=0.03,
|
||||
bf16=True,
|
||||
gradient_checkpointing=not a.no_grad_ckpt,
|
||||
logging_steps=a.logging_steps,
|
||||
eval_strategy="steps",
|
||||
eval_steps=a.eval_steps,
|
||||
save_strategy="epoch",
|
||||
report_to="none",
|
||||
)
|
||||
|
||||
trainer = SFTTrainer(model=MODEL, args=cfg,
|
||||
train_dataset=train, eval_dataset=val,
|
||||
peft_config=peft)
|
||||
trainer.train()
|
||||
trainer.save_model(a.out)
|
||||
print("saved adapter ->", a.out)
|
||||
print("final metrics:", trainer.state.log_history[-1])
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
22
mlb/verify_base.py
Normal file
22
mlb/verify_base.py
Normal file
@@ -0,0 +1,22 @@
|
||||
import os
|
||||
HERE = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
os.environ["HF_HOME"] = os.path.join(HERE, "hf_cache")
|
||||
|
||||
import torch
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
MODEL = "Qwen/Qwen3-0.6B"
|
||||
device = "mps" if torch.backends.mps.is_available() else "cpu"
|
||||
|
||||
tok = AutoTokenizer.from_pretrained(MODEL)
|
||||
model = AutoModelForCausalLM.from_pretrained(MODEL, dtype=torch.float32).to(device)
|
||||
|
||||
print("device:", device)
|
||||
print("params:", sum(p.numel() for p in model.parameters()))
|
||||
print("chat_template set:", tok.chat_template is not None)
|
||||
|
||||
msgs = [{"role": "user", "content": "Reply with exactly: pong"}]
|
||||
text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True, enable_thinking=False)
|
||||
ids = tok(text, return_tensors="pt").to(device)
|
||||
out = model.generate(**ids, max_new_tokens=16, do_sample=False)
|
||||
print("OUT:", tok.decode(out[0][ids.input_ids.shape[1]:], skip_special_tokens=True))
|
||||
Reference in New Issue
Block a user