垂直领域问答助手开发

文章来源声明: 原文作者:warrior; 来源站点:掘金; 原文链接:https://juejin.cn/post/7687015527131234319; 本文基于上述来源整理/加工,觅优补充点评,仅供技术学习交流。版权归原作者所有。
觅优短评

适合有Python基础、想低成本落地私有领域问答的开发者,链路完整、代码可复制,踩坑清单尤其实用;若追求开箱即用,还需自行补充数据与评测环节。

垂直领域问答助手开发 ----------

前段时间做了一个垂直领域的问答助手项目,从零把"数据合成→参数高效微调→量化压缩→本地部署"整条链路跑通了。网上碎片教程很多,但很多坑只有亲手踩过才知道,这里把完整方案和踩坑记录分享出来。基座模型用 Qwen3,显存 8G 就能玩。

一、整体方案

原始文档(txt)
    │  LangChain + Ollama(qwen3) 批量生成问答对
    ▼
instruction-tuning 数据集 (JSON, 200+条)
    │  bitsandbytes 4bit + LoRA (peft)
    ▼
QLoRA 微调后的基座模型
    │  llama.cpp 转 GGUF → INT4 量化
    ▼
1GB 左右的量化模型, CPU/低端GPU 可跑

核心思路:没有现成语料就自己合成。用强模型(qwen3:7b)从非结构化文档里"出题",生成指令微调数据,再反哺小模型——这就是所谓的 Self-Instruct / 弱到强蒸馏的平民版。

二、环境

<span># 云端 GPU 机器</span>
pip install langchain-core langchain-community langchain-text-splitters langchain-ollama
pip install torch transformers peft datasets bitsandbytes accelerate
pip install llama-cpp-python gguf

  • 基座:Qwen3-1.7B(练习)/ 7B(生产),HuggingFace 可下
  • 推理服务:Ollama,模型 qwen3:7b
  • 显存参考:1.7B 4bit 微调 ~6GB;7B 4bit 微调 ~12GB(batch 调小可行)

三、第一步:用 LangChain 从文档合成指令数据

3.1 分块

文档先切块(chunk_size≈512),逐块喂给模型出题:

<span>from</span> langchain_text_splitters <span>import</span> RecursiveCharacterTextSplitter

splitter = RecursiveCharacterTextSplitter(
    chunk_size=<span>512</span>, chunk_overlap=<span>50</span>,
    separators=[<span>"\n\n"</span>, <span>"\n"</span>, <span>"。"</span>, <span>","</span>, <span>" "</span>, <span>""</span>],
)
chunks = splitter.split_text(text)

中文分隔符一定要给 "。",否则切出来的块会从句子中间断开。

3.2 Prompt 模板三要素

一个稳定的出题模板 = 角色设定 + 任务说明 + 严格的输出格式约束

PROMPT = PromptTemplate.from_template(
    <span>"你是一个擅长总结和出题的AI助手,精通{domain}领域知识。\n"</span>
    <span>"任务说明:请根据下面的【文档片段】,生成 {n} 个具体的问答对。\n"</span>
    <span>"要求:\n"</span>
    <span>"1. 问题必须基于片段中的事实,答案必须能在片段中找到依据;\n"</span>
    <span>"2. 问题要具体、多样,不要重复;\n"</span>
    <span>"3. 严格输出 JSON 格式,不要输出任何其他内容,格式如下:\n"</span>
    <span>'[{{"instruction": "问题1", "input": "", "output": "答案1"}}, ...]\n\n'</span>
    <span>"【文档片段】\n{chunk}\n"</span>
)

3.3 ⚠️ 最大的坑:Qwen3 的 <think> 标签

Qwen3 默认开启思考模式,输出是:

<think>用户想要...我应该...</think>
[{"instruction": ...}]

直接 json.loads 必炸。解析前必须剥离思考段:

<span>def</span> <span>strip_think</span>(<span>text</span>):
    <span>return</span> re.sub(<span>r"<think>.*?</think>"</span>, <span>""</span>, text, flags=re.DOTALL).strip()

<span>def</span> <span>extract_json</span>(<span>text</span>):
    text = strip_think(text)
    m = re.search(<span>r"\[.*\]"</span>, text, flags=re.DOTALL)  <span># 模型可能用```json```包裹</span>
    <span>if</span> <span>not</span> m: <span>return</span> <span>None</span>
    <span>try</span>:
        data = json.loads(m.group(<span>0</span>))
        <span>return</span> data <span>if</span> <span>isinstance</span>(data, <span>list</span>) <span>else</span> <span>None</span>
    <span>except</span> json.JSONDecodeError:
        <span>return</span> <span>None</span>

3.4 数据质量控制

合成数据≠能用数据,三道过滤必须有:

<span>def</span> <span>valid_item</span>(<span>item, domain, seen</span>):
    <span>if</span> <span>not</span> <span>isinstance</span>(item, <span>dict</span>): <span>return</span> <span>None</span>
    ins = <span>str</span>(item.get(<span>"instruction"</span>,<span>""</span>)).strip()
    out = <span>str</span>(item.get(<span>"output"</span>,<span>""</span>)).strip()
    <span>if</span> <span>len</span>(ins) < <span>5</span> <span>or</span> <span>len</span>(out) < <span>10</span>: <span>return</span> <span>None</span>   <span># 太短=无效</span>
    key = ins[:<span>40</span>]
    <span>if</span> key <span>in</span> seen: <span>return</span> <span>None</span>                      <span># 去重</span>
    seen.add(key)
    <span>return</span> {<span>"instruction"</span>: ins, <span>"input"</span>: <span>""</span>, <span>"output"</span>: out, <span>"domain"</span>: domain}

