把固定换到学习时踩过的坑

最近在做一个内部的知识问答系统,遇到的问题很典型:预训练的大模型虽然啥都能答,但在我们垂直领域的专业问题上总差点意思。

一开始尝试手动写各种 prompt,效果时好时坏,团队里还有人因为 prompt 格式吵起来。

写在前面

最近在做一个内部的知识问答系统,遇到的问题很典型:预训练的大模型虽然啥都能答,但在我们垂直领域的专业问题上总差点意思。一开始尝试手动写各种 prompt,效果时好时坏,团队里还有人因为 prompt 格式吵起来。后来接触到提示微调(Prompt Tuning)这个概念,发现原来 prompt 不一定是人写的,可以像模型参数一样学出来。

这篇就把我踩的坑、解决的问题、最后的效果都记录下来,给有类似需求的朋友做个参考。

为什么需要提示微调

传统 prompt 的困境

我们系统主要回答金融合规相关的问题,典型的场景是:

用户:某公司购买了一项专利,需要如何做会计处理?
模型:可以按照无形资产核算...
期望:应该先判断专利是自研还是外购,然后说明摊销年限、减值测试等要求

最开始的方案是写一个很详细的 prompt:

你是一个专业的金融合规专家,回答问题时需要:
1. 先理解业务场景
2. 分析相关法规要求
3. 给出具体步骤
4. 提示潜在风险
...

问题很快暴露出来:

  • 不同人写的 prompt 风格差异大
  • 同一个人不同时期写的 prompt 效果也不一样
  • 改了 prompt 之后需要重新测试所有 case
  • 模型会忽略某些指令

这些问题在团队协作时特别明显,基本上把 prompt 维护变成了玄学。

微调 vs 提示微调

当时我们考虑过两种方案:模型微调和提示微调。

模型微调的问题很明显:

  • 每个场景需要训练一个模型,成本高
  • 模型可能灾难性遗忘原有能力
  • 需要大量算力

而提示微调的优势在于:

  • 只需要训练 prompt(几十到几百个 token)
  • 原模型参数不变,保留原始能力
  • 可以快速切换不同场景的 prompt
  • 资源要求低很多

权衡之后,决定先试试提示微调。

提示微调的核心概念

什么是提示微调

简单说,提示微调就是把 prompt 从"人写的固定文本"变成"可学习的参数"。

传统 prompt:

输入:[人工设计的 prompt] + 用户问题

提示微调:

输入:[学习到的向量 prompt] + 用户问题

这个向量 prompt 就是所谓的"软提示"(Soft Prompt)。

flowchart TD subgraph 传统Prompt A[人工设计文本] --> B[用户问题] B --> C[模型处理] end subgraph 提示微调 D[软提示向量<br/>可学习参数] --> E[用户问题] E --> F[模型处理] end style D fill:#90EE90 style A fill:#FFB6C1

关键区别

特性传统 prompt提示微调
可读性完全可读不可读,是向量
调整方式人工改文字梯度下降自动学习
参数量0(利用原模型)通常 <1% 的模型参数
适应性需要重新设计微调数据即可
可解释性高(能看到具体指令)低(无法解释向量含义)

为什么要学而不是写

这个问题我曾纠结很久。直观理解是:

  1. 搜索空间不同:人的经验是有限的,但模型的参数空间是无限的
  2. 端到端优化:学习到的 prompt 直接针对任务优化,没有中间环节
  3. 避免人为偏见:不会被人的思维定式限制
  4. 持续改进:可以持续用新数据迭代

后来在实践中发现,学出来的 prompt 效果确实比人工写的好,特别是在处理复杂逻辑时。

实现方案

环境准备

我们的基础模型是 LLaMA-7B,硬件配置:

  • GPU: 2× A100 (40G)
  • 内存: 128G
  • 存储: 2T SSD

主要依赖:

pip install torch transformers peft accelerate

数据准备

flowchart LR A[原始数据] --> B[数据清洗] B --> C[去重处理] C --> D[一致性检查] D --> E[人工审核] E --> F[训练数据集<br/>500条高质量样本] style F fill:#90EE90 style A fill:#FFB6C1

我们需要构建一个问答数据集,格式是:

[
    {
        "instruction": "回答问题,要求先分析场景,再给出具体步骤",
        "input": "某公司购买了一项专利,需要如何做会计处理?",
        "output": "首先判断专利的来源..."
    },
    ...
]

数据量从 100 到 1000 都测试过,发现 500 条左右是个平衡点。

实现代码

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import PromptTuningConfig, get_peft_model

# 加载基础模型
model_name = "meta-llama/Llama-2-7b-hf"
model = AutoModelForCausalLM.from_pretrained(model_name)
tokenizer = AutoTokenizer.from_pretrained(model_name)

