在大型语言模型落地过程中,全参数微调往往需要保存与主干模型同样大小的梯度、优化器状态和检查点文件。PEFT 系列方法通过冻结预训练权重,只训练少量新增参数,把显存和存储成本降到可接受范围。Prompt Tuning 和 Prefix Tuning 是其中两种容易被混淆的代表:一个把学习对象放在输入序列的嵌入层,另一个把可训练向量塞进每一层 Transformer 的注意力计算里。二者在推理阶段的行为并不完全相同,下面从推理链路、参数结构和部署差异三个角度展开。

一、PEFT推理的底层逻辑:冻结主干与参数注入
PEFT 的目标不是重新训练所有参数,而是通过少量可训练参数来适配下游任务。推理阶段,主模型权重从公共检查点加载,新增的 prompt 或 prefix 参数作为独立张量参与计算。这种设计让同一个基座模型可以服务多个任务,每个任务只需保存几十到几百兆的增量参数。相比全参数微调,这种方式可以显著降低多租户场景下的存储和切换成本。
从注入位置看,PEFT 方法可以分为输入端、层间和输出端三类。Prompt Tuning 属于输入端,Prefix Tuning 属于层间。输入端注入最直观,它把额外信息放进序列 token embedding 前面;层间注入则把额外状态拼到自注意力的 key 和 value 上,使模型每一层都能直接读取前缀信号。推理时二者的计算图长度不同:Prompt Tuning 只在首层前增加长度,Prefix Tuning 在每一层前向中增加额外的拼接维度。
理解这点对部署很重要。如果只需要微调一个极小的文本分类器,Prompt Tuning 更轻;如果要在生成任务中保持稳定的格式约束,Prefix Tuning 常常更有效。下面分别展开讨论这两种方法的实现细节。
二、Prompt Tuning:把提示词训练成连续向量
传统 prompt 是离散 token,比如“请判断情感倾向:”这类文本。Prompt Tuning 不用选择具体词,而是初始化一个长度为 L、维度为 hidden_size 的参数张量,把它当作提示词嵌入序列。前向传播时,这些向量与输入文本的词嵌入直接拼接,然后送入冻结的 Transformer。因为是连续向量,优化器可以像更新普通参数一样调整它们,而不受离散词汇表限制。
可训练参数量为 L × hidden_size。以 7B 模型、hidden_size=4096、prompt_length=20 为例,参数量约 8 万,远小于全参数微调。值得注意的是,Prompt Tuning 对初始化和学习率比较敏感,使用离散 prompt 的 embedding 作为初始化通常比随机初始化更稳定。代码实现如下:
import torch
import torch.nn as nn
class PromptTuningModel(nn.Module):
def __init__(self, base_model, prompt_length=20, hidden_size=4096):
super().__init__()
self.base_model = base_model
for param in self.base_model.parameters():
param.requires_grad = False
self.prompt_embeddings = nn.Parameter(
torch.randn(prompt_length, hidden_size) * 0.02
)
def forward(self, input_ids, attention_mask=None):
input_embeds = self.base_model.get_input_embeddings()(input_ids)
batch_size = input_embeds.size(0)
prompt_embeds = self.prompt_embeddings.unsqueeze(0).expand(
batch_size, -1, -1
)
combined_embeds = torch.cat([prompt_embeds, input_embeds], dim=1)
prompt_mask = torch.ones(
batch_size, self.prompt_embeddings.size(0),
dtype=attention_mask.dtype,
device=attention_mask.device
)
if attention_mask is not None:
attention_mask = torch.cat([prompt_mask, attention_mask], dim=1)
outputs = self.base_model(
inputs_embeds=combined_embeds,
attention_mask=attention_mask
)
return outputs
推理时只需把训练好的 prompt_embeddings 载入并拼到输入前,模型其他部分完全复用。相比 Prefix Tuning,Prompt Tuning 不改变每一层的计算结构,因此容易与 FlashAttention、量化推理等优化兼容。但它的缺点是 prompt 长度受限,且只影响首层输入,对深层语义的控制能力偏弱,在一些复杂生成任务上可能不够稳定。
三、Prefix Tuning:为每一层注入键值前缀
Prefix Tuning 把可训练参数放到每一层 Transformer 的注意力模块中。它初始化两组前缀向量:prefix_keys 和 prefix_values,形状都为 prefix_length × hidden_size。在每一层计算自注意力时,它们被拼接到当前 key 和 value 序列的前面,query 保持不变。这样模型每一层都能在注意力聚合中吸收前缀信息,而不仅仅是第一层。对于深层语义控制和生成质量,这种设计通常比 Prompt Tuning 更有效。
可训练参数量大约是 num_layers × prefix_length × 2 × hidden_size。比如 12 层 Transformer、prefix_length=10、hidden_size=768 时,新增约 18 万个参数,仍然极小。Prefix Tuning 的一个核心技巧是使用 MLP 重参数化来稳定训练:真正学习的可能是一个低秩向量,再通过 MLP 映射成每层的 prefix。推理时可以直接保存映射后的 prefix,丢弃 MLP,减少额外参数。下面的示例省略 MLP 步骤,直接训练 prefix 张量:
import torch
import torch.nn as nn
class PrefixAttention(nn.Module):
def __init__(self, hidden_size, num_heads, prefix_length=10):
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.prefix_length = prefix_length
self.prefix_keys = nn.Parameter(
torch.randn(prefix_length, hidden_size) * 0.02
)
self.prefix_values = nn.Parameter(
torch.randn(prefix_length, hidden_size) * 0.02
)
def forward(self, query, key, value):
batch_size = query.size(0)
prefix_k = self.prefix_keys.unsqueeze(0).expand(batch_size, -1, -1)
prefix_v = self.prefix_values.unsqueeze(0).expand(batch_size, -1, -1)
key = torch.cat([prefix_k, key], dim=1)
value = torch.cat([prefix_v, value], dim=1)
attn_weights = torch.matmul(query, key.transpose(-2, -1))
attn_weights = attn_weights / (self.head_dim ** 0.5)
attn_probs = torch.softmax(attn_weights, dim=-1)
attn_output = torch.matmul(attn_probs, value)
return attn_output
由于 Prefix Tuning 在每层注意力中插入了前缀,推理时的显存占用比 Prompt Tuning 略高。序列长度增加的部分等于 prefix_length,Key/Value 缓存也会相应增大。但对大多数短序列任务,这个开销可以接受。它的优势在于更直接地影响深层表示,在文本生成、摘要等任务中往往比 Prompt Tuning 更稳定,尤其是需要对输出格式进行长期约束时。
四、推理实现对比与选择建议
从推理角度对比,Prompt Tuning 的增量参数仅存在于 embedding 输入阶段,因此可以很方便地作为虚拟提示加入请求,模型内部各层完全不需要感知 prompt 存在。Prefix Tuning 需要在模型加载时逐层注册 prefix 参数,或者在转换 ONNX/TensorRT 时重写每层 attention 的拼接逻辑。对于已经高度封装的大模型推理框架,Prompt Tuning 的接入成本更低。
显存方面,Prefix Tuning 每层多出 prefix_length 个 token 的 KV 缓存。若 batch=1、prefix_length=10、层数=32、hidden_size=4096,多出的 KV 缓存大约为 10×32×2×4096×2字节(FP16),约 5MB,实际影响有限。但如果是长序列任务,Prefix Tuning 的固定前缀不会随输入长度进一步增加,所以相对总缓存占比会下降。Prompt Tuning 则在 KV 缓存上几乎不增加额外负担,因为它只扩展输入序列长度,且通常 prompt 长度较短。
选择建议:在少样本分类、实体抽取或追求极致部署简易度时优先用 Prompt Tuning;在生成任务、风格控制、领域适配且能接受少量结构改造时用 Prefix Tuning。两者可以与其他 PEFT 方法如 LoRA 叠加,但要注意参数初始化与学习率解耦,否则新增参数之间可能产生梯度干扰。推理部署时建议将 prompt 或 prefix 参数冻结为常量,避免构建反向图,从而进一步提升吞吐并减少显存占用。
PEFTPrompt TuningPrefix Tuning修改时间:2026-10-07 04:33:50