实测 1.7B 模型重复率明显高于 7B,目标 200 条时每个文档要把分块循环复用(块不够就 round_i % len(chunks) 轮转),直到每域攒够 60+ 条。

四、第二步:QLoRA 微调

4.1 数据模板与 labels 对齐

### 指令:\n{instruction}\n### 回答:\n{output} 的朴素模板(比 chat template 更适合让小模型学格式),只对答案部分计算 loss

<span>def</span> <span>build_example</span>(<span>tokenizer, item, max_len</span>):
    prompt = <span>f"### 指令:\n<span>{item[<span>'instruction'</span>]}</span>\n### 回答:\n"</span>
    target = item[<span>"output"</span>] + tokenizer.eos_token
    prompt_ids = tokenizer(prompt, add_special_tokens=<span>False</span>)[<span>"input_ids"</span>]
    full_ids = tokenizer(prompt + target, add_special_tokens=<span>False</span>)[<span>"input_ids"</span>][:max_len]
    labels = [-<span>100</span>] * <span>min</span>(<span>len</span>(prompt_ids), <span>len</span>(full_ids)) + full_ids[<span>len</span>(prompt_ids):]
    <span>return</span> {<span>"input_ids"</span>: full_ids, <span>"labels"</span>: labels, <span>"attention_mask"</span>: [<span>1</span>]*<span>len</span>(full_ids)}

4.2 4bit 加载 + LoRA 挂载

bnb_config = BitsAndBytesConfig(
    load_in_4bit=<span>True</span>,
    bnb_4bit_quant_type=<span>"nf4"</span>,            <span># NF4 精度优于普通 int4</span>
    bnb_4bit_compute_dtype=torch.float16,
    bnb_4bit_use_double_quant=<span>True</span>,
)
model = AutoModelForCausalLM.from_pretrained(base, quantization_config=bnb_config, device_map=<span>"auto"</span>)
model = prepare_model_for_kbit_training(model)

lora_config = LoraConfig(
    r=<span>8</span>, lora_alpha=<span>16</span>, lora_dropout=<span>0.05</span>,
    target_modules=[<span>"q_proj"</span>,<span>"k_proj"</span>,<span>"v_proj"</span>,<span>"o_proj"</span>],
    task_type=<span>"CAUSAL_LM"</span>,
)
model = get_peft_model(model, lora_config)
<span># trainable params: 3.2M / 1.7B = 0.19%</span>

0.19% 的可训练参数,单卡轻松跑。

4.3 三个编译不过的坑(transformers 5.x 实测)

  1. apply_chat_template 返回 BatchEncoding 不是 tensor:推理时要取 enc["input_ids"].to(device),直接 .to() 会 AttributeError
  2. warmup_ratio 参数被移除:用 warmup_steps=10(新旧版本都兼容)
  3. 自定义 labels 的数据集不能用 DataCollatorForLanguageModeling(报 excessive nesting),换 DataCollatorForSeq2Seq(tokenizer, padding=True)——它会自动用 -100 pad labels

训练参数(200条数据实测):

TrainingArguments(
    num_train_epochs=<span>3</span>, per_device_train_batch_size=<span>4</span>,
    gradient_accumulation_steps=<span>4</span>, learning_rate=<span>2e-4</span>,
    lr_scheduler_type=<span>"cosine"</span>, warmup_steps=<span>10</span>,
    bf16=<span>True</span>, eval_strategy=<span>"steps"</span>, eval_steps=<span>50</span>,
)
<span># 1 epoch 15 秒, 3 epoch 后 loss 3.2 → 1.1</span>

4.4 合并权重保存

4bit 模型不能直接 merge。要重新以 FP16 加载基座再合并:

<span>del</span> model, trainer; torch.cuda.empty_cache()   <span># 先释放</span>

fp16 = AutoModelForCausalLM.from_pretrained(base, torch_dtype=torch.float16, device_map=<span>"auto"</span>)
merged = PeftModel.from_pretrained(fp16, <span>"lora_adapter"</span>).merge_and_unload()
merged.save_pretrained(<span>"checkpoint-best"</span>, safe_serialization=<span>True</span>)
tokenizer.save_pretrained(<span>"checkpoint-best"</span>)

检查产物:config.json + model.safetensors + tokenizer.json 三件套齐,from_pretrained 直接可用。

五、第三步:GGUF INT4 量化

FP16 的 1.7B 有 3.4GB,量化后 1GB,这是能塞进 8GB 显存/纯 CPU 部署的关键。

<span># 1. HF 格式 → GGUF (convert.py 是旧名, 新版叫 convert_hf_to_gguf.py)</span>
python llama.cpp/convert_hf_to_gguf.py checkpoint-best \
       --outfile base.gguf --outtype f16
<span># 2. FP16 → INT4</span>
./llama.cpp/build/bin/llama-quantize base.gguf model_int4.gguf q4_0

实测坑:

  • f16 大文件在容器/overlay 文件系统可能写失败(报 Not enough free space 但 df 显示充足)→ 改 --outtype q8_0,quantize 加 --allow-requantize
  • quantize 完记得对比体积:3.44GB → 1.05GB(3.3x)