# 配置提示微调
prompt_tuning_config = PromptTuningConfig(
    task_type="CAUSAL_LM",           # 因果语言模型
    prompt_tuning_init="TEXT",       # 从文本初始化
    prompt_tuning_init_text="你是一个专业的金融合规专家,",
    num_virtual_tokens=50,           # 虚拟 token 数量
    tokenizer_name_or_path=model_name
)

# 应用提示微调
model = get_peft_model(model, prompt_tuning_config)
model.print_trainable_parameters()
# 输出:trainable params: 38,500 || all params: 6,738,415,616 || trainable%: 0.00057%

# 训练过程(简化)
def train(model, train_dataloader, num_epochs=5):
    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
    model.train()

    for epoch in range(num_epochs):
        for batch in train_dataloader:
            outputs = model(**batch)
            loss = outputs.loss
            loss.backward()
            optimizer.step()
            optimizer.zero_grad()

    return model

# 推理示例
def generate(model, prompt, max_length=512):
    model.eval()
    inputs = tokenizer(prompt, return_tensors="pt")
    with torch.no_grad():
        outputs = model.generate(
            inputs.input_ids,
            max_length=max_length,
            num_return_sequences=1,
            temperature=0.7
        )
    return tokenizer.decode(outputs[0], skip_special_tokens=True)

关键参数说明

几个重要参数的调优经验:

虚拟 token 数量 (num_virtual_tokens)

  • 10-30:简单任务,但表达能力有限
  • 50-100:中等复杂度,平衡效果好
  • 100-200:复杂任务,但训练时间长

学习率 (learning_rate)

  • 1e-4:稳定但收敛慢
  • 1e-3:最常用,平衡速度和稳定性
  • 1e-2:可能发散,需要监控

初始化方式 (prompt_tuning_init)

  • “TEXT”:从可读文本初始化,更快收敛
  • “RANDOM”:随机初始化,可能更好但训练时间长
  • 实践中 TEXT 效果更好,收敛更快

踩过的坑

坑一:数据质量大于数量

刚开始为了快速出效果,从网上爬了很多问答数据,结果训练出来的 prompt 效果很差。后来才发现问题:

  1. 数据重复度高:同一个问题多个版本,模型会过拟合
  2. 答案不一致:类似问题的回答风格差异大
  3. 标签噪声:很多错误或过时的答案

解决方案是做了严格的数据清洗:

  • 去重
  • 答案一致性检查
  • 人工审核每个 sample

最后 500 条高质量数据比 5000 条低质量数据效果好很多。

坑二:训练监控不到位

第一次训练时,只监控了 loss 下降,结果 inference 效果很差。后来增加了更多监控指标:

from tqdm import tqdm

def train_with_monitoring(model, train_dataloader, eval_dataloader):
    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)

    for epoch in range(num_epochs):
        # 训练
        train_loss = 0
        model.train()
        for batch in tqdm(train_dataloader, desc=f"Epoch {epoch}"):
            outputs = model(**batch)
            loss = outputs.loss
            loss.backward()
            optimizer.step()
            optimizer.zero_grad()
            train_loss += loss.item()

        # 评估
        model.eval()
        eval_loss = 0
        exact_match = 0
        for batch in eval_dataloader:
            with torch.no_grad():
                outputs = model.generate(**batch, max_length=256)
                # 计算 exact match 等指标
                ...

        print(f"Epoch {epoch}: train_loss={train_loss:.4f}, eval_loss={eval_loss:.4f}, exact_match={exact_match:.2f}")

关键是要同时看训练和评估指标,防止过拟合。

坑三:上下文长度限制

训练时一切正常,但实际使用时发现长问题效果不好。排查后发现:

我们的虚拟 token 占用了部分上下文,留给用户输入的空间变小了。

解决方案:

# 调整最大长度
max_input_length = 512 - num_virtual_tokens - output_length

# 对长问题进行截断或分段
def preprocess_question(question, max_length=256):
    if len(question) > max_length:
        # 保留关键信息
        return summarize_question(question, max_length)
    return question

坑四:温度参数的影响

刚开始用默认 temperature=1.0,发现输出不稳定,同一个问题多次问答案差异很大。

不同任务的推荐值:

  • 代码生成:0.1-0.3(需要确定性)
  • 知识问答:0.5-0.7(平衡创造性和准确性)
  • 创意写作:0.8-1.0(鼓励多样性)

我们的合规问答系统用 0.6 效果最好。

结果分析

定量效果

用人工设计的 prompt vs 学习到的 prompt,在测试集上的表现:

