【Bug已解决】Add PSOFT to PEFT 解决方案
【Bug已解决】Add PSOFT to PEFT 解决方案一、现象长什么样PEFT 已经支持 LoRA、PrefixTuning、PromptTuning、P-Tuning、IA³、AdaLoRA 等多种方法但社区里有人提议再加入一种叫PSOFTSoft Prompt 类的高效微调方法。提这个 issue 的动机很实际现有的PromptTuning是把一段可学习的软提示向量拼到输入 embedding 最前面但它对所有层都只在输入处加某些任务上不如“在每层都注入软提示”的方法用户想要一种比 PrefixTuning 更轻、比 PromptTuning 更灵活的软提示变体并且希望它和 LoRA 一样能通过get_peft_model/PeftModel.from_pretrained统一管理现在如果自己实现得手写nn.Module去改写模型forward、手动挂到 embedding 之后无法复用 PEFT 的print_trainable_parameters/merge_and_unload/load_adapter等生态复用save_pretrained时软提示权重没有统一的adapter_config.json描述跨模型加载容易键名错配参考前面关于get_peft_model与load_adapter键名不一致的讨论。所以“Add PSOFT to PEFT”本质是一个集成方法论问题如何把一个自定义的软提示方法按 PEFT 的扩展规范接进去让它成为一等公民。二、背景PEFT 的扩展规范以现有 tuner 为参考需要三件套XxxConfig继承PeftConfig描述超参如软提示长度soft_prompt_length、初始化方式、作用在哪些层、是否只作用在输入层。XxxModel继承BaseTuner/PreTuner持有被包装的基座管理注入的软提示参数重写forward把软提示拼到合适位置。mapping.py注册把PeftType.PSOFT映射到XxxModel让get_peft_model/PeftModel能通过peft_type字符串查到它。软提示的核心数学很简单构造一个可学习矩阵P ∈ ℝ^{L×d}L 是软提示长度d 是 embedding 维度在输入 token 的 embedding 序列前拼接P即E [P; E_input]其余前向不变。区别在于 PSOFT 可以指定“只在第 0 层注入”还是“在多个指定层注入”以及初始化是用随机、还是用真实词嵌入vocab 里某些 token 的 embedding做 warm-start。下面用可运行代码给出一个自包含的 PSOFT 实现不依赖 PEFT 内部 API但结构对齐其规范以及一个把它接进 PEFT 风格的骨架。三、根因这里的“问题”是缺失集成规范导致的重复造轮子没有统一的 Config/Tuner/Mapping 三件套每个用户自己重写 forward无法复用 PEFT 生态。软提示参数没有统一 state_dict 命名跨模型save/load键名错配。与 LoRA 等方法无法组合不能在一个模型上既用 LoRA 又用 PSOFT。修复方向按 PEFT 的扩展规范实现PSOFTConfigPSOFTModelmapping注册让软提示成为可管理的一等公民并能与现有方法组合。四、最小可运行复现下面给一个自包含的 PSOFT 实现软提示拼到输入 embedding 前可选 vocab warm-start可独立运行验证数学正确性。import torch import torch.nn as nn import torch.nn.functional as F class PSOFTConfig: def __init__(self, soft_prompt_length10, init_from_vocabNone, target_layers(0,)): self.soft_prompt_length soft_prompt_length self.init_from_vocab init_from_vocab # 可选用某些 token 的 embedding 初始化 self.target_layers target_layers class PSOFTWrapper(nn.Module): 把可学习软提示拼到指定层的输入 embedding 前。 def __init__(self, base_model, config: PSOFTConfig, vocab_embedNone): super().__init__() self.base base_model d base_model.get_input_embeddings().embedding_dim # 软提示参数 [L, d] self.soft_prompt nn.Parameter(torch.randn(config.soft_prompt_length, d) * 0.01) if config.init_from_vocab is not None and vocab_embed is not None: # 用真实词嵌入 warm-start取前 L 个 token with torch.no_grad(): idx config.init_from_vocab[:config.soft_prompt_length] self.soft_prompt.copy_(vocab_embed.weight[idx]) def forward(self, input_ids, layer_idx0): emb self.base.get_input_embeddings()(input_ids) # [B, T, d] if layer_idx in self.target_layers_tuple(): P self.soft_prompt.unsqueeze(0).expand(emb.size(0), -1, -1) emb torch.cat([P, emb], dim1) # 前拼软提示 return emb def target_layers_tuple(self): # 简化这里假设只在输入层注入多层的实现需改 base 各层 forward return (0,) # ---- 验证软提示确实改变了 embedding 序列长度 ---- class ToyModel(nn.Module): def __init__(self, vocab100, d16): super().__init__() self.emb nn.Embedding(vocab, d) def get_input_embeddings(self): return self.emb torch.manual_seed(0) base ToyModel() cfg PSOFTConfig(soft_prompt_length5) wrap PSOFTWrapper(base, cfg) ids torch.randint(0, 100, (2, 8)) # [B, T8] out wrap(ids, layer_idx0) print(原始 embedding 形状:, base.emb(ids).shape) # [2, 8, 16] print(加软提示后形状: , out.shape) # [2, 13, 16] 85 print(可训练参数含 soft_prompt:, any(soft_prompt in n for n, p in wrap.named_parameters() if p.requires_grad))运行后输出序列长度从 8 变成 13加了 5 个软提示 token且soft_prompt是唯一可训练参数基座冻结验证了 PSOFT 的核心行为。五、解决方案第一层最小直接修复修复 1用 vocab warm-start 初始化软提示随机初始化软提示收敛慢用真实词嵌入如任务相关 tokenwarm-start 更快cfg PSOFTConfig(soft_prompt_length5, init_from_vocabtorch.arange(5)) wrap PSOFTWrapper(base, cfg, vocab_embedbase.emb)修复 2冻结基座只训练软提示for p in base.parameters(): p.requires_grad_(False) # soft_prompt 保持 requires_gradTrue wrap.train()修复 3注意注意力掩码要同步延长拼接软提示后序列变长若模型用因果掩码/位置编码要相应扩展否则位置错位# 软提示位置用独立位置 id如 0..L-1原 token 位置整体 L pos_ids torch.cat([torch.arange(L), torch.arange(T) L])六、解决方案第二层结构性改进改进 1按 PEFT 规范写成PSOFTConfigPSOFTModelfrom peft import PeftConfig, PeftType class PSOFTConfig(PeftConfig): def __init__(self, soft_prompt_length10, init_from_vocabNone, **kwargs): super().__init__(peft_typePeftType.PSOFT, **kwargs) self.soft_prompt_length soft_prompt_length self.init_from_vocab init_from_vocab # PEFTModel 风格的包装示意 from peft.tuners import BaseTuner class PSOFTModel(BaseTuner): def __init__(self, model, config, adapter_namedefault): super().__init__(model, config, adapter_name) # 在这里注入 soft_prompt 参数、重写 forward 拼接逻辑 def forward(self, *args, **kwargs): # 把软提示拼进 embedding 再交给 base return self.model(*args, **kwargs)改进 2在mapping.py注册# peft/mapping.py 中增加 from .psft import PSOFTModel PEFT_TYPE_TO_MODEL_MAPPING[PeftType.PSOFT] PSOFTModel这样get_peft_model(base, PSOFTConfig(...))就能通过peft_type查到并实例化。改进 3与 LoRA 组合软提示作用于 embedding 入口LoRA 作用于线性层两者正交可在同一模型并存model get_peft_model(base, LoraConfig(r8, target_modules[q_proj,v_proj])) model.add_adapter(soft1, PSOFTConfig(soft_prompt_length10)) model.set_adapter([default, soft1]) # 同时激活七、解决方案第三层断言 / CI 守护import torch import torch.nn as nn import pytest class ToyModel(nn.Module): def __init__(self, vocab100, d16): super().__init__() self.emb nn.Embedding(vocab, d) def get_input_embeddings(self): return self.emb class PSOFTWrapper(nn.Module): def __init__(self, base, length): super().__init__() self.base base d base.get_input_embeddings().embedding_dim self.soft_prompt nn.Parameter(torch.randn(length, d) * 0.01) def forward(self, ids): emb self.base.get_input_embeddings()(ids) P self.soft_prompt.unsqueeze(0).expand(emb.size(0), -1, -1) return torch.cat([P, emb], dim1) def _trainable(m): return sum(p.numel() for p in m.parameters() if p.requires_grad) def test_soft_prompt_extends_seq(): torch.manual_seed(0) base ToyModel() wrap PSOFTWrapper(base, 5) out wrap(torch.randint(0, 100, (2, 8))) assert out.shape (2, 13, 16) def test_only_soft_prompt_trainable(): torch.manual_seed(1) base ToyModel() wrap PSOFTWrapper(base, 5) for p in base.parameters(): p.requires_grad_(False) assert _trainable(wrap) 5 * 16 def test_warm_start_changes_init(): torch.manual_seed(2) base ToyModel() w1 PSOFTWrapper(base, 3) with torch.no_grad(): w1.soft_prompt.copy_(base.emb.weight[:3]) assert torch.allclose(w1.soft_prompt, base.emb.weight[:3])这三个测试守护“软提示扩展序列、仅软提示可训练、warm-start 生效”。八、排查清单要把 PSOFT 接进 PEFT 时按序查三件套齐全PSOFTConfig继承PeftConfigPSOFTModel继承BaseTunermapping注册。state_dict 命名统一软提示参数用固定键名如soft_prompt.default避免跨模型加载错配。基座冻结只训练soft_prompt否则就失去“高效”意义。位置编码/掩码同步序列变长后位置 id 与注意力掩码要相应扩展。vocab warm-start用真实词嵌入初始化加速收敛。与 LoRA 组合两者正交可并存用add_adapterset_adapter管理。复用生态接好后即可用print_trainable_parameters/save_pretrained/load_adapter。测试守护序列扩展、仅软提示可训练、warm-start 三项必须有测试。九、小结Add PSOFT to PEFT的本质不是代码崩溃而是集成规范缺失用户各自重写 forward、软提示参数没有统一命名、无法复用 PEFT 生态、也无法与 LoRA 等方法组合。最小修复是实现自包含 PSOFT软提示拼到 embedding 前、vocab warm-start、冻结基座、掩码同步结构性改进是按 PEFT 规范写PSOFTConfigPSOFTModelmapping注册让它成为一等公民并能与 LoRA 组合最后用测试守护“软提示扩展序列、仅软提示可训练、warm-start 生效”。这样 PSOFT 就能像其它 PEFT 方法一样被get_peft_model统一管理。