六、第四步:部署与冒烟测试

6.1 llama-cpp-python 推理

<span>from</span> llama_cpp <span>import</span> Llama
llm = Llama(model_path=<span>"model_int4.gguf"</span>, n_ctx=<span>2048</span>, n_gpu_layers=<span>99</span>, verbose=<span>False</span>)
out = llm(<span>"### 指令:\n什么是金融危机?\n### 回答:\n"</span>,
          max_tokens=<span>96</span>, temperature=<span>0.1</span>,
          repeat_penalty=<span>1.15</span>,        <span># 防复读, 必加!</span>
          stop=[<span>"###"</span>])
<span>print</span>(out[<span>"choices"</span>][<span>0</span>][<span>"text"</span>].strip())

repeat_penalty=1.15 是血泪教训:不加的话小模型很容易输出"用口语化表达。用简单词汇。用口语化表达..."的复读机模式。

6.2 更省事的替代:GGUF 挂进 Ollama

<span>echo</span> <span>"FROM ./model_int4.gguf"</span> > Modelfile
ollama create my-int4 -f Modelfile
ollama run my-int4 <span>"你是谁?"</span>

6.3 冒烟测试怎么写

固定 3 个问题(身份题+通识题+领域题),每题记录:输出文本(查乱码/复读/断裂)、token 数、耗时、tok/s,落一份报告文件。实测量化后 155 tok/s(4090),与量化前持平,质量无肉眼可见差异。

七、总结

环节耗时关键点
数据合成 216 条~40min剥 think 标签 + 去重过滤
QLoRA 3 epochs~2minlabels 对齐、只学答案
GGUF 量化~3minf16 写盘失败就 q8\_0
冒烟部署~5minrepeat\_penalty 防复读

整条链路最贵的不是算力,是数据质量:合成数据宁可少不可脏,去重和长度过滤做扎实,微调才有意义。


欢迎评论区交流,踩坑互助。如果对你有帮助,点赞收藏支持一下~

八、完整代码附录(复制即用)

8.1 数据合成 gen_dataset.py(全文)

<span># -*- coding: utf-8 -*-</span>
<span>"""
任务一:数据整理 —— langchain + ollama(qwen3) 批量生成指令微调数据集
用法:
    python gen_dataset.py                          # 默认 qwen3:7b, 每域60条
    python gen_dataset.py --model qwen3:0.5b       # 快速测试用小模型
依赖:
    pip install langchain-core langchain-community langchain-ollama
输出:
    /home/user/workspace/model_b/data/dataset.json  (>=200条, 每域>=60条)
要点(考场易错):
    1. qwen3 默认输出 <think>...</think> 思考过程 -> 解析前必须剥离, 否则JSON解析失败
    2. 模型可能把 JSON 包在 ```json ```代码块里 -> 用正则提取第一个 [...] 块
    3. 读取 txt 一律 encoding="utf-8"
    4. 去重与校验: instruction 非空且不重复, output 长度>10 才算有效数据
"""</span>
<span>import</span> argparse
<span>import</span> json
<span>import</span> re
<span>from</span> pathlib <span>import</span> Path

<span>from</span> langchain_community.llms <span>import</span> Ollama
<span>from</span> langchain_core.prompts <span>import</span> PromptTemplate
<span>from</span> langchain_text_splitters <span>import</span> RecursiveCharacterTextSplitter

<span># ---------------- 配置 ----------------</span>
DATA_DIR = Path(<span>"/home/user/workspace/data"</span>)
OUT_PATH = Path(<span>"/home/user/workspace/model_b/data/dataset.json"</span>)

<span># 三个领域文档 -> 主题标签</span>
DOCS = {
    <span>"金融"</span>: DATA_DIR / <span>"美国次贷危机.txt"</span>,
    <span>"医疗"</span>: DATA_DIR / <span>"尿毒性心包炎.txt"</span>,
    <span>"法规"</span>: DATA_DIR / <span>"生成式人工智能服务管理暂行办法.txt"</span>,
}

<span># ============== 步骤1-1: Prompt 模板(角色设定+任务说明+格式要求) ==============</span>
PROMPT_TEMPLATE = PromptTemplate.from_template(
    <span>"你是一个擅长总结和出题的AI助手,精通{domain}领域知识。\n"</span>
    <span>"任务说明:请根据下面的【文档片段】,生成 {n} 个具体的问答对(一个问题和它的标准答案)。\n"</span>
    <span>"要求:\n"</span>
    <span>"1. 问题必须基于片段中的事实,答案必须能在片段中找到依据;\n"</span>
    <span>"2. 问题要具体、多样,不要重复;\n"</span>
    <span>"3. 严格输出 JSON 格式,不要输出任何其他内容,格式如下:\n"</span>
    <span>'[{{"instruction": "问题1", "input": "", "output": "答案1"}},\n'</span>
    <span>' {{"instruction": "问题2", "input": "", "output": "答案2"}}]\n\n'</span>
    <span>"【文档片段】\n{chunk}\n"</span>
)


<span>def</span> <span>strip_think</span>(<span>text: <span>str</span></span>) -> <span>str</span>:
    <span>"""剥离 qwen3 的 <think>...</think> 思考内容"""</span>
    <span>return</span> re.sub(<span>r"<think>.*?</think>"</span>, <span>""</span>, text, flags=re.DOTALL).strip()


