news 2026/9/20 3:22:56

PEFT 实战:Context-aware Prompt Tuning(CPT)——基于对抗方法的少样本上下文学习提示微调

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PEFT 实战:Context-aware Prompt Tuning(CPT)——基于对抗方法的少样本上下文学习提示微调

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/)。它取三条技术路线的各自长处:

  1. 继承 ICL 的拼接策略:将示例拼接在输入之前,构造统一的上下文;
  2. 引入 Prompt Tuning 式的学习:通过迭代优化对上下文嵌入(context embedding)进行精炼,从训练示例中提取更深层信息;
  3. 借鉴对抗攻击(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.pyCPTConfig,继承PromptLearningConfig
嵌入模块src/peft/tuners/cpt/model.pyCPTEmbedding,投影与损失计算
注册入口src/peft/tuners/cpt/init.pyregister_peft_method(name="cpt", ...)
模型集成src/peft/peft_model.py_cpt_forward前向与损失组装
类型枚举src/peft/utils/peft_types.pyCPT = "CPT"
示例examples/cpt_finetuning/cpt_train_and_inference.ipynbSST-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输入句子 tokeninput 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_typeTaskType必填仅支持TaskType.CAUSAL_LM,否则抛ValueError(config.py L83-L84)
cpt_token_idsOptional[list[int]]None(→[0]CPT 提示的 token ID 序列,长度即虚拟 token 数
cpt_maskOptional[list[int]]None(→ 全 1)应用于 CPT token 的掩码(对应上下文 attention mask)
cpt_tokens_type_maskOptional[list[int]]None(→ 全 1)每个 CPT token 的类型掩码(见 2.2 节)
opt_weighted_loss_typeOptional[Literal["none", "decay"]]"none"加权损失类型,"decay"启用指数衰减
opt_loss_decay_factorOptional[float]1.0损失权重指数衰减因子(Notebook 用0.95
opt_projection_epsilonOptional[float]0.1输入 token 的投影 epsilon(Notebook 用0.2
opt_projection_format_epsilonOptional[float]0.1输入/输出模板 token 的投影 epsilon(Notebook 用0.1
tokenizer_name_or_pathOptional[str]None用于 prompt tuning 初始化的 tokenizer 名称或路径

__post_init__中还有三条重要逻辑:

  1. 固定 PEFT 类型self.peft_type = PeftType.CPTnum_transformer_submodules = 1
  2. 任务类型校验:非CAUSAL_LM直接报错(测试 tests/test_cpt.py 也专门验证了SEQ_CLS与缺省task_type两种情况);
  3. 默认值与一致性校验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_embeddings

4.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-10

get_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把上述机制串起来:

  1. 从 kwargs 取出labelsinput_type_mask(缺省时全部置 4);
  2. get_prompt拿到上下文嵌入并与inputs_embeds拼接;
  3. 生成前缀标签cpt_token_ids,并把输入侧input_type_mask整体平移prefix_type_mask.max(),避免与上下文掩码冲突;
  4. 依据cpt_type_mask % 4 == 0判定标签位置,其余位置标签置-100
  5. 前向基座模型后,调用CPTEmbedding.calculate_loss覆写loss并返回。

这一流程意味着训练时必须把input_type_mask作为额外输入传给模型,评估(labels=None)时则直接返回基座输出。


五、端到端实战:SST-2 上的训练与评估

5.1 环境与数据集

Notebook 使用的基座模型为bigscience/bloom-1b7,数据集为 GLUE 的sst2load_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_idsattention_maskinput_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=500training_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)场景设计的,原因有二:

  1. 自注意力的二次复杂度:上下文示例越多,input_ids越长,self-attention 的显存开销随序列长度增长;
  2. 额外损失项:CPT 需要在上下文上计算带类型掩码与衰减权重的损失,进一步增加计算负担。

因此当数据集较大时,建议限制上下文示例数量,把剩余样本仅用于优化(即作为训练数据而非拼入上下文),从而在显存与效果之间取得平衡。这一约束也解释了为什么示例代码默认MAX_ICL_SAMPLES = 10per_device_train_batch_size = 1


七、测试与可复现性

仓库在 tests/test_cpt.py 中提供了可运行的验证用例,可直接作为最小复现脚本:

  • config_text(L46-L60):显式给出 8 个 token 的cpt_token_idscpt_maskcpt_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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/20 3:20:47

如何导出微信聊天记录永久保存:WeChatMsg 免费上手指南

如何导出微信聊天记录永久保存:WeChatMsg 免费上手指南 【免费下载链接】WeChatMsg 提取微信聊天记录,将其导出成HTML、Word、CSV文档永久保存,对聊天记录进行分析生成年度聊天报告 项目地址: https://gitcode.com/GitHub_Trending/we/WeCh…

作者头像 李华
网站建设 2026/9/20 3:17:23

MBA写作降AI率实测:8类工具原理、效果与避坑指南

MBA课程作业最让人头疼的事之一,就是案例分析报告写完之后,打开学校系统里的AI检测一跑,显示“疑似AI生成比例:42%”。更离谱的是,有些段落明明是自己一个字一个字敲的,也被标红。于是大家开始到处找所谓的…

作者头像 李华
网站建设 2026/9/20 3:16:50

RapidOCR 在 Python 3.12 装不上?2 步绕开依赖冲突

RapidOCR 在 Python 3.12 装不上?2 步绕开依赖冲突 【免费下载链接】RapidOCR 📄 Awesome OCR multiple programing languages toolkits based on ONNX Runtime, OpenVINO, MNN, PaddlePaddle, TensorRT and PyTorch. 项目地址: https://gitcode.com/G…

作者头像 李华
网站建设 2026/9/20 3:15:20

Bandizip 免费版实测:从下载安装到高效压缩配置全指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华