详细讲解工具调用语言模型的微调全流程:轨迹解析、结构化提取、Qwen兼容ChatML格式及LoRA高效微调实现。
在本教程中,我们实现了一个端到端的监督微调流水线,使用 XYZ-Aquila-SFT 数据集、Hugging Face Transformers、PyTorch 和 PEFT。我们对数据集进行流式读取和检查,解析多轮工具调用轨迹,提取结构化工具调用,分析语料库特征,并保留嵌入的推理和观察模式。随后,我们在消息嵌入格式和结构化格式之间转换工具模式,渲染 Qwen 兼容的 ChatML(仅对 assistant 部分计算 loss),准备自定义 PyTorch 数据集和数据整理器(collator),并使用 LoRA 对 Qwen3-0.6B 进行微调。最后,我们在训练前后评估工具调用预测效果,并导出转换后的数据集和语料库统计数据以供进一步实验。
import os, sys, subprocess
CFG = dict(
REPO = "XYZAILab/XYZ-Aquila-SFT",
LANG = "en",
N_STREAM = 400,
N_EVAL = 40,
MODEL_ID = "Qwen/Qwen3-0.6B",
MAX_SEQ_LEN = 2048,
LENGTH_POLICY = "truncate",
RUN_TRAINING = True,
MAX_STEPS = 30,
GRAD_ACCUM = 8,
LR = 1e-4,
LORA_R = 16,
RUN_EVAL = True,
N_EVAL_PROBES = 24,
OUT_DIR = "/content/aquila_out",
SEED = 0,
)
os.makedirs(CFG["OUT_DIR"], exist_ok=True)
def pip(*pkgs):
subprocess.run([sys.executable, "-m", "pip", "install", "-q", "-U", *pkgs], check=False)
pip("datasets>=3.0.0", "transformers>=4.51.0", "peft>=0.13.0", "accelerate>=1.0.0")
import json, re, math, random, statistics as stats
from collections import Counter, defaultdict
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional
import torch
import matplotlib.pyplot as plt
from datasets import load_dataset
from transformers import AutoTokenizer, AutoModelForCausalLM, get_cosine_schedule_with_warmup
random.seed(CFG["SEED"]); torch.manual_seed(CFG["SEED"])
DEV = "cuda" if torch.cuda.is_available() else "cpu"
BF16 = DEV == "cuda" and torch.cuda.is_bf16_supported()
print(f"device={DEV} bf16={BF16} torch={torch.__version__}")
print(f"\n[1] streaming {CFG['REPO']}:{CFG['LANG']} ...")
stream = load_dataset(CFG["REPO"], CFG["LANG"], split="train", streaming=True)
RAW: List[Dict[str, Any]] = list(stream.take(CFG["N_STREAM"]))
print(f" pulled {len(RAW)} rows; keys = {list(RAW[0].keys())}")
_r = RAW[0]
print(f" question[:110] : {_r['question'][:110]}...")
print(f" answer : {_r['answer'][:80]}")
print(f" number of tool calls : {_r['number of tool calls']}")
print(f" trajectory len : {len(_r['trajectory'])} msgs")
print(f" role sequence (first8): {[m['role'] for m in _r['trajectory'][:8]]}")
我们为完整工作流配置数据集、模型、训练参数、输出目录和可重复性设置。我们安装所需的 Hugging Face、PEFT、Accelerate 和 PyTorch 相关依赖,并检测是否有 CUDA GPU 以及是否支持 BF16。随后,我们流式读取限定数量的 XYZ-Aquila-SFT 示例,检查数据集模式,并查看第一条工具调用轨迹的结构。
TOOLS_BLOCK_RE = re.compile(r"<tools>\s*(.*?)\s*</tools>", re.S)
THINK_RE = re.compile(r"<think>(.*?)</think>", re.S)
TOOL_RESP_RE = re.compile(r"<tool_response>\s*(.*?)\s*</tool_response>", re.S)
TOOLS_HDR_RE = re.compile(r"\n\n# Tools\n\n")
def iter_json_objects(text: str, limit: int = 1):
"""Nesting-safe JSON scanner. Regex like r'\\{.*?\\}' breaks on nested
`arguments` objects, which every real tool call has."""
dec, i, n, out = json.JSONDecoder(), 0, len(text), []
while i < n and len(out) < limit:
while i < n and text[i] not in "{[":
i += 1
if i >= n:
break
try:
obj, end = dec.raw_decode(text, i)
except json.JSONDecodeError:
i += 1
continue
out.append(obj); i = end
return out
def parse_tool_calls(content: str) -> List[Dict[str, Any]]:
calls = []
for m in re.finditer(r"<tool_call>", content):
got = iter_json_objects(content[m.end():], limit=1)
if got:
calls.append(got[0])
return calls
@dataclass
class Trajectory:
question: str
answer: str
declared_calls: int
messages: List[Dict[str, str]]
system_core: str = ""
tools: List[Dict[str, Any]] = field(default_factory=list)
tools_suffix: str = ""
calls: List[Dict[str, Any]] = field(default_factory=list)
n_observations: int = 0
n_think: int = 0
@property
def tool_names(self): return [c.get("name", "?") for c in self.calls]
@property
def depth(self): return len(self.messages)
def parse_row(row: Dict[str, Any]) -> Trajectory:
msgs = [{"role": m["role"], "content": m["content"]} for m in row["trajectory"]]
t = Trajectory(row["question"], row["answer"], row["number of tool calls"], msgs)
if msgs and msgs[0]["role"] == "system":
sysmsg = msgs[0]["content"]
split = TOOLS_HDR_RE.search(sysmsg)
if split:
t.system_core = sysmsg[:split.start()]
t.tools_suffix = sysmsg[split.start():]
else:
t.system_core = sysmsg
blk = TOOLS_BLOCK_RE.search(sysmsg)
if blk:
t.tools = iter_json_objects(blk.group(1), limit=64)
for m in msgs:
if m["role"] == "assistant":
t.calls += parse_tool_calls(m["content"])
t.n_think += len(THINK_RE.findall(m["content"]))
else:
t.n_observations += len(TOOL_RESP_RE.findall(m["content"]))
return t
TRAJ = [parse_row(r) for r in RAW]
t0 = TRAJ[0]
print(f"\n[2] parsed {len(TRAJ)} trajectories")
print(f" tool schemas found : {[fn.get('function', fn).get('name') for fn in t0.tools]}")
print(f" parsed calls : {len(t0.calls)} (declared {t0.declared_calls})")
print(f" observations : {t0.n_observations} think blocks: {t0.n_think}")
if t0.calls:
print(f" sample call : {json.dumps(t0.calls[0], ensure_ascii=False)[:200]}")
agree = sum(len(t.calls) == t.declared_calls for t in TRAJ)
print(f" parser vs 'number of tool calls': {agree}/{len(TRAJ)} exact match")
calls_per = [len(t.calls) for t in TRAJ]
depth_per = [t.depth for t in TRAJ]
chars_per = [sum(len(m["content"]) for m in t.messages) for t in TRAJ]
name_freq = Counter(n for t in TRAJ for n in t.tool_names)
argkey_freq = defaultdict(Counter)
for t in TRAJ:
for c in t.calls:
args = c.get("arguments", {})
if isinstance(args, dict):
for k in args: argkey_freq[c.get("name", "?")][k] += 1
def q(xs, p):
xs = sorted(xs); return xs[min(len(xs) - 1, int(p * len(xs)))]
print("\n[3] corpus statistics")
print(f" tool calls / traj : mean {stats.mean(calls_per):.1f} p50 {q(calls_per,.5)} "
f"p90 {q(calls_per,.9)} max {max(calls_per)}")
print(f" messages / traj : mean {stats.mean(depth_per):.1f} p90 {q(depth_per,.9)} max {max(depth_per)}")
print(f" chars / traj : mean {stats.mean(chars_per):,.0f} p90 {q(chars_per,.9):,}")
print(f" tool distribution : {dict(name_freq)}")
for k, v in argkey_freq.items(
我们定义了对嵌套安全的工具函数,用于从每个对话中提取 JSON 工具调用、推理块、观察结果和嵌入的工具模式。我们将每条原始数据行转换为结构化的轨迹对象,并验证解析出的工具调用数量是否与数据集中声明的值匹配。随后,我们计算语料库级别的统计数据,并可视化工具调用、消息深度、轨迹大小和工具使用频率的分布。
QWEN3_TOOLS_TMPL = ( "You are provided with function signatures within <tools></tools> XML tags:\n<tools>\n" "{lines}\n</tools>\n\nFor each function call, return a json object with function name " "and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n" '{{"name": <function-name>, "arguments": <args-json-object>}}\n</tool_call>' ) def extract_tools(t: Trajectory) -> Dict[str, Any]: """message-embedded schemas -> {'messages': [...], 'tools': [...]}""" msgs = [dict(m) for m in t.messages] if msgs and msgs[0]["role"] == "system": msgs[0]["content"] = t.system_core return {"messages": msgs, "tools": t.tools, "question": t.question, "answer": t.answer} def render_tools(rec: Dict[str, Any]) -> List[Dict[str, str]]: """inverse: structured tools -> schemas re-embedded in the system message""" msgs = [dict(m) for m in rec["messages"]] if rec["tools"] and msgs and msgs[0]["role"] == "system": lines = "\n".join(json.dumps(x, ensure_ascii=False) for x in rec["tools"]) msgs[0]["content"] = msgs[0]["content"] + QWEN3_TOOLS_TMPL.format(lines=lines) return msgs _rt = render_tools(extract_tools(t0)) exact = _rt[0]["content"] == t0.messages[0]["content"] print(f"\n[4] extract->render byte-exact: {exact}") if not exact: print(" template drift detected -> using verbatim tools_suffix for render()") a, b = t0.messages[0]["content"], _rt[0]["content"] i = next((i for i in range(min(len(a), len(b))) if a[i] != b[i]), min(len(a), len(b))) print(f" first divergence @{i}: {a[i:i+70]!r} vs {b[i:i+70]!r}") tok = AutoTokenizer.from_pretrained(CFG["MODEL_ID"]) if tok.pad_token is None: tok.pad_token = tok.eos_token IM_START, IM_END, NL = "<|im_start|>", "<|im_end|>", "\n" def render_and_mask(t: Trajectory, max_len: int, policy: str): """Manual ChatML so we control masking token-exactly. WHY NOT apply_chat_template(): Qwen3's template deletes <think>...</think> from every assistant turn except the last. On this dataset that silently destroys most of the reasoning supervision you are paying to train on. """ ids, labels = [], [] for m in t.messages: head = tok(f"{IM_START}{m['role']}{NL}", add_special_tokens=False).input_ids body = tok(m["content"], add_special_tokens=False).input_ids tail = tok(f"{IM_END}{NL}", add_special_tokens=False).input_ids seg = head + body + tail if m["role"] == "assistant": lab = [-100] * len(head) + body + tail else: lab = [-100] * len(seg) ids += seg; labels += lab if len(ids) > max_len: if policy == "drop": return None ids, labels = ids[:max_len], labels[:max_len] if all(l == -100 for l in labels): return None return {"input_ids": ids, "labels": labels} _probe = [{"role": "system", "content": "S"}, {"role": "user", "content": "U"}, {"role": "assistant", "content": "A"}] _mine = "".join(f"{IM_START}{m['role']}{NL}{m['content']}{IM_END}{NL}" for m in _probe) _theirs = tok.apply_chat_template(_probe, tokenize=False, add_generation_prompt=False) print(f"\n[5] manual ChatML == chat_template on tool-free probe: {_mine == _theirs}") if _mine != _theirs: print(f" mine : {_mine!r}\n theirs: {_theirs!r} (informational only)") ENC = [e for e in (render_and_mask(t, CFG["MAX_SEQ_LEN"], CFG["LENGTH_POLICY"]) for t in TRAJ) if e] sup = [sum(1 for x in e["labels"] if x != -100) / len(e["labels"]) for e in ENC] print(f" encoded {len(ENC)}/{len(TRAJ)} examples") print(f" supervised-token ratio: mean {stats.mean(sup):.3f} p10 {q(sup,.1):.3f} p90 {q(sup,.9):.3f}") over = sum(1 for t in TRAJ if sum(len(tok(m['content'], add_special_tokens=False).input_ids) for m in t.messages[:3]) > CFFG["MAX_SEQ_LEN"]) print(f" trajectories whose first 3 msgs alone exceed MAX_SEQ_LEN: {over}") SPLIT = len(ENC) - min(CFG["N_EVAL"], len(ENC)//5) TRAIN_ENC, EVAL_TRAJ = ENC
我们将嵌入的工具定义提取为结构化格式,然后重建它们以测试转换是否保留了原始系统消息。我们手动以 ChatML 格式渲染每条轨迹,以保留所有推理内容,并仅对助手生成的 token 应用损失。我们还对样本进行分词、执行所选的序列长度策略、创建训练和评估划分,并准备一个带填充的 PyTorch DataLoader。
def build_probes(trajs, n):
"""Teacher-forced probes: cut the trajectory right before an assistant turn
that issues a tool call; the gold label is that call."""
probes = []
for t in trajs:
for i, m in enumerate(t.messages):
if m["role"] != "assistant":
continue
gold = parse_tool_calls(m["content"])
if not gold:
continue
prefix = "".join(f"{IM_START}x['role']{NL}" for x in [])
prefix = "".join(f"{IM_START}{p['role']}{NL}{p['content']}{IM_END}{NL}"
for p in t.messages[:i]) + f"{IM_START}assistant{NL}"
if len(tok(prefix, add_special_tokens=False).input_ids) > CFG["MAX_SEQ_LEN"] - 160:
continue
probes.append({"prefix": prefix, "gold": gold[0]})
break
if len(probes) >= n:
break
return probes
@torch.no_grad()
def eval_tool_calls(model, probes, tag):
model.eval()
name_hit = arg_f1 = parsed = 0
for p in probes:
enc = tok(p["prefix"], return_tensors="pt", add_special_tokens=False).to(model.device)
out = model.generate(**enc, max_new_tokens=160, do_sample=False,
pad_token_id=tok.pad_token_id)
gen = tok.decode(out[0][enc.input_ids.shape[1]:], skip_special_tokens=True)
pred = (parse_tool_calls(gen) or iter_json_objects(gen, limit=1) or [None])[0]
if not isinstance(pred, dict):
continue
parsed += 1
g = p["gold"]
name_hit += int(pred.get("name") == g.get("name"))
pk = set((pred.get("arguments") or {}).keys()) if isinstance(pred.get("arguments"), dict) else set()
gk = set((g.get("arguments") or {}).keys()) if isinstance(g.get("arguments"), dict) else set()
if pk or gk:
inter = len(pk & gk)
arg_f1 += 0.0 if inter == 0 else 2*inter/(len(pk)+len(gk))
n = max(1, len(probes))
print(f" [{tag}] parseable {parsed}/{n} | tool-name acc {name_hit/n:.3f} | arg-key F1 {arg_f1/n:.3f}")
return dict(parsed=parsed/n, name_acc=name_hit/n, arg_f1=arg_f1/n)
PROBES = build_probes(EVAL_TRAJ, CFG["N_EVAL_PROBES"])
print(f" built {len(PROBES)} teacher-forced probes")
results = {}
if CFG["RUN_TRAINING"]:
from peft import LoraConfig, get_peft_model
dtype = torch.bfloat16 if BF16 else torch.float32
model = AutoModelForCausalLM.from_pretrained(
CFG["MODEL_ID"], torch_dtype=dtype, attn_implementation="sdpa").to(DEV)
model.config.use_cache = False
model.gradient_checkpointing_enable()
model.enable_input_require_grads()
if CFG["RUN_EVAL"] and PROBES and DEV == "cuda":
print("\n[8] baseline eval")
results["before"] = eval_tool_calls(model, PROBES, "base")
model = get_peft_model(model, LoraConfig(
r=CFG["LORA_R"], lora_alpha=2*CFG["LORA_R"], lora_dropout=0.05,
bias="none", task_type="CAUSAL_LM",
model.print_trainable_parameters()
opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad],
lr=CFG["LR"], weight_decay=0.0, betas=(0.9, 0.95))
sched = get_cosine_schedule_with_warmup(opt, 5, CFG["MAX_STEPS"])
scaler = torch.amp.GradScaler("cuda", enabled=(DEV == "cuda" and not BF16))
amp_dt = torch.bfloat16 if BF16 else torch.float16
print(f"\n[7] training {CFG['MAX_STEPS']} steps "
f"(bs1 x accum{CFG['GRAD_ACCUM']} = {CFG['GRAD_ACCUM']} traj/step)")
model.train(); step = 0; run = None; it = iter(loader)
while step < CFG["MAX_STEPS"]:
opt.zero_grad(set_to_none=True); acc = 0.0
for _ in range(CFG["GRAD_ACCUM"]):
try: batch = next(it)
except StopIteration:
it = iter(loader); batch = next(it)
batch = {k: v.to(DEV) for k, v in batch.items()}
with torch.autocast(DEV, dtype=amp_dt, enabled=(DEV == "cuda")):
loss = model(**batch).loss / CFG["GRAD_ACCUM"]
scaler.scale(loss).backward() if scal
我们通过在包含工具调用的助手回合之前切断轨迹来构建教师强制评估探针。我们加载 Qwen3-0.6B,测量其基线工具调用性能,附加 LoRA 适配器,并使用梯度累积、混合精度、梯度检查点、梯度裁剪和余弦学习率调度对模型进行微调。然后我们评估适配后的模型,将其指标与基线进行比较,并保存训练好的 LoRA 适配器和分词器。
```php
struct_path = f"{CFG['OUT_DIR']}/aquila_{CFG['LANG']}_structured_tools.jsonl"
with open(struct_path, "w", encoding="utf-8") as f:
for t in TRAJ:
f.write(json.dumps(extract_tools(t), ensure_ascii=False) + "\n")
stats_path = f"{CFG['OUT_DIR']}/corpus_stats.json"
with open(stats_path, "w") as f:
json.dump({"n": len(TRAJ), "tool_freq": dict(name_freq),
"calls_mean": stats.mean(calls_per), "calls_max": max(calls_per),
"depth_p90": q(depth_per, .9), "encoded": len(ENC),
"supervised_ratio_mean": stats.mean(sup), "eval": results}, f, indent=2)
print(f"\n[9] wrote:\n {struct_path}\n {stats_path}")
print("done.")
我们将每条解析后的轨迹导出为结构化的 JSONL 记录,其中包含消息、工具模式(schemas)、问题及答案。同时保存一份 JSON 报告,涵盖语料库规模、工具频率、轨迹统计、有监督 token 比例以及可用的评估结果。整个工作流以可复用的数据集产物、分析输出和模型文件结束,这些内容存储在配置好的输出目录中。
我们完成了一条完整的实用流水线,用于对 XYZ-Aquila-SFT 数据集中复杂的工具使用轨迹进行分析、转换、微调和评估。我们保留了原始对话结构,仅在助手回复上应用 token 级别的监督,并使用 LoRA 在兼容 Colab 的 GPU 上高效适配 Qwen3-0.6B。我们还通过教师强制评估(teacher-forced evaluation)比较了基线和微调后的工具调用性能,并导出了可复用的结构化记录、模型适配器以及分析统计数据。这套工作流为我们大规模扩展工具感知的监督微调、测试替代性序列长度策略以及训练更强大的智能体语言模型奠定了坚实基础。
查看完整代码。同时,欢迎关注我们的 Twitter,记得加入我们的 15 万+ ML SubReddit 并订阅我们的Newsletter。等一下!你用 Telegram 吗?现在也可以加入我们了。
需要与我们合作推广你的 GitHub 仓库、Hugging Face 页面、产品发布或网络研讨会吗?联系我们吧。
Sana Hassan,Marktechpost 咨询实习生,同时是马德拉斯理工学院(IIT Madras)的双学位学生,热爱将技术与 AI 应用于解决现实世界挑战。凭借解决实际问题的浓厚兴趣,他为 AI 与现实生活解决方案的交叉领域带来了全新的视角。