<span>def</span> <span>extract_json</span>(<span>text: <span>str</span></span>):
    <span>"""从模型输出中稳健地提取 JSON 列表(容错: 代码块包裹/前后废话)"""</span>
    text = strip_think(text)
    m = re.search(<span>r"\[.*\]"</span>, text, flags=re.DOTALL)   <span># 找第一个 [...] 块</span>
    <span>if</span> <span>not</span> m:
        <span>return</span> <span>None</span>
    <span>try</span>:
        data = json.loads(m.group(<span>0</span>))
        <span>return</span> data <span>if</span> <span>isinstance</span>(data, <span>list</span>) <span>else</span> <span>None</span>
    <span>except</span> json.JSONDecodeError:
        <span>return</span> <span>None</span>


<span>def</span> <span>valid_item</span>(<span>item, domain_label, seen</span>):
    <span>"""校验一条数据是否'有效', 并打上主题标签"""</span>
    <span>if</span> <span>not</span> <span>isinstance</span>(item, <span>dict</span>):
        <span>return</span> <span>None</span>
    ins = <span>str</span>(item.get(<span>"instruction"</span>, <span>""</span>)).strip()
    out = <span>str</span>(item.get(<span>"output"</span>, <span>""</span>)).strip()
    <span>if</span> <span>len</span>(ins) < <span>5</span> <span>or</span> <span>len</span>(out) < <span>10</span>:      <span># 太短视为无效</span>
        <span>return</span> <span>None</span>
    key = ins[:<span>40</span>]                          <span># 粗粒度去重</span>
    <span>if</span> key <span>in</span> seen:
        <span>return</span> <span>None</span>
    seen.add(key)
    <span>return</span> {<span>"instruction"</span>: ins, <span>"input"</span>: <span>""</span>, <span>"output"</span>: out,
            <span>"domain"</span>: domain_label}         <span># domain 字段便于统计, 训练时忽略</span>


<span>def</span> <span>main</span>():
    ap = argparse.ArgumentParser()
    ap.add_argument(<span>"--model"</span>, default=<span>"qwen3:7b"</span>, <span>help</span>=<span>"ollama 模型名"</span>)
    ap.add_argument(<span>"--per-domain"</span>, <span>type</span>=<span>int</span>, default=<span>60</span>, <span>help</span>=<span>"每个领域目标条数"</span>)
    ap.add_argument(<span>"--batch-per-chunk"</span>, <span>type</span>=<span>int</span>, default=<span>4</span>, <span>help</span>=<span>"每次调用生成条数"</span>)
    ap.add_argument(<span>"--chunk-size"</span>, <span>type</span>=<span>int</span>, default=<span>512</span>, <span>help</span>=<span>"文档分块大小"</span>)
    args = ap.parse_args()

    <span># ============== 步骤1-2: 分块(chunk_size约512) + 循环调用 ==============</span>
    splitter = RecursiveCharacterTextSplitter(
        chunk_size=args.chunk_size, chunk_overlap=<span>50</span>,
        separators=[<span>"\n\n"</span>, <span>"\n"</span>, <span>"。"</span>, <span>","</span>, <span>" "</span>, <span>""</span>],
    )
    llm = Ollama(model=args.model, temperature=<span>0.7</span>)

    seen = <span>set</span>()
    dataset = []
    <span>for</span> domain, path <span>in</span> DOCS.items():
        text = path.read_text(encoding=<span>"utf-8"</span>)
        chunks = splitter.split_text(text)
        <span>print</span>(<span>f"[<span>{domain}</span>] <span>{path.name}</span>: <span>{<span>len</span>(text)}</span>字 -> <span>{<span>len</span>(chunks)}</span> 块"</span>)
        got = <span>0</span>                      <span># 该领域已生成条数</span>
        round_i = <span>0</span>
        <span>while</span> got < args.per_domain <span>and</span> round_i < <span>len</span>(chunks) * <span>6</span>:
            chunk = chunks[round_i % <span>len</span>(chunks)]
            round_i += <span>1</span>
            prompt = PROMPT_TEMPLATE.<span>format</span>(
                domain=domain, n=args.batch_per_chunk, chunk=chunk)
            <span>try</span>:
                resp = llm.invoke(prompt)
            <span>except</span> Exception <span>as</span> e:
                <span>print</span>(<span>f"  调用失败(重试): <span>{e}</span>"</span>)
                <span>continue</span>
            <span>for</span> item <span>in</span> (extract_json(resp) <span>or</span> []):
                v = valid_item(item, domain, seen)
                <span>if</span> v:
                    dataset.append(v)
                    got += <span>1</span>
            <span>print</span>(<span>f"  [<span>{domain}</span>] 已生成 <span>{got}</span>/<span>{args.per_domain}</span> 条"</span>)
        <span>if</span> got < args.per_domain:
            <span>print</span>(<span>f"  警告: [<span>{domain}</span>] 只生成 <span>{got}</span> 条(可增大重试上限或换模型)"</span>)

    <span># ============== 步骤1-3: 保存 ==============</span>
    OUT_PATH.parent.mkdir(parents=<span>True</span>, exist_ok=<span>True</span>)
    OUT_PATH.write_text(
        json.dumps(dataset, ensure_ascii=<span>False</span>, indent=<span>2</span>), encoding=<span>"utf-8"</span>)

    <span># 统计报告</span>
    <span>from</span> collections <span>import</span> Counter
    c = Counter(d[<span>"domain"</span>] <span>for</span> d <span>in</span> dataset)
    <span>print</span>(<span>f"\n完成! 共 <span>{<span>len</span>(dataset)}</span> 条 -> <span>{OUT_PATH}</span>"</span>)
    <span>print</span>(<span>"领域分布:"</span>, <span>dict</span>(c))
    <span>assert</span> <span>len</span>(dataset) >= <span>200</span>, <span>"总数不足200条!"</span>
    <span>for</span> k, n <span>in</span> c.items():
        <span>assert</span> n >= <span>60</span>, <span>f"<span>{k}</span> 不足60条!"</span>