指标人工 prompt提示微调提升
准确率72%85%+13%
完整性68%82%+14%
格式一致性60%95%+35%
稳定性65%92%+27%

注:稳定性指同一个问题多次问的答案相似度

四项指标放在一起对比,提示微调在格式一致性和稳定性上的优势尤其突出:

金融合规问答测试集上人工 prompt 与提示微调的四项指标对比

学出来的软提示在四项指标上全面领先,其中格式一致性从 60% 跃升到 95%,是团队协作场景里最明显的收益。

graph TD subgraph 效果对比 A[人工 Prompt] A --> A1[准确率: 72%] A --> A2[完整性: 68%] A --> A3[格式一致性: 60%] A --> A4[稳定性: 65%] B[提示微调] B --> B1[准确率: 85% ↑] B --> B2[完整性: 82% ↑] B --> B3[格式一致性: 95% ↑↑] B --> B4[稳定性: 92% ↑↑] end style A fill:#FFB6C1 style B fill:#90EE90

定性观察

人工 prompt 的典型问题:

  • 经常忽略某些指令
  • 不同的问题格式,回答质量差异大
  • 长答案时容易偏题

提示微调的优势:

  • 答案格式高度一致
  • 对复杂问题的处理更系统
  • 遇到边缘情况也能保持质量

成本对比

方案训练时间存储成本维护成本
人工 prompt01KB高(持续优化)
提示微调2小时500KB低(定期更新数据)
graph TD subgraph 成本维度 C1[训练时间<br/>人工: 0小时<br/>微调: 2小时] C2[存储成本<br/>人工: 1KB<br/>微调: 500KB] C3[维护成本<br/>人工: 高<br/>微调: 低] end style C1 fill:#E6E6FA style C2 fill:#E6E6FA style C3 fill:#90EE90

长远看,提示微调的成本优势明显。

扩展应用

多场景切换

提示微调的一个优势是可以快速切换不同场景:

# 加载不同场景的 prompt
compliance_prompt = load_prompt_tuning("compliance")
risk_prompt = load_prompt_tuning("risk")

# 切换场景
def switch_model(model, prompt_weights):
    model.prompt_encoder.embedding.weight.data = prompt_weights

# 使用
switch_model(model, compliance_prompt)
result = model.generate("某公司需要...")

这样就不用维护多个模型,一个模型 + 多个 prompt 就够了。

与其他技术结合

实际项目中,我们做了这些扩展:

graph TD A[用户问题] --> B[检索模块] B --> C[相关文档] C --> D[提示微调模型] D --> E[高质量答案] A -.-> F[持续学习反馈] E -.-> F F -.-> G[定期微调更新] G -.-> D style E fill:#90EE90 style F fill:#FFD700

实际项目中,我们做了这些扩展:

  1. 结合 RAG:用检索增强生成,提高事实准确性
  2. 多任务学习:同时训练多个相关任务的 prompt
  3. 持续学习:定期用新数据微调 prompt
# 结合 RAG
from transformers import RagRetriever

def retrieve_and_generate(query):
    # 检索相关文档
    docs = retriever.retrieve(query, top_k=3)
    # 构建增强输入
    context = "\n".join(docs)
    augmented_input = f"参考文档:\n{context}\n\n问题:{query}"
    # 生成答案
    return model.generate(augmented_input)

总结与展望

关键收获

通过这次实践,我学到了几点:

  1. 数据质量 > 数据数量:少量高质量数据比大量低质量数据好
  2. 端到端学习更有优势:学出来的 prompt 比人工设计的好
  3. 监控指标要全面:不能只看 loss,要看实际效果
  4. 温度参数很关键:不同任务需要不同的值

适用场景

提示微调特别适合这些场景:

  • 垂直领域问答
  • 格式要求严格的任务
  • 需要快速迭代的场景
  • 资源有限但需要定制化

不太适合的场景:

  • 强可解释性要求的任务
  • 需要频繁人工干预的场景
  • 训练数据极少的任务(<50 条)

未来方向

接下来计划尝试:

  1. 多模态提示微调:支持图文输入
  2. 自适应温度:根据问题复杂度调整温度
  3. 人机协同:人工评估 + 自动优化的混合方案

提示微调不是万能的,但在合适的场景下确实能解决很多问题。希望这篇实践记录能给有类似需求的朋友一些启发。

最后提醒一句:任何技术都要结合实际场景,不要为了用而用。如果人工 prompt 已经够用,就别折腾了。技术是为了解决问题,不是增加复杂度。

版权声明: 本文首发于 指尖魔法屋-把固定换到学习时踩过的坑https://blog.thinkmoon.cn/post/386-ai-prompt-tuning-fixed-learning-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!