在 Transformers 中定制模型组件:以 SAM 注意力机制拆分与 LoRA 微调实战
在 Transformers 中定制模型组件以 SAM 注意力机制拆分与 LoRA 微调实战【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers导读在 Transformers 中定制模型除了完整重写一个全新架构之外还有一条更轻量的路径——直接修改现成模型的组成部件。本文以图像分割模型 SAMSegment Anything为实战对象手把手演示如何将 Vision Encoder 中合并的qkv注意力投影拆分为独立的q、k、v线性层并在拆分后的q、v上应用 LoRALow-Rank Adaptation进行高效参数微调。读完本文你将掌握继承原注意力类 自定义加载钩子拆分权重 替换模型模块 挂接 PEFT 适配器这一整套可复用的模型定制方法论并理解其底层实现原理与源码位置。为什么选择修改组件而不是重写模型模型定制的核心诉求是以最小改动让一个预训练模型适配特定使用场景——例如新增一层、优化架构中的注意力机制。与从零编写新模型相比直接在 Transformers 模型上修改组件有一个关键优势定制结果是直接作用于模型对象本身的因此Trainer、PreTrainedModel以及 PEFT 库的整套生态能力训练循环、checkpoint 存取、适配器管理等都可以继续无缝使用无需另起炉灶。本文选用的示例对象是 Segment AnythingSAM——一种输入图像即可预测任意目标分割掩码的图像分割模型。SAM 的注意力机制将 query、key、value 合并为一个qkv投影以节省参数与计算量但这种合并结构恰恰阻碍了 LoRA 按模块精准挂接。因此实战的第一步就是把qkv拆开。迭代开发利器clear_import_cache 热重载在反复修改模型源码并调试的过程中Python 的模块导入缓存会记住旧代码导致改动不生效通常只能重启环境。Transformers 为此提供了clear_import_cache工具函数定义于 src/transformers/utils/import_utils.py#L3333它会清空sys.modules中所有以transformers.开头的模块缓存对_LazyModule还会重置其内部对象缓存并强制importlib.reload主模块从而让修改后的代码无需重启进程即可重新导入。from transformers import AutoModel from transformers.utils.import_utils import clear_import_cache model AutoModel.from_pretrained(bert-base-uncased) # 修改模型源码例如 src/transformers/models/... 下的实现 # 清除缓存以重新加载修改后的代码 clear_import_cache() # 重新导入此时生效的是更新后的代码 model AutoModel.from_pretrained(bert-base-uncased)在编写自定义注意力类、反复调试的迭代过程中这一工具能显著提升开发效率。实战拆分 SAM 的 qkv 注意力投影背景SAM 注意力机制的合并投影结构从源码看SamVisionAttention定义于 src/transformers/models/sam/modeling_sam.py#L701。其__init__中关键结构如下self.qkv nn.Linear(config.hidden_size, config.hidden_size * 3, biasconfig.qkv_bias)modeling_sam.py#L717单个线性层一次性输出 3 倍宽度的合并投影self.proj nn.Linear(config.hidden_size, config.hidden_size)modeling_sam.py#L718注意力输出投影头数与缩放num_attention_heads config.num_attention_headshead_dim hidden_size // num_attention_headsscale head_dim ** -0.5modeling_sam.py#L712-L714相对位置编码当config.use_rel_posTrue时初始化rel_pos_h、rel_pos_w两个可学习参数modeling_sam.py#L720-L727。原始forwardmodeling_sam.py#L803-L831将qkv输出reshape为(batch_size, height * width, 3, num_attention_heads, -1)再permute与unbind得到独立的query、key、value随后计算缩放点积注意力并在use_rel_pos开启时叠加分解式相对位置偏置。要降低可训练参数数量与计算开销可以把 LoRA 只挂到q、v上而k保持冻结。但这要求先把合并的qkv拆成三个独立投影——这正是下面自定义类的目标。第一步继承并改造 SamVisionAttention创建 SamVisionAttentionSplit自定义类继承原始SamVisionAttention同时混入nn.Module以确保模块注册在__init__中删除合并的self.qkv替换为三个独立线性层import torch import torch.nn as nn from transformers.models.sam.modeling_sam import SamVisionAttention class SamVisionAttentionSplit(SamVisionAttention, nn.Module): def __init__(self, config, window_size): super().__init__(config, window_size) # 删除合并的 qkv 投影 del self.qkv # 为 q、k、v 分别创建独立投影 self.q nn.Linear(config.hidden_size, config.hidden_size, biasconfig.qkv_bias) self.k nn.Linear(config.hidden_size, config.hidden_size, biasconfig.qkv_bias) self.v nn.Linear(config.hidden_size, config.hidden_size, biasconfig.qkv_bias) self._register_load_state_dict_pre_hook(self.split_q_k_v_load_hook)注意两个细节window_size参数SamVisionAttention.__init__的签名是(config, window_size)传入 0 时表示使用全局注意力输入尺寸由image_size // patch_size推算传入非 0 值时则为窗口注意力modeling_sam.py#L706-L710。自定义类必须保持相同的构造签名因为下方替换时调用方会原样传参。_register_load_state_dict_pre_hook注册一个加载 checkpoint 之前的钩子用于在状态字典进入模块前改写其中的权重键——这是第二步的核心机制。第二步用加载钩子把预训练 qkv 权重一分为三SamVisionAttentionSplit的结构与原版不同直接加载facebook/sam-vit-base等预训练 checkpoint 会因键名不匹配qkv.weightvsq.weight/k.weight/v.weight而失败。钩子函数split_q_k_v_load_hook在权重加载时把形状为(3 * hidden_size, hidden_size)的合并权重沿第 0 维chunk(3, dim0)切成三份分别挂到q.、k.、v.键下并删除原qkv.键def split_q_k_v_load_hook(self, state_dict, prefix, *args): keys_to_delete [] for key in list(state_dict.keys()): if qkv. in key: # 从合并投影中切出 q、k、v q, k, v state_dict[key].chunk(3, dim0) # 替换为独立的 q、k、v 投影权重 state_dict[key.replace(qkv., q.)] q state_dict[key.replace(qkv., k.)] k state_dict[key.replace(qkv., v.)] v # 将旧 qkv 键标记为待删除 keys_to_delete.append(key) # 删除旧的 qkv 键 for key in keys_to_delete: del state_dict[key]由于chunk(3, dim0)按行均匀切片只要原始qkv的拼接顺序是q || k || v源码中 modeling_sam.py#L805-L812 的reshapeunbind(0)即按此约定取出三份拆分后的q、k、v就与原版前向计算完全等价从而保证与任意 SAM 预训练 checkpoint 兼容。第三步重写 forward独立计算 q、k、v拆分后前向传播改为分别调用self.q、self.k、self.v其余注意力计算流程与原版保持一致def forward(self, hidden_states: torch.Tensor, output_attentionsFalse) - torch.Tensor: batch_size, height, width, _ hidden_states.shape qkv_shapes (batch_size * self.num_attention_heads, height * width, -1) query self.q(hidden_states).reshape((batch_size, height * width,self.num_attention_heads, -1)).permute(0,2,1,3).reshape(qkv_shapes) key self.k(hidden_states).reshape((batch_size, height * width,self.num_attention_heads, -1)).permute(0,2,1,3).reshape(qkv_shapes) value self.v(hidden_states).reshape((batch_size, height * width,self.num_attention_heads, -1)).permute(0,2,1,3).reshape(qkv_shapes) attn_weights (query * self.scale) key.transpose(-2, -1) attn_weights torch.nn.functional.softmax(attn_weights, dtypetorch.float32, dim-1).to(query.dtype) attn_probs nn.functional.dropout(attn_weights, pself.dropout, trainingself.training) attn_output (attn_probs value).reshape(batch_size, self.num_attention_heads, height, width, -1) attn_output attn_output.permute(0, 2, 3, 1, 4).reshape(batch_size, height, width, -1) attn_output self.proj(attn_output) if output_attentions: outputs (attn_output, attn_weights) else: outputs (attn_output, None) return outputs几点实现说明形状推导沿用了原版约定每个头的特征维度为hidden_size // num_attention_headsqkv_shapes把张量整理成(batch_size * num_attention_heads, height * width, -1)的多头形式对照 modeling_sam.py#L803-L812softmax 固定使用torch.float32计算后再转回query的 dtype与原版一致modeling_sam.py#L823从代码结构看该简化版forward未包含原版在use_rel_posTrue时叠加分解式相对位置偏置的分支modeling_sam.py#L816-L821。若你的使用场景需要保留相对位置编码带来的精度应在attn_weights计算后补回get_decomposed_rel_pos的偏置逻辑——这点在自行定制时需特别注意。替换注意力模块两种挂接方式与时序陷阱自定义类写好后需要把它真正装进 SAM 模型。这里有两种方式方式 A加载后逐层替换面向已加载模型from transformers import SamModel # 加载预训练 SAM 模型 model SamModel.from_pretrained(facebook/sam-vit-base) # 替换视觉编码器各层中的注意力模块 for layer in model.vision_encoder.layers: if hasattr(layer, attn): layer.attn SamVisionAttentionSplit(model.config.vision_config, model.config.vision_config.window_size)方式 B加载前替换类注册表推荐保证权重正确加载Transformers 通过一个注意力类注册表来实例化各注意力实现。在 modeling_sam.py#L885-L888 中SAM_VISION_ATTENTION_CLASSES { eager: SamVisionAttention, sdpa: SamVisionSdpaAttention, }而视觉编码器层正是通过SAM_VISION_ATTENTION_CLASSESconfig._attn_implementation创建注意力模块modeling_sam.py#L895。因此可以在加载模型之前把eager实现替换为自定义类from transformers import SamModel from transformers.models.sam import modeling_sam # 替换视觉层实例化时使用的注意力类 modeling_sam.SAM_VISION_ATTENTION_CLASSES[eager] SamVisionAttentionSplit # 加载预训练 SAM 模型钩子在加载过程中完成 qkv 权重拆分 model SamModel.from_pretrained(facebook/sam-vit-base, attn_implementationeager)必须注意时序问题方式 A 在模型加载完成后再替换模块此时split_q_k_v_load_hook已经错过 checkpoint 加载阶段新创建的q、k、v层会保持随机初始化模型输出将不再等价于预训练模型方式 B 在from_pretrained之前替换类钩子随 checkpoint 加载同步执行q、k、v才能正确继承预训练权重。若坚持使用方式 A需要自行手动搬运并拆分权重。另外由于自定义实现基于 eager 注意力加载时建议显式传入attn_implementationeager避免默认的sdpa分支影响。应用 LoRA精准挂接 q 与 v拆分完成并获得正确初始化的q、k、v后就可以用 PEFT 库对q、v应用 LoRA 了。LoRA 的核心思想是冻结原始权重仅在旁路训练低秩矩阵从而把可训练参数量压缩到极小比例。首先创建LoraConfig指定秩r、缩放因子lora_alpha、dropout、任务类型以及最关键的目标模块from peft import LoraConfig, get_peft_model config LoraConfig( r16, lora_alpha32, # 只对 q 和 v 应用 LoRA target_modules[q, v], lora_dropout0.1, task_typeFEATURE_EXTRACTION )参数说明r秩低秩矩阵的秩决定 LoRA 适配器的参数量与表达能力。r16表示旁路矩阵维度为16 × hidden_size属于中等规模设置实际使用时可按任务复杂度在 464 之间调节。lora_alpha缩放因子LoRA 旁路输出按alpha / r缩放后与冻结权重相加。alpha32、r16时缩放系数为 2一般与r同量级即可具体数值影响微调强度。target_modulesLoRA 要挂接的模块名列表。这里填入[q, v]是因为前面拆分出的两个独立线性层恰好名为q与v——这也是必须拆分qkv的根本原因。k未被纳入从而进一步节省参数若需要也可加入k或输出投影proj。lora_dropoutLoRA 旁路输入的 dropout 概率用于缓解过拟合训练时生效。task_type任务类型标记。FEATURE_EXTRACTION适用于特征提取/微调场景若在Trainer中做分类等任务可对应调整。随后把模型与配置一并交给get_peft_model完成适配器包装model get_peft_model(model, config)调用print_trainable_parameters可以查看实际训练的参数规模与占比model.print_trainable_parameters() trainable params: 589,824 || all params: 94,274,096 || trainable%: 0.6256以上述facebook/sam-vit-base为例LoRA 只训练约 58.98 万参数而模型总参数约 9427 万可训练比例仅约 0.63%——这正是拆分注意力 选择性挂接 LoRA 的价值所在用极小成本完成对 SAM 视觉编码器的定向适配。总结与注意事项本文演示的组件级定制流程可概括为四步继承原始注意力类并改造结构 → 用加载钩子做权重映射 → 在正确时机替换模块 → 用 PEFT 精准挂接 LoRA。该模式不仅适用于 SAM 的qkv拆分也可以推广到其他模型的自定义改造场景。实践中的关键注意点保持构造签名一致自定义类需沿用(config, window_size)这类原始构造参数才能被模型层的实例化逻辑原样调用modeling_sam.py#L895权重兼容性是硬约束任何结构改动都要配套load_state_dict钩子做键名/形状映射否则预训练权重无法复用替换时机决定权重是否有效务必在from_pretrained之前替换类注册表加载完成后替换会导致新层随机初始化功能保真重写forward时对照原实现逐行核对缩放、softmax 的 float32 计算、相对位置偏置、dropout 等避免定制后行为漂移迭代开发时善用clear_import_cachesrc/transformers/utils/import_utils.py#L3333免去频繁重启环境的开销。通过这条路径你既保留了 Transformers 完整的训练与推理生态Trainer、PreTrainedModel 等又能按需对模型内部结构进行精准外科手术式的定制实现参数效率与任务适配的双赢。【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考