<span>if</span> __name__ == <span>"__main__"</span>:
    main()

8.2 QLoRA 微调 qlora.py(全文)

<span># -*- coding: utf-8 -*-</span>
<span>"""
步骤2-2:QLoRA 微调脚本(50分)
对 /home/user/workspace/QwenPretrain/ 的 QWEN3 模型做 QLoRA 4-bit 微调。
训练数据: /home/user/workspace/model_b/data/dataset.json (任务一生成)
模板: 严格遵循 /home/user/workspace/tmp.doc:
    "### 指令:\n{instruction}\n### 回答:\n{output}"
训练完合并 LoRA 权重, 保存完整模型到 /home/user/workspace/model_b/checkpoint-best/

用法:
    python qlora.py                          # 完整微调
    python qlora.py --epochs 1 --max-len 256 # 快速跑通(练习时先用)
依赖:
    pip install torch transformers peft datasets bitsandbytes accelerate
"""</span>
<span>import</span> argparse
<span>import</span> json
<span>import</span> os
<span>import</span> torch
<span>from</span> datasets <span>import</span> Dataset
<span>from</span> transformers <span>import</span> (AutoModelForCausalLM, AutoTokenizer, Trainer,
                          TrainingArguments, BitsAndBytesConfig, DataCollatorForSeq2Seq)
<span>from</span> peft <span>import</span> LoraConfig, get_peft_model, prepare_model_for_kbit_training, PeftModel

BASE_MODEL = <span>"/home/user/workspace/QwenPretrain"</span>
DATA_PATH = <span>"/home/user/workspace/model_b/data/dataset.json"</span>
ADAPTER_DIR = <span>"/home/user/workspace/model_b/lora_adapter"</span>       <span># LoRA 适配器</span>
MERGED_DIR = <span>"/home/user/workspace/model_b/checkpoint-best"</span>     <span># 合并后完整模型(步骤2-3)</span>


<span># ---------- 数据处理: 拼模板 + labels 对齐 ----------</span>
<span>def</span> <span>build_example</span>(<span>tokenizer, item, max_len</span>):
    <span>"""按 tmp.doc 模板拼接; 提示部分 label=-100 只学答案"""</span>
    instruction = item[<span>"instruction"</span>]
    inp = item.get(<span>"input"</span>, <span>""</span>)
    output = item[<span>"output"</span>]

    prompt = <span>f"### 指令:\n<span>{instruction}</span>\n"</span>
    <span>if</span> inp:
        prompt += <span>f"### 输入:\n<span>{inp}</span>\n"</span>
    prompt += <span>"### 回答:\n"</span>
    target = output + tokenizer.eos_token

    prompt_ids = tokenizer(prompt, add_special_tokens=<span>False</span>)[<span>"input_ids"</span>]
    full_ids = tokenizer(prompt + target, add_special_tokens=<span>False</span>)[<span>"input_ids"</span>]
    full_ids = full_ids[:max_len]
    labels = [-<span>100</span>] * <span>min</span>(<span>len</span>(prompt_ids), <span>len</span>(full_ids)) + full_ids[<span>len</span>(prompt_ids):]
    <span>return</span> {<span>"input_ids"</span>: full_ids, <span>"labels"</span>: labels,
            <span>"attention_mask"</span>: [<span>1</span>] * <span>len</span>(full_ids)}


