PEFT 实战:Context-aware Prompt Tuning(CPT)——基于对抗方法的少样本上下文学习提示微调
【免费下载链接】peft🤗 PEFT: State-of-the-art Parameter-Efficient Fine-Tuning.项目地址: https://gitcode.com/gh_mirrors/pe/peft
本文以 PEFT 仓库中 examples/cpt_finetuning/README.md 为骨架,结合 CPT 源码实现、官方示例 Notebook 与 测试用例,系统讲解如何在少样本分类场景下用 CPT(Context-aware Prompt Tuning)训练因果语言模型:从模板化数据拼接、
cpt_tokens_type_mask类型掩码的语义,到CPTConfig全部关键参数、投影梯度下降(PGD)与衰减损失的底层实现,再到基于 Hugging FaceTrainer的完整训练与评估流程。读完本文,你将能够直接在 PEFT 中复现并定制自己的 CPT 少样本训练方案。
一、CPT 是什么:融合 ICL、Prompt Tuning 与对抗攻击
1.1 背景:少样本学习的两种路径与各自的痛点
大语言模型(LLM)的少样本学习主要有两条技术路线:
- 基于优化的方法(optimization-based):对模型参数进行梯度更新。缺点是数据量有限而待更新参数众多,极易过拟合;
- 上下文学习(In-Context Learning, ICL):把若干示例拼接到输入之前,无需更新参数。缺点在于通常精度低于优化方法,且对示例的选择、顺序与格式高度敏感。
1.2 CPT 的核心思想
Context-aware Prompt Tuning(CPT)正是为了同时缓解上述两个问题而提出的(论文 arXiv:2410.17222,仓库对应实现见 src/peft/tuners/cpt/)。它取三条技术路线的各自长处:
- 继承 ICL 的拼接策略:将示例拼接在输入之前,构造统一的上下文;
- 引入 Prompt Tuning 式的学习:通过迭代优化对上下文嵌入(context embedding)进行精炼,从训练示例中提取更深层信息;
- 借鉴对抗攻击(adversarial attacks):依据上下文中存在的标签调整输入,同时保留用户数据本身的语义价值。
CPT 只优化上下文中特定 token 的嵌入,模型其余部分保持完全冻结(这与 Prompt Tuning 只训练虚拟 token 的思路一脉相承,但优化对象是真实示例的上下文嵌入)。为了防止过拟合并保证优化稳定性,CPT 采用投影梯度下降(Projected Gradient Descent, PGD),把 token 嵌入约束在接近其原始值的范围内,从而保护上下文质量。
仓库官方文档 docs/source/package_reference/cpt.md 还补充了另一个关键动机:为缓解近因偏差(recency bias)——即上下文靠后的示例往往比靠前的示例更被模型优先关注——CPT 在损失上引入了衰减系数(decay loss factor),本文第四章会结合源码详解。
1.3 仓库中的实现位置
| 组件 | 文件 | 说明 |
|---|---|---|
| 配置类 | src/peft/tuners/cpt/config.py | CPTConfig,继承PromptLearningConfig |
| 嵌入模块 | src/peft/tuners/cpt/model.py | CPTEmbedding,投影与损失计算 |
| 注册入口 | src/peft/tuners/cpt/init.py | register_peft_method(name="cpt", ...) |
| 模型集成 | src/peft/peft_model.py | _cpt_forward前向与损失组装 |
| 类型枚举 | src/peft/utils/peft_types.py | CPT = "CPT" |
| 示例 | examples/cpt_finetuning/cpt_train_and_inference.ipynb | SST-2 情感二分类完整流程 |
| 测试 | tests/test_cpt.py | 初始化与训练正确性验证 |
在 PEFT 中,PeftType.CPT被归类为 prompt learning 方法(CPTConfig.is_prompt_learning = True),并与 Prompt Tuning、Multitask Prompt Tuning 共享同一套 prompt encoder 装配逻辑(见 src/peft/peft_model.py)。
二、数据准备与拼接:模板化 Tokenization
2.1 模板的角色
模板定义了输入-输出对的结构,使模型能在统一上下文中理解任务。原文档给出三类要素:
- 输入模板(Input Templates):如
"input: {sentence}",{sentence}占位符替换为真实输入文本; - 输出模板(Output Templates):如
"output: {label}",格式化标签(如positive/negative); - 分隔符(Separator Tokens):区分输入文本与标签、以及上下文中的不同示例。
Notebook 中定义的具体模板如下(见 cpt_train_and_inference.ipynb):
templates = { 'input': 'input: {}', # 输入模板,占位符替换为句子 'intra_seperator': ' ', # 示例内部:输入与输出之间的分隔符 'output': 'output: {}', # 输出模板,占位符替换为标签 'inter_seperator': '\n' # 示例之间:换行分隔符 }2.2 类型掩码:CPT 感知上下文结构的钥匙
原文档强调,CPT 借助编码在cpt_tokens_type_mask中的上下文结构来实现高效优化:根据 token 的不同角色区别对待,部分 token 被更新,其余 token 仅用于优化(作为约束或损失锚点)。掩码的数值语义由 Notebook 的preprocess_sentence与源码set_updated_tokens共同确认:
| 掩码值 | 角色 | 是否参与嵌入更新(后向钩子) | 投影 epsilon 类型 |
|---|---|---|---|
0 | 分隔符 / EOS | 否(掩码取余后非 1/2/3) | 极小值1e-10 |
1 | 输入模板 token(如"input: ") | 是 | format epsilon |
2 | 输入句子 token | 是 | input epsilon |
3 | 输出模板 token(如"output: ") | 是 | format epsilon |
4 | 标签 token(如positive) | 否(保持不变) | 极小值1e-10 |
单个样本的掩码由 cpt_train_and_inference.ipynb 中的CPTDataset.preprocess_sentence按顺序生成:
input_type_mask = ( [1] * len(input_template_tokenized_part1) # 输入模板前半 "input: " + [2] * len(input_tokenized) # 输入句子 + [1] * len(input_template_tokenized_part2) # 输入模板后半 + [0] * len(sep_tokenized) # 输入/输出间分隔符 + [3] * len(label_template_part1_tokenized) # 输出模板前半 "output: " + [4] * len(label_tokenized) # 标签文本 + [3] * len(label_template_part2_tokenized) # 输出模板后半 + [0] * len(eos) # 句末 EOS ) assert len(input_type_mask) == len(input_ids) == len(attention_mask)这条掩码与论文3.1 / 3.2 / 3.3 节一一对应:区分上下文各组成部分的结构(3.1)、仅更新特定 token(3.2)、按 token 类型施加不同投影范数(3.3)。
2.3 上下文拼接:多示例的掩码递增机制
CPT 会把多个示例依次拼接为一段长上下文(ICL 风格)。拼接时掩码不能简单复制,否则模型无法区分"这是第几个示例的标签"。Notebook 采用按示例递增 4的策略:
context_ids = [] context_attention_mask = [] context_input_type_mask = [] first_type_mask = 0 for i in range(len(context_dataset)): context_ids += cpt_context_dataset[i]['input_ids'] context_attention_mask += cpt_context_dataset[i]['attention_mask'] context_input_type_mask += [ i + first_type_mask if i > 0 else 0 # 类型索引动态递增 for i in cpt_context_dataset[i]['input_type_mask'] ] first_type_mask += 4 # 每个示例后偏移 +4这样第 k 个示例的掩码整体落在[4k, 4k+4]区间,对 4 取余后仍能还原其角色语义(这正是源码中大量使用remainder(mask, 4)判断类型的原因)。其中0保持不变,确保分隔符/EOS 永远不会被更新。
2.4 数据规模与 quick_review
Notebook 默认使用小规模子集保证快速运行(quick_review = True):前 10 条样本作为上下文示例(MAX_ICL_SAMPLES = 10),接下来 100 条作为训练集(NUM_TRAINING_SAMPLES = 100),验证集取前 100 条用于评估。完整评估时可将quick_review设为False。
三、CPTConfig 配置详解
CPTConfig继承自PromptLearningConfig,在 src/peft/tuners/cpt/config.py 中定义。完整参数如下:
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
task_type | TaskType | 必填 | 仅支持TaskType.CAUSAL_LM,否则抛ValueError(config.py L83-L84) |
cpt_token_ids | Optional[list[int]] | None(→[0]) | CPT 提示的 token ID 序列,长度即虚拟 token 数 |
cpt_mask | Optional[list[int]] | None(→ 全 1) | 应用于 CPT token 的掩码(对应上下文 attention mask) |
cpt_tokens_type_mask | Optional[list[int]] | None(→ 全 1) | 每个 CPT token 的类型掩码(见 2.2 节) |
opt_weighted_loss_type | Optional[Literal["none", "decay"]] | "none" | 加权损失类型,"decay"启用指数衰减 |
opt_loss_decay_factor | Optional[float] | 1.0 | 损失权重指数衰减因子(Notebook 用0.95) |
opt_projection_epsilon | Optional[float] | 0.1 | 输入 token 的投影 epsilon(Notebook 用0.2) |
opt_projection_format_epsilon | Optional[float] | 0.1 | 输入/输出模板 token 的投影 epsilon(Notebook 用0.1) |
tokenizer_name_or_path | Optional[str] | None | 用于 prompt tuning 初始化的 tokenizer 名称或路径 |
__post_init__中还有三条重要逻辑:
- 固定 PEFT 类型:
self.peft_type = PeftType.CPT,num_transformer_submodules = 1; - 任务类型校验:非
CAUSAL_LM直接报错(测试 tests/test_cpt.py 也专门验证了SEQ_CLS与缺省task_type两种情况); - 默认值与一致性校验:
cpt_token_ids缺省为[0],cpt_mask/cpt_tokens_type_mask缺省为全 1,且三者长度必须等于num_virtual_tokens = len(cpt_token_ids),否则抛ValueError(config.py L86-L100)。
Notebook 中的完整配置示例(SST-2 + BLOOM-1.7B):
from peft import CPTConfig, TaskType, get_peft_model config = CPTConfig( task_type=TaskType.CAUSAL_LM, cpt_token_ids=context_ids, # 拼接后的上下文 token ID cpt_mask=context_attention_mask, # 拼接后的注意力掩码 cpt_tokens_type_mask=context_input_type_mask, # 拼接后的类型掩码 opt_weighted_loss_type='decay', # 启用指数衰减损失 opt_loss_decay_factor=0.95, # 衰减因子 opt_projection_epsilon=0.2, # 输入 token 投影半径 opt_projection_format_epsilon=0.1, # 模板 token 投影半径 tokenizer_name_or_path=model_id, # 用于文本初始化 ) model = get_peft_model(base_model, config)注意:
cpt_token_ids传入的是整段拼接后的上下文token ID(可达数十字条样本 × 每条数十 token),因此num_virtual_tokens会远大于普通 Prompt Tuning 的虚拟 token 数——这也解释了为什么 CPT 更适合少样本场景(见第六章)。
四、源码级原理:CPTEmbedding 如何"只动该动的嵌入"
4.1 双嵌入结构:冻结基底 + 可学习增量
CPTEmbedding(src/peft/tuners/cpt/model.py)维护两套嵌入:
embedding(冻结):用cpt_token_ids对应的基座词嵌入初始化,requires_grad_(False),保证上下文语义不被破坏;delta_embedding(可学习):零初始化,前向时prompt_embeddings + delta_prompt_embeddings即最终上下文嵌入,训练只产生 delta。
def forward(self, indices): with torch.no_grad(): prompt_embeddings = self.embedding(indices) # 冻结的原始嵌入 self.delta_embedding.weight.data = self.get_projection() # 每次前向先做 PGD 投影 delta_prompt_embeddings = self.delta_embedding(indices) return prompt_embeddings + delta_prompt_embeddings4.2 选择性梯度更新(对应论文"不更新标签 token")
set_updated_tokens(model.py L87-L102)在delta_embedding.weight上注册后向钩子,按cpt_tokens_type_mask % 4决定哪些梯度被保留:
mask_input_template = torch.remainder(tensor_ICL_mask, 4) == 1 mask_input = torch.remainder(tensor_ICL_mask, 4) == 2 mask_output_template = torch.remainder(tensor_ICL_mask, 4) == 3 mask = mask_input_template | mask_input | mask_output_template def backward_hook(grad): grad = grad * mask.to(grad.device) # 只保留类型 1/2/3 的梯度 return grad即输入模板、输入句子、输出模板被更新,而标签 token(余数 0)与分隔符(0)的梯度被置零。标签包含宝贵的、不可修改的信息,因此始终保持原样——这正是原文档强调的"Refrain from Updating Label Tokens"。
4.3 类型感知投影范数(PGD 的落点)
get_epsilon(model.py L104-L124)先按token_dim对 epsilon 做归一化,再按类型分发:
normalized_format_eps = self.config.opt_projection_format_epsilon * torch.sqrt(torch.Tensor([self.config.token_dim / 2048])) normalized_input_eps = self.config.opt_projection_epsilon * torch.sqrt(torch.Tensor([self.config.token_dim / 2048])) epsilon[(mask > 0) & (remainder(mask, 4) == 1)] = normalized_format_eps # 输入模板 epsilon[(mask > 0) & (remainder(mask, 4) == 3)] = normalized_format_eps # 输出模板 epsilon[(mask > 0) & (remainder(mask, 4) == 2)] = normalized_input_eps # 输入句子 # 其余(标签/分隔符)保持 MIN_VALUE = 1e-10get_projection(model.py L126-L142)随即对 delta 嵌入做范数投影,把每个 token 的扰动限制在其类型对应的 epsilon 球内:
token_norm = torch.norm(new_embeddings_weights, p=2, dim=1) projection_mask = token_norm > 0 new_embeddings_weights[projection_mask] *= ( epsilon[projection_mask] / (token_norm[projection_mask].clamp(min=epsilon[projection_mask])) ).view(-1, 1)当token_norm ≤ epsilon时该式缩放因子为 1(不截断);当token_norm > epsilon时被缩放回 epsilon 半径。于是输入 token 允许更大的扰动(opt_projection_epsilon),模板 token 扰动更小(opt_projection_format_epsilon),标签几乎零扰动——既吸取了对抗方法的"受限扰动"思想,又通过保留用户原始数据保证泛化,从而降低少样本过拟合。
4.4 衰减损失:缓解近因偏差
calculate_loss(model.py L144-L202)先做标准的因果语言建模 shift,用CrossEntropyLoss(reduction="none", ignore_index=-100)计算每个 token 的损失;随后对每个样本中类型余数为 0 的标签位置按出现顺序施加指数衰减权重:
idx_labels = (shift_cpt_type_mask[i] > 0) & (shift_cpt_type_mask[i] % 4 == 0) labels_ids = shift_cpt_type_mask[i][idx_labels].unique() decay_value = 1 for label_mask_idx in torch.flip(labels_ids, [0]): # 从最后一个示例开始 exponential_decay[shift_cpt_type_mask[i] == label_mask_idx] = decay_value decay_value *= config.opt_loss_decay_factor if config.opt_weighted_loss_type == "decay": shift_labels_weights[i] *= exponential_decay由于掩码随示例递增,labels_ids按升序排列,flip后最后一个示例获得权重 1,越靠前的示例权重越小(0.95^k),从而对抗"上下文末尾示例被过度优先"的近因偏差。opt_weighted_loss_type="none"时则退化为普通加权平均。
4.5 前向装配:_cpt_forward
src/peft/peft_model.py 中的_cpt_forward把上述机制串起来:
- 从 kwargs 取出
labels与input_type_mask(缺省时全部置 4); - 用
get_prompt拿到上下文嵌入并与inputs_embeds拼接; - 生成前缀标签
cpt_token_ids,并把输入侧input_type_mask整体平移prefix_type_mask.max(),避免与上下文掩码冲突; - 依据
cpt_type_mask % 4 == 0判定标签位置,其余位置标签置-100; - 前向基座模型后,调用
CPTEmbedding.calculate_loss覆写loss并返回。
这一流程意味着训练时必须把input_type_mask作为额外输入传给模型,评估(labels=None)时则直接返回基座输出。
五、端到端实战:SST-2 上的训练与评估
5.1 环境与数据集
Notebook 使用的基座模型为bigscience/bloom-1b7,数据集为 GLUE 的sst2(load_dataset('glue', 'sst2')),并把数值标签映射为字符串:
def add_string_labels(example): example['label_text'] = "positive" if example['label'] == 1 else "negative" return example context_dataset = dataset['train'].select(range(MAX_ICL_SAMPLES)).map(add_string_labels) # 前 10 条做上下文 train_dataset = dataset['train'].select(range(MAX_ICL_SAMPLES, NUM_TRAINING_SAMPLES + MAX_ICL_SAMPLES)).map(add_string_labels) # 后 100 条训练模型加载与 CPT 装配:
base_model = AutoModelForCausalLM.from_pretrained( model_id, cache_dir='.', dtype=torch.float16, device_map='auto' ) model = get_peft_model(base_model, config) # config 见第三章5.2 自定义 Data Collator
由于 CPT 需要同时传递input_ids、attention_mask与input_type_mask(训练时还需labels),Notebook 定义了CPTDataCollatorForLanguageModeling:按 batch 内最大长度补齐三种序列、labels = input_ids.clone(),并在评估模式下附加sample_mask(用于按样本聚合指标)。
5.3 训练
training_args = TrainingArguments( output_dir='../.', use_cpu=False, auto_find_batch_size=False, learning_rate=1e-4, logging_steps=100, per_device_train_batch_size=1, # 上下文较长,batch 设 1 save_total_limit=1, remove_unused_columns=False, # 关键:保留 input_type_mask 列 num_train_epochs=5, fp16=True, save_strategy='no', report_to="none" ) trainer = Trainer( model=model, args=training_args, train_dataset=cpt_train_dataset, data_collator=CPTDataCollatorForLanguageModeling(tokenizer, training=True, mlm=False) ) trainer.train()Notebook 记录的本次运行结果为:global_step=500、training_loss≈0.0982、训练耗时约 90.7 秒(约 5.51 样本/秒,5 个 epoch),损失曲线从 step 100 的 0.4008 稳步下降到 step 500 的 0.0116。注意remove_unused_columns=False是必须项,否则input_type_mask会被Trainer默认清理导致前向失败。
5.4 评估
评估时逐条样本前向,通过input_type_mask == 4定位标签 token,再在候选标签词表(['negative', 'positive'])上取argmax判定类别:
model.eval() device = model.device for i in range(len(test_dataset)): input_ids, input_type_mask = cpt_test_dataset[i]['input_ids'], cpt_test_dataset[i]['input_type_mask'] outputs = model( input_ids=torch.Tensor(input_ids).long().to(device).view(1, -1), labels=torch.Tensor(input_ids).long().to(device).view(1, -1), input_type_mask=torch.Tensor(input_type_mask).long().to(device).view(1, -1) ) shifted_logits = outputs.logits[..., :-1, :].contiguous()[0, -len(input_ids) + 1:] shift_labels = torch.Tensor(input_ids).long().to(device).view(1, -1)[0, 1:].contiguous() shifted_input_type_mask = torch.Tensor(input_type_mask).long().to(device).view(1, -1)[..., 1:].contiguous() mask = shifted_input_type_mask.view(-1) == 4 # 只取标签位置 logit, label = shifted_logits[mask], shift_labels[mask] all_labels = torch.Tensor([tokenizer(w, add_special_tokens=False)["input_ids"] for w in ['negative', 'positive']]).long().to(device).view(-1) prediction = logit[0, all_labels].argmax() # 在标签词表内比较 prediction_text = 'negative' if prediction == 0 else 'positive'Notebook 在 100 条 quick_review 验证子集上记录到90.0% 的准确率(个别样本如"holden caulfield did it better ."被判错,属正常误差范围)。
六、局限性:面向少样本的设计取舍
原文档明确强调 CPT 是为少样本(few-shot)场景设计的,原因有二:
- 自注意力的二次复杂度:上下文示例越多,
input_ids越长,self-attention 的显存开销随序列长度增长; - 额外损失项:CPT 需要在上下文上计算带类型掩码与衰减权重的损失,进一步增加计算负担。
因此当数据集较大时,建议限制上下文示例数量,把剩余样本仅用于优化(即作为训练数据而非拼入上下文),从而在显存与效果之间取得平衡。这一约束也解释了为什么示例代码默认MAX_ICL_SAMPLES = 10、per_device_train_batch_size = 1。
七、测试与可复现性
仓库在 tests/test_cpt.py 中提供了可运行的验证用例,可直接作为最小复现脚本:
config_text(L46-L60):显式给出 8 个 token 的cpt_token_ids、cpt_mask与cpt_tokens_type_mask,使用"decay"损失与投影参数0.2 / 0.1;config_random(L63-L74):不指定 token ID,验证默认值路径;test_model_initialization_text/test_model_initialization_random(L216-L229):验证get_peft_model装配成功;test_model_initialization_wrong_task_type_raises(L232-L239):验证非CAUSAL_LM任务类型抛出ValueError。
测试使用的模型是peft-internal-testing/tiny-random-OPTForCausalLM,模板与 Notebook 一致(TEMPLATE = {"input": "input: {}", "intra_separator": " ", "output": "output: {}", "inter_separator": "\n"}),便于在无 GPU 的 CI 环境中快速回归。
八、引用
@article{ blau2025cpt, title={Context-Aware Prompt Tuning: Advancing In-Context Learning with Adversarial Methods}, author={Tsachi Blau, Moshe Kimhi, Yonatan Belinkov, Alexander Bronstein, Chaim Baskin}, journal={arXiv preprint arXiv:2410.17222}, year={2025} }如需进一步查阅 CPT 的官方 API 文档(CPTConfig/CPTEmbedding的 autodoc),可参见 docs/source/package_reference/cpt.md;与本方法关系最紧密的基线实现可对照 Prompt Tuning 文档 理解其继承关系。
【免费下载链接】peft🤗 PEFT: State-of-the-art Parameter-Efficient Fine-Tuning.项目地址: https://gitcode.com/gh_mirrors/pe/peft
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考