
SHiRA 稀疏高秩适配器实战指南基于 PEFT 直接微调少量基座参数【免费下载链接】peft PEFT: State-of-the-art Parameter-Efficient Fine-Tuning.项目地址: https://gitcode.com/gh_mirrors/pe/peftSHiRASparse High Rank Adapters稀疏高秩适配器是 PEFT 中与 LoRA 并列的一种适配器范式它绕过低秩分解思路通过稀疏掩码直接微调基座模型的一小部分原始权重约 1%–2%在视觉与语言任务上取得优于 LoRA 的精度并显著缓解多适配器融合时的概念丢失concept loss问题。本指南将带你从零掌握 SHiRA 的配置参数、训练脚本、自定义稀疏掩码函数与模型加载并深入其底层源码实现最终能够在自己的任务上直接复现和二次开发。SHiRA 是什么从低秩适配到稀疏高秩与 LoRA 在权重矩阵旁引入低秩分解旁路不同SHiRA 直接在预训练权重矩阵上选中少量元素并将其设为可训练参数其余权重保持不变。由于被选中的参数分散在整张 m×n 权重矩阵中其有效秩远高于 LoRA 的固定秩 r因此被称为高秩适配器又因只有极少数参数参与训练适配器本身是高度稀疏的故称稀疏高秩。根据 官方文档 与 README 的说明SHiRA 相对 LoRA 的主要优势包括精度更高在一系列视觉与语言任务上获得优于 LoRA 的准确率更优的多适配器融合显著降低低秩适配器在多适配器同时使用时常见的概念丢失问题零推理开销与快速切换稀疏的 delta 权重可直接合并进基座权重fused 模式既能免去推理开销又能在融合模式下直接切换不同适配器。当前 SHiRA 有一个明确的约束仅支持nn.Linear层见 package_reference/shira.md 与 model.py 中对目标模块类型的检查。核心原理稀疏掩码与可训练参数设计SHiRA 的参数规模设计与 LoRA 对齐便于公平对比。在 config.py 中r 的语义是对于形状为 m×n 的目标张量SHiRA 参数数量按r(mn)计算这与一个 LoRA 适配器的参数量相同但 SHiRA 是高秩适配器设置 r 并不限制其实际秩。从 layer.py 可以看到底层的具体实现训练开始时通过掩码函数生成一张与base_layer.weight同形状的 0/1 稀疏掩码掩码中值为 1 的位置通过torch.where(mask 1.0)提取为shira_indices真正可训练的shira_weight是一个长度恰为r(mn)的一维nn.Parameter默认零初始化其位置与掩码选中的索引一一对应前向时通过torch.sparse_coo_tensor将shira_weight组装成稀疏 delta 权重并叠加到基座权重上见get_delta_weight。这种一维参数 稀疏索引的设计避开了将torch.sparse_coo_tensor直接作为 Parameter 的兼容性问题同时保证了与 LoRA 相当的训练速度和更低的峰值显存占用。ShiraConfig 配置参数详解ShiraConfig定义于 src/peft/tuners/shira/config.py继承自PeftConfig。核心参数如下参数默认值说明r32目标模块的 SHiRA 参数量按r(mn)计算与 LoRA 参数量一致SHiRA 为高秩适配器此值不限制实际秩mask_typerandom掩码函数类型默认使用随机稀疏掩码也可实例化config后手动赋值config.mask_fn传入自定义掩码函数random_seedNone随机掩码使用的 torch generator 的随机种子target_modulesNone要替换为 SHiRA 的模块名列表或正则表达式例如[q, v]或.*decoder.*(SelfAttention\|EncDecAttention).*(q\|v)$仅支持线性层fan_in_fan_outFalse目标层若按 (fan_in, fan_out) 存储权重如 GPT-2 的Conv1D需设为Trueinit_weightsTrue为True时 SHiRA 权重零初始化为False时使用 randn 初始化仅用于测试modules_to_saveNone除 SHiRA 层外需要可训练并保存的模块如分类任务中随机初始化的classifier/score层__post_init__中会根据mask_type自动绑定掩码函数当mask_type random时绑定random_mask否则给出告警并置mask_fn None等待用户手动赋值config.py。快速上手SFTTrainer 一键微调README 给出了最精简的入门示例使用 TRL 的SFTTrainer对facebook/opt-350m做指令微调import torch from peft import ShiraConfig, get_peft_model from transformers import AutoTokenizer, AutoModelForCausalLM from trl import SFTConfig, SFTTrainer from datasets import load_dataset model AutoModelForCausalLM.from_pretrained(facebook/opt-350m, dtypetorch.bfloat16, device_mapauto) tokenizer AutoTokenizer.from_pretrained(facebook/opt-350m) dataset load_dataset(imdb, splittrain[:1%]) shira_config ShiraConfig( r32, ) peft_model get_peft_model(model, shira_config) training_args SFTConfig(dataset_text_fieldtext, max_length128) trainer SFTTrainer( modelpeft_model, train_datasetdataset, processing_classtokenizer, ) trainer.train() peft_model.save_pretrained(shira-opt-350m)流程要点get_peft_model(model, shira_config)会遍历模型、按target_modules此处为默认映射将所有匹配的线性层替换为ShiraLayer并仅将稀疏选中的参数置为可训练训练结束后save_pretrained保存的是稀疏索引shira_indices与一维权重shira_weight文件体积与同参数量的 LoRA 相当SFTConfig(dataset_text_fieldtext, max_length128)是 TRL 侧配置与 PEFT 无关但需配套设置。使用示例脚本训练含全部命令行参数仓库提供了功能更完整的训练脚本 examples/shira_finetuning/shira_finetuning.py基于transformers.Trainer与 Alpaca 指令数据支持 DDP、自定义掩码函数等。直接运行python3 examples/shira_finetuning/shira_finetuning.py --base_model facebook/opt-350m脚本支持的全部命令行参数均可在 shira_finetuning.py 中找到参数默认值说明--base_modelpath/to/model基座模型名称或本地路径--data_pathyahma/alpaca-cleaned数据集名称load_dataset可解析的格式--output_dirshira输出目录--batch_size16每设备训练 batch 大小--num_epochs1训练轮数--learning_rate3e-4学习率--cutoff_len256序列截断长度--val_set_size16验证集大小从训练集切分--eval_step100评估间隔步数--save_step100保存间隔步数--device_mapauto设备映射策略--shira_r32SHiRA 的 r 值--shira_target_modulesNone目标模块列表逗号分隔或正则默认使用内置映射--dtypefloat16模型加载精度torch属性名--seedNone随机种子--use_custom_random_mask_function_with_custom_kwargsFalseflag是否启用示例中的自定义掩码函数使用 accelerate 运行 DDP脚本会自动检测环境变量WORLD_SIZE/PMI_SIZE判断是否处于多卡环境当世界大小大于 1 时会将device_map自动设为{: Accelerator().process_index}以配合 DDP见 shira_finetuning.py。运行方式accelerate config accelerate launch examples/shira_finetuning/shira_finetuning.py --base_model facebook/opt-350m若要在 CPU 上微调请追加--device_map cpupython3 examples/shira_finetuning/shira_finetuning.py --base_model facebook/opt-350m --device_map cpu自定义稀疏掩码函数掩码函数的契约PEFT 对掩码函数有明确签名要求见 mask_functions.py必需位置参数base_layer挂载 SHiRA 适配器的线性层、r用于按 LoRA 一致的方式确定适配器参数量关键字参数可按需追加返回值与base_layer.weight同形状的torch.tensor取值只能为 0 或 1且dtype 与 device 必须与base_layer.weight一致掩码中 1 的数量必须严格等于r(mn)否则 layer.py 中的维度校验会直接报错。默认的random_mask实现思路mask_functions.py用torch.randperm在全部m*n个权重位置中随机抽取r(mn)个索引通过scatter_置 1 生成掩码并支持通过random_seed控制生成器的随机性可通过ShiraConfig.random_seed配置。示例带自定义参数的掩码函数当自定义掩码需要额外关键字参数时脚本提供了custom_random_mask_function_with_custom_kwargs作为模板shira_finetuning.py它通过闭包把自定义参数custom_arg示例中设为 120捕获进mask_fn用该值派生随机种子其余逻辑与random_mask一致def custom_random_mask_function_with_custom_kwargs(custom_arg): def mask_fn(base_layer, r): new_seed custom_arg shape base_layer.weight.shape num_shira_weights r * (shape[0] shape[1]) random_generator torch.Generator() random_generator.manual_seed(new_seed) idx (torch.randperm(base_layer.weight.numel(), generatorrandom_generator)[:num_shira_weights]).to( base_layer.weight.device ) val torch.ones_like(idx.type(base_layer.weight.dtype)) mask torch.zeros_like(base_layer.weight.view(1, -1)) mask mask.scatter_(1, idx.unsqueeze(0), val.unsqueeze(0)).view(shape) return mask return mask_fn脚本内部会据此构建配置shira_finetuning.pyconfig ShiraConfig(rshira_r, mask_typecustom, target_modulesshira_target_modules, task_typeCAUSAL_LM) custom_mask_fn custom_random_mask_function_with_custom_kwargs(120) config.mask_fn custom_mask_fn运行命令python3 examples/shira_finetuning/shira_finetuning.py --base_model facebook/opt-350m --use_custom_random_mask_function_with_custom_kwargs不加该参数时SHiRA 默认使用mask_typerandom的随机稀疏掩码。这一自定义流程同样有测试覆盖tests/test_shira.py中的test_save_load_custom_mask_function使用mask_typecustom与自定义掩码函数验证了训练、保存与加载全链路test_shira.py。加载与使用微调后的模型微调产物与其他 PEFT 适配器完全一致可用PeftModel.from_pretrained直接加载from peft import PeftModel from transformers import AutoTokenizer, AutoModelForCausalLM model AutoModelForCausalLM.from_pretrained(facebook/opt-350m) tokenizer AutoTokenizer.from_pretrained(facebook/opt-350m) shira_model PeftModel.from_pretrained(model, shira-opt-350m)加载后即可按常规方式推理或继续训练。注意 SHiRA 的 delta 权重是稀疏张量因此supports_lora_conversion返回False稀疏权重不适用于 SVD 分解见 layer.py即不支持转换为 LoRA 权重。源码级深入SHiRA 在 PEFT 中的实现前向与合并ShiraLayer.Linear.forward的默认路径会克隆基座权重并叠加各激活适配器的稀疏 delta 权重再执行F.linearlayer.py。因此挂在基座层上的 forward/backward hook 会被 SHiRA 的 forward 实现忽略如需 hook 请挂在 PEFT 模块上_warn_once_about_module_hooks会给出一次性告警。merge/unmerge方法支持把适配器 delta 合并进基座权重safe_mergeTrue时先复制再校验 NaN合并后推理零额外开销set_scale可在推理时调整各适配器的缩放系数训练时默认 scaling 为 1.0。量化支持从 model.py 看SHiRA 通过resolve_quantization_backend识别量化层并复用Linear实现当量化后端不支持权重合并时forward 走基座结果 稀疏 delta 结果的分支避免反量化整张权重layer.py。跨平台保存兼容shira_indices是整数张量而 Windows 平台在 safetensors 中保存整数存在问题。为此ShiraModel._get_adapter_state_dict在 Windows 上先将索引转成float32保存加载时再转回intmodel.py因此 Windows 与其他平台的 checkpoint 可以互通。测试与验证仓库的 tests/test_shira.py 覆盖了 SHiRA 的关键行为可作为二次开发的参考基线test_mlp_single_adapter_shapes单适配器下shira_weight形状与索引维度的正确性test_multiple_adapters_save_load多适配器保存/加载test_save_load_custom_mask_function自定义掩码函数的全链路验证含init_weightsFalsetest_save_load_default_random_mask_with_seed_function带随机种子默认掩码的保存加载test_shira_dtypes不同 dtype 下的兼容性test_shira_warns_about_hooks基座层 hook 被忽略时的一次性告警。引用若在你的研究或产品中使用 SHiRA请引用论文inproceedings{NEURIPS2024_18c0102c, author {Bhardwaj, Kartikeya and Pandey, Nilesh Prasad and Priyadarshi, Sweta and Ganapathy, Viswanath and Kadambi, Shreya and Esteves, Rafael and Borse, Shubhankar and Whatmough, Paul and Garrepalli, Risheek and Van Baalen, Mart and Teague, Harris and Nagel, Markus}, booktitle {Advances in Neural Information Processing Systems}, editor {A. Globerson and L. Mackey and D. Belgrave and A. Fan and U. Paquet and J. Tomczak and C. Zhang}, pages {13685--13715}, publisher {Curran Associates, Inc.}, title {Sparse High Rank Adapters}, volume {37}, year {2024} }小结SHiRA 用稀疏选择 直接微调原始权重重新定义了参数高效微调的范式参数量与 LoRA 对齐、精度更高、支持快速切换与高质量多适配器融合。通过本指南你可以基于ShiraConfig快速接入 PEFT 生态借助 shira_finetuning.py 一键训练并通过自定义mask_fn探索任意稀疏结构若需深入实现细节可继续阅读 config.py、mask_functions.py、layer.py 与 model.py 四个核心文件。【免费下载链接】peft PEFT: State-of-the-art Parameter-Efficient Fine-Tuning.项目地址: https://gitcode.com/gh_mirrors/pe/peft创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考