<span>def</span> <span>main</span>():
    ap = argparse.ArgumentParser()
    ap.add_argument(<span>"--base-model"</span>, default=BASE_MODEL)
    ap.add_argument(<span>"--epochs"</span>, <span>type</span>=<span>int</span>, default=<span>3</span>)
    ap.add_argument(<span>"--batch"</span>, <span>type</span>=<span>int</span>, default=<span>4</span>)
    ap.add_argument(<span>"--grad-accum"</span>, <span>type</span>=<span>int</span>, default=<span>4</span>)
    ap.add_argument(<span>"--lr"</span>, <span>type</span>=<span>float</span>, default=<span>2e-4</span>)
    ap.add_argument(<span>"--max-len"</span>, <span>type</span>=<span>int</span>, default=<span>512</span>)
    args = ap.parse_args()

    tokenizer = AutoTokenizer.from_pretrained(args.base_model, trust_remote_code=<span>True</span>)
    <span>if</span> tokenizer.pad_token <span>is</span> <span>None</span>:
        tokenizer.pad_token = tokenizer.eos_token

    <span># ---------- 1. 4-bit 量化加载 (bitsandbytes) ----------</span>
    bnb_config = BitsAndBytesConfig(
        load_in_4bit=<span>True</span>,
        bnb_4bit_quant_type=<span>"nf4"</span>,              <span># NF4 精度优于普通 int4</span>
        bnb_4bit_compute_dtype=torch.float16,
        bnb_4bit_use_double_quant=<span>True</span>,
    )
    model = AutoModelForCausalLM.from_pretrained(
        args.base_model,
        quantization_config=bnb_config,
        device_map=<span>"auto"</span>,
        trust_remote_code=<span>True</span>,
    )
    model.config.use_cache = <span>False</span>

    <span># ---------- 2. LoRA 配置 ----------</span>
    model = prepare_model_for_kbit_training(model)
    lora_config = LoraConfig(
        r=<span>8</span>,
        lora_alpha=<span>16</span>,
        lora_dropout=<span>0.05</span>,
        bias=<span>"none"</span>,
        task_type=<span>"CAUSAL_LM"</span>,
        target_modules=[<span>"q_proj"</span>, <span>"k_proj"</span>, <span>"v_proj"</span>, <span>"o_proj"</span>],
    )
    model = get_peft_model(model, lora_config)
    model.print_trainable_parameters()

    <span># ---------- 3. 加载数据集 ----------</span>
    <span>with</span> <span>open</span>(DATA_PATH, encoding=<span>"utf-8"</span>) <span>as</span> f:
        raw = json.load(f)
    ds = Dataset.from_list([
        build_example(tokenizer, it, args.max_len) <span>for</span> it <span>in</span> raw
    ]).train_test_split(test_size=<span>0.05</span>, seed=<span>42</span>)
    <span>print</span>(<span>f"训练 <span>{<span>len</span>(ds[<span>'train'</span>])}</span> 条 / 验证 <span>{<span>len</span>(ds[<span>'test'</span>])}</span> 条"</span>)

    <span># ---------- 4. 训练参数 ----------</span>
    targs = TrainingArguments(
        output_dir=<span>"/home/user/workspace/model_b/qlora_out"</span>,
        num_train_epochs=args.epochs,
        per_device_train_batch_size=args.batch,
        per_device_eval_batch_size=args.batch,
        gradient_accumulation_steps=args.grad_accum,
        learning_rate=args.lr,
        lr_scheduler_type=<span>"cosine"</span>,
        warmup_steps=<span>10</span>,   <span># 注: transformers 5.x 移除了 warmup_ratio, 用 steps 新旧版均兼容</span>
        logging_steps=<span>10</span>,
        eval_strategy=<span>"steps"</span>,       <span># 旧版 transformers 用 evaluation_strategy</span>
        eval_steps=<span>50</span>,
        save_strategy=<span>"no"</span>,          <span># 只存最终 adapter, 省 IO</span>
        bf16=torch.cuda.is_bf16_supported(),
        report_to=<span>"none"</span>,
        remove_unused_columns=<span>False</span>,
    )
    trainer = Trainer(
        model=model, args=targs,
        train_dataset=ds[<span>"train"</span>], eval_dataset=ds[<span>"test"</span>],
        data_collator=DataCollatorForSeq2Seq(tokenizer, padding=<span>True</span>, pad_to_multiple_of=<span>8</span>),
    )

    <span># ---------- 5. 训练 + 保存 LoRA ----------</span>
    trainer.train()
    os.makedirs(ADAPTER_DIR, exist_ok=<span>True</span>)
    model.save_pretrained(ADAPTER_DIR)
    tokenizer.save_pretrained(ADAPTER_DIR)
    <span>print</span>(<span>f"LoRA adapter 已保存: <span>{ADAPTER_DIR}</span>"</span>)

    <span># ---------- 6. 合并 LoRA + 基座, 保存完整模型 (步骤2-3) ----------</span>
    <span># 4bit 模型不能直接 merge -> 重新以 fp16 加载基座再合并</span>
    <span>del</span> model, trainer
    torch.cuda.empty_cache()

    fp16_model = AutoModelForCausalLM.from_pretrained(
        args.base_model, torch_dtype=torch.float16, device_map=<span>"auto"</span>,
        trust_remote_code=<span>True</span>)
    merged = PeftModel.from_pretrained(fp16_model, ADAPTER_DIR)
    merged = merged.merge_and_unload()
    os.makedirs(MERGED_DIR, exist_ok=<span>True</span>)
    merged.save_pretrained(MERGED_DIR, safe_serialization=<span>True</span>)   <span># 存 .safetensors</span>
    tokenizer.save_pretrained(MERGED_DIR)
    <span>print</span>(<span>f"合并模型已保存: <span>{MERGED_DIR}</span>"</span>)
    <span>print</span>(<span>"验证: python inference.py --model"</span>, MERGED_DIR)


<span>if</span> __name__ == <span>"__main__"</span>:
    main()

8.3 推理验证 inference.py(全文)

<span># -*- coding: utf-8 -*-</span>
<span>"""
步骤2-1:基础模型加载验证(20分)
从 /home/user/workspace/QwenPretrain/ 加载 QWEN3 模型和 Tokenizer,
当输入 "你是谁?" 时输出合理自我介绍。裁判用此脚本检查基础模型可正常调用。

用法:
    python inference.py                     # 默认问 "你是谁?"
    python inference.py --q "什么是次贷危机?"
    python inference.py --model /path/to/checkpoint-best   # 也可验证微调后模型
依赖:
    pip install torch transformers
"""</span>
<span>import</span> argparse
<span>import</span> torch
<span>from</span> transformers <span>import</span> AutoModelForCausalLM, AutoTokenizer

MODEL_PATH = <span>"/home/user/workspace/QwenPretrain"</span>   <span># 考场: /home/user/workspace/QwenPretrain/</span>


<span>def</span> <span>main</span>():
    ap = argparse.ArgumentParser()
    ap.add_argument(<span>"--model"</span>, default=MODEL_PATH, <span>help</span>=<span>"模型目录"</span>)
    ap.add_argument(<span>"--q"</span>, default=<span>"你是谁?"</span>, <span>help</span>=<span>"输入问题"</span>)
    ap.add_argument(<span>"--max-new"</span>, <span>type</span>=<span>int</span>, default=<span>256</span>)
    args = ap.parse_args()

    <span># 1. 加载 tokenizer 和模型(GPU 可用则用 GPU)</span>
    tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=<span>True</span>)
    device = <span>"cuda"</span> <span>if</span> torch.cuda.is_available() <span>else</span> <span>"cpu"</span>
    model = AutoModelForCausalLM.from_pretrained(
        args.model,
        torch_dtype=torch.float16 <span>if</span> device == <span>"cuda"</span> <span>else</span> torch.float32,
        device_map=device,
        trust_remote_code=<span>True</span>,
    )
    model.<span>eval</span>()
    <span>print</span>(<span>f"[已加载模型] <span>{args.model}</span>  设备: <span>{device}</span>"</span>)

    <span># 2. 用 chat 模板构造输入(Qwen3 是对话模型,直接拼字符串效果差)</span>
    <span># 注意 transformers 5.x: apply_chat_template 返回 BatchEncoding, 取 input_ids</span>
    messages = [{<span>"role"</span>: <span>"user"</span>, <span>"content"</span>: args.q}]
    enc = tokenizer.apply_chat_template(
        messages, add_generation_prompt=<span>True</span>,
        tokenize=<span>True</span>, return_tensors=<span>"pt"</span>, return_dict=<span>True</span>,
        enable_thinking=<span>False</span>,          <span># 关闭 qwen3 思考模式, 直接出答案</span>
    )
    inputs = enc[<span>"input_ids"</span>].to(model.device)

    <span># 3. 推理</span>
    <span>with</span> torch.no_grad():
        out = model.generate(
            inputs,
            max_new_tokens=args.max_new,
            do_sample=<span>False</span>,            <span># 贪心解码, 输出稳定</span>
            temperature=<span>None</span>, top_p=<span>None</span>, top_k=<span>None</span>,   <span># do_sample=False 时禁用采样参数</span>
            pad_token_id=tokenizer.eos_token_id,
        )
    response = tokenizer.decode(out[<span>0</span>][inputs.shape[<span>1</span>]:], skip_special_tokens=<span>True</span>)
    <span># 若仍带思考标签则剥离</span>
    <span>import</span> re
    response = re.sub(<span>r"<think>.*?</think>"</span>, <span>""</span>, response, flags=re.DOTALL).strip()

    <span>print</span>(<span>f"问题: <span>{args.q}</span>"</span>)
    <span>print</span>(<span>f"回答: <span>{response}</span>"</span>)


<span>if</span> __name__ == <span>"__main__"</span>:
    main()

8.4 冒烟自检 self_check.py(全文)

<span># -*- coding: utf-8 -*-</span>
<span>"""
步骤3-3:本地冒烟测试与性能记录(15分)
加载量化后的 GGUF 模型, 回答 3 个固定问题, 记录输出与推理速度,
与原始模型(FP16, transformers 加载)对比, 写入 self_check_report.txt

用法:
    pip install llama-cpp-python
    python self_check.py                      # 默认测 model_int4.gguf
依赖: pip install llama-cpp-python transformers torch
"""</span>
<span>import</span> json
<span>import</span> time
<span>from</span> pathlib <span>import</span> Path

<span>from</span> llama_cpp <span>import</span> Llama

GGUF_PATH = Path(<span>"/home/user/workspace/model_b/quantization/model_int4.gguf"</span>)
BASE_GGUF = Path(<span>"/home/user/workspace/model_b/quantization/base.gguf"</span>)  <span># FP16 对照(存在则测)</span>
REPORT_PATH = Path(<span>"/home/user/workspace/model_b/quantization/self_check_report.txt"</span>)

QUESTIONS = [
    <span>"你是谁?"</span>,
    <span>"请简要介绍人工智能。"</span>,
    <span>"什么是金融危机?"</span>,
]


<span>def</span> <span>ask</span>(<span>llm, question, max_new=<span>96</span></span>):
    <span>"""单次推理, 返回 (回答文本, 生成token数, 耗时秒)"""</span>
    prompt = <span>f"### 指令:\n<span>{question}</span>\n### 回答:\n"</span>   <span># 与微调模板一致</span>
    t0 = time.time()
    out = llm(prompt, max_tokens=max_new, temperature=<span>0.1</span>,
              stop=[<span>"###"</span>], echo=<span>False</span>)
    dt = time.time() - t0
    text = out[<span>"choices"</span>][<span>0</span>][<span>"text"</span>].strip()
    n_tok = out[<span>"choices"</span>][<span>0</span>].get(<span>"tokens_evaluated"</span>, <span>0</span>) <span>or</span> <span>len</span>(text) // <span>2</span>
    <span>return</span> text, n_tok, dt


<span>def</span> <span>bench</span>(<span>path, n_ctx=<span>2048</span>, n_gpu_layers=<span>99</span></span>):   <span># GPU offload 全层</span>
    <span>print</span>(<span>f"\n===== 测试: <span>{path.name}</span> ====="</span>)
    llm = Llama(model_path=<span>str</span>(path), n_ctx=n_ctx,
                n_gpu_layers=n_gpu_layers, verbose=<span>False</span>)
    results = []
    <span>for</span> q <span>in</span> QUESTIONS:
        text, n_tok, dt = ask(llm, q)
        speed = n_tok / dt <span>if</span> dt > <span>0</span> <span>else</span> <span>0.0</span>
        results.append({<span>"q"</span>: q, <span>"a"</span>: text, <span>"tokens"</span>: n_tok, <span>"sec"</span>: dt, <span>"speed"</span>: speed})
        <span>print</span>(<span>f"Q: <span>{q}</span>\nA: <span>{text[:<span>80</span>]}</span>...\n   [<span>{n_tok}</span> tok / <span>{dt:<span>.1</span>f}</span>s = <span>{speed:<span>.1</span>f}</span> tok/s]"</span>)
    <span>return</span> results


<span>def</span> <span>main</span>():
    lines = [<span>"="</span> * <span>60</span>, <span>"量化模型本地自检报告 (self_check_report.txt)"</span>, <span>"="</span> * <span>60</span>]
    int4 = bench(GGUF_PATH)
    base = <span>None</span>
    <span>if</span> BASE_GGUF.exists():
        base = bench(BASE_GGUF)   <span># FP16 GGUF 作为"原始模型"对照</span>

    lines.append(<span>f"\n量化模型: <span>{GGUF_PATH.name}</span> (INT4 q4_0)"</span>)
    s4 = [r[<span>"speed"</span>] <span>for</span> r <span>in</span> int4]
    avg4 = <span>sum</span>(s4) / <span>len</span>(s4)
    lines.append(<span>f"平均推理速度: <span>{avg4:<span>.1</span>f}</span> tok/s"</span>)
    lines.append(<span>f"文件大小: <span>{GGUF_PATH.stat().st_size/<span>1e9</span>:<span>.2</span>f}</span> GB"</span>)
    <span>for</span> r <span>in</span> int4:
        lines.append(<span>f"  Q: <span>{r[<span>'q'</span>]}</span>"</span>)
        lines.append(<span>f"  A: <span>{r[<span>'a'</span>][:<span>120</span>]}</span>"</span>)
        lines.append(<span>f"  (<span>{r[<span>'tokens'</span>]}</span> tok / <span>{r[<span>'sec'</span>]:<span>.1</span>f}</span>s = <span>{r[<span>'speed'</span>]:<span>.1</span>f}</span> tok/s)"</span>)
        ok = <span>"正常"</span> <span>if</span> <span>len</span>(r[<span>"a"</span>]) > <span>4</span> <span>and</span> <span>not</span> _garbled(r[<span>"a"</span>]) <span>else</span> <span>"异常!"</span>
        lines.append(<span>f"  输出质量: <span>{ok}</span>"</span>)

    <span>if</span> base:
        sb = [r[<span>"speed"</span>] <span>for</span> r <span>in</span> base]
        avgb = <span>sum</span>(sb) / <span>len</span>(sb)
        lines.append(<span>f"\n原始模型对照: base.gguf (FP16)"</span>)
        lines.append(<span>f"平均推理速度: <span>{avgb:<span>.1</span>f}</span> tok/s"</span>)
        lines.append(<span>f"文件大小: <span>{BASE_GGUF.stat().st_size/<span>1e9</span>:<span>.2</span>f}</span> GB"</span>)
        lines.append(<span>f"压缩比: <span>{BASE_GGUF.stat().st_size/GGUF_PATH.stat().st_size:<span>.1</span>f}</span>x"</span>)
        lines.append(<span>f"速度对比: INT4 是 FP16 的 <span>{avg4/avgb:<span>.2</span>f}</span>x"</span>)

    lines.append(<span>"\n结论: "</span> + (
        <span>"三个问题输出均连贯, 无乱码/重复/语义断裂, 量化模型可正常工作。"</span>
        <span>if</span> <span>all</span>(<span>len</span>(r[<span>"a"</span>]) > <span>4</span> <span>and</span> <span>not</span> _garbled(r[<span>"a"</span>]) <span>for</span> r <span>in</span> int4)
        <span>else</span> <span>"存在异常输出, 需返回步骤3-2调整量化参数!"</span>))

    REPORT_PATH.write_text(<span>"\n"</span>.join(lines), encoding=<span>"utf-8"</span>)
    <span>print</span>(<span>f"\n报告已写入: <span>{REPORT_PATH}</span>"</span>)


<span>def</span> <span>_garbled</span>(<span>text</span>):
    <span>"""粗查乱码: 常见替换字符/控制字符"""</span>
    <span>return</span> (<span>"�"</span> <span>in</span> text) <span>or</span> <span>any</span>(<span>0</span> < <span>ord</span>(c) < <span>9</span> <span>for</span> c <span>in</span> text)


<span>if</span> __name__ == <span>"__main__"</span>:
    main()