ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Axolotl 集成 Cut Cross Entropy 指南:用低显存交叉熵损失优化大词表模型微调

Axolotl 集成 Cut Cross Entropy 指南:用低显存交叉熵损失优化大词表模型微调 Axolotl 集成 Cut Cross Entropy 指南用低显存交叉熵损失优化大词表模型微调【免费下载链接】axolotlGo ahead and axolotl questions项目地址: https://gitcode.com/GitHub_Trending/ax/axolotl本文是 Axolotl 中 Cut Cross EntropyCCE集成插件的完整技术指南。CCE 是 Apple 提出的交叉熵损失优化实现通过对 loss 计算阶段的优化显著降低训练显存占用尤其适合词表规模巨大的 LLM 微调场景。读完本文你将掌握 CCE 的安装方式、Axolotl 配置方法、可用模型范围、底层插件机制与常见约束并能在自己的训练配置中直接落地。什么是 Cut Cross EntropyCut Cross Entropy 的核心思想是在损失cross-entropy计算阶段做优化从而减少训练时的显存VRAM占用。在标准的因果语言模型训练中模型的lm_head会将隐藏状态映射到词表大小的 logits例如 128K 词表这一巨大张量是显存的主要消耗来源之一。CCE 通过特殊的 kernel 实现避免了在训练过程中物化完整的 logits 张量从而在 forward/backward 阶段都节省显存。相关论文为Cut Your Losses in Large-Vocabulary Language ModelsWijmans 等2024仓库提供了 BibTeX 引用条目见 集成 README。Axolotl 将其封装为插件CutCrossEntropyPlugin并维护了带 transformers 支持的 fork 版本axolotl-ai-cloud/ml-cross-entropy以适配 Axolotl 所支持的众多模型架构。环境要求PyTorch ≥ 2.4.0这是 CCE 的最低版本要求。插件在运行时也会再次校验见下文“插件机制与校验逻辑”。需要fp16/bf16 混合精度训练CCE 的 backward pass 依赖半精度因此在配置中必须开启bf16或fp16否则配置校验会直接报错。安装安装分为两种场景两者都需要安装带transformersextra 的cut-cross-entropy包。开发环境仓库内仓库提供了安装辅助脚本 scripts/cutcrossentropy_install.py它会根据当前环境自动输出正确的安装命令用法为python scripts/cutcrossentropy_install.py | sh该脚本的行为如下若 PyTorch 版本 2.4.0输出空内容并退出不安装检查cut_cross_entropy是否已安装若已安装但缺少cut_cross_entropy.transformers子模块则输出先卸载的命令前缀pip uninstall -y cut-cross-entropy 最终输出的安装命令为pip install cut-cross-entropy[transformers] githttps://github.com/axolotl-ai-cloud/ml-cross-entropy.git4dfa522若在脚本后追加--uv参数python scripts/cutcrossentropy_install.py --uv | sh则使用uv pip install形式适配 uv 环境。pip 直接安装如果环境中已装过旧版或非 transformers 版本的包建议先卸载再安装官方推荐的 fork 版本pip3 uninstall -y cut-cross-entropy pip3 install cut-cross-entropy[transformers] githttps://github.com/axolotl-ai-cloud/ml-cross-entropy.git4dfa522说明Axolotl 插件会校验安装的是否为 Axolotl 的 fork见下文因此请务必使用上述带 commit 固定版本的安装方式而不是官方原版。配置启用在 Axolotl 的 YAML 配置中通过plugins字段注册插件即可plugins: - axolotl.integrations.cut_cross_entropy.CutCrossEntropyPlugin注册插件后还需要满足以下配置约束由 args.py 中的 pydantic 校验保证配置项要求说明cut_cross_entropytrue默认开启插件自带该参数默认即为true用于控制是否应用 CCE 补丁bf16/fp16至少一个为trueCCE 的 backward pass 需要半精度否则校验报错Cut Cross Entropy requires fp16/bf16 training for backward passchunked_cross_entropy必须为false或不设置CCE 与 chunked cross entropy 互斥同时开启会报错Cut Cross Entropy does not support chunked cross entropy此外在更上层的配置校验 validation.py 中cut_cross_entropy、chunked_cross_entropy、liger_cross_entropy、liger_fused_linear_cross_entropy这四种交叉熵优化选项同一时间只能启用一个同时启用多个会抛出校验错误。这意味着如果你同时注册了 Liger 插件并使用其融合交叉熵 kernel需要先关闭其中一个。一个完整的参考配置见 examples/ministral/ministral-small-qlora.yamlbase_model: mistralai/Ministral-8B-Instruct-2410 plugins: - axolotl.integrations.cut_cross_entropy.CutCrossEntropyPlugin load_in_4bit: true adapter: qlora sequence_len: 2048 sample_packing: true micro_batch_size: 2 gradient_accumulation_steps: 4 num_epochs: 1 optimizer: adamw_bnb_8bit lr_scheduler: cosine learning_rate: 0.0002 bf16: auto gradient_checkpointing: true attn_implementation: flash_attention_2支持的模型CCE 插件已在下列模型架构上完成适配模型支持注册表中声明了cut_cross_entropy能力。完整列表见 集成 READMEafmoe、apertus、arcee、cohere、cohere2、cohere2_moe、cohere2_vision、cohere_compass、cohere_compass_text、deepseek_v2、deepseek_v3、deepseek_v4、exaone4、exaone4_5、exaone_moe、gemma、gemma2、gemma3、gemma3_text、gemma3n、gemma3n_text、gemma4、gemma4_text、gemma4_unified、gemma4_unified_text、glm、glm4、glm4_moe、glm4_moe_lite、glm46v、glm4v、glm4v_moe、glm_image、glm_moe_dsa、gpt_oss、granite、granitemoe、granitemoehybrid、granitemoeshared、hunyuan_v1_dense、hunyuan_v1_moe、internvl、kimi_linear、lfm2、lfm2_moe、lfm2_vl、llama、llama4、llama4_text、llava、minimax、minimax_m2、ministral、ministral3、mistral、mistral3、mistral4、mixtral、mllama、muse_glimmer、nemotron_h、olmo、olmo2、olmo3、olmoe、phi、phi3、phi4_multimodal、qwen2、qwen2_5_vl、qwen2_moe、qwen2_vl、qwen3、qwen3_5、qwen3_5_text、qwen3_5_moe、qwen3_5_moe_text、qwen3_moe、qwen3_next、qwen3_vl、qwen3_vl_moe、qwen4_exp、qwen4_exp_text、seed_oss、smollm3、step3p5、step3p7、voxtral。除 README 中的清单外仓库的模型支持注册表还针对部分模型给出了 CCE 适配的特殊说明例如cohere_compass 注册表声明cut_cross_entropy: Supported并注明在 North-Micro-Vision-Instruct 上其 loss 与未打补丁版本在 bf16 噪声水平上一致含logit_scalemuse_glimmer 注册表说明 CCE 直接打补丁到MuseGlimmerForConditionalGenerationfork 会把output_multiplier折叠进 hidden states并把final_logit_softcapping传给apply_lce另有一些模型如 paddleocr_vl 注册表明确将cut_cross_entropy标记为Unsupported训练时应关闭该选项。在配置时如果所选model_config_type未被声明支持插件会在加载模型前调用check_capability并抛出提示Disable cut_cross_entropy for this model.。插件机制与校验逻辑CutCrossEntropyPlugin的实现位于 src/axolotl/integrations/cut_cross_entropy/init.py其关键流程如下参数注入get_input_args()返回axolotl.integrations.cut_cross_entropy.CutCrossEntropyArgs即 args.py 中定义的CutCrossEntropyArgs包含cut_cross_entropy: bool True并带有前文所述的两条 pydantic 校验规则。加载前校验_check_requirements在pre_model_load阶段执行校验 PyTorch 版本 ≥ 2.4.0否则抛出ImportError校验cut_cross_entropy包已安装否则提示安装校验cut_cross_entropy.transformers子模块存在否则提示安装带 transformers extra 的版本校验是否为 Axolotl 的 fork尝试从cut_cross_entropy.transformers.patch导入AXOLOTL_CCE_FORK标志若为假或导入失败则报错提示使用官方推荐的安装命令。模型能力检查pre_model_load若cfg.cut_cross_entropy为真先调用check_capability(get_model_support(cfg.model_config_type), cut_cross_entropy, ...)确认当前模型类型支持 CCE然后执行_check_requirements()再调用cce_patch(model_type, remote_model_id...)应用补丁。若模型设置了trust_remote_code还会把base_model作为remote_model_id传入。通用补丁回退patch_llama_like对于 fork 中尚未登记补丁函数的模型类型插件会注册一个通用补丁动态导入transformers.models.{model_type}.modeling_{model_type}获取对应的{prefix}ForCausalLM类将其forward替换为cut_cross_entropy.transformers.llama.cce_forward。此路径被明确标注为实验性的日志中提示Generic Cut Cross Entropy {model_type} support is experimental and may not work as expected.因此新架构建议先确认官方支持列表。加载流程集成在模型加载器 src/axolotl/loaders/model.py 中插件通过PLUGIN_MANAGER.pre_model_load(self.cfg)在模型加载前被调用同时该文件在判断是否需要将 embedding 层转换为 fp16/bf16 时会显式检查self.cfg.cut_cross_entropy因为 CCE 要求 embedding 层保持半精度以支持 backward pass。端到端测试验证仓库在 tests/e2e/integrations/test_cut_cross_entropy.py 中提供了完整的 e2e 测试覆盖了以下场景Llama CCE基于HuggingFaceTB/SmolLM2-135Mbf16: auto40 步训练断言模型输出存在且 TensorBoard loss 下降Qwen2 CCE基于axolotl-ai-co/tiny-qwen2-129m50 步训练同样断言 loss 下降Llama CCE 不同注意力实现参数化测试flash_attention与sdp_attention两种注意力后端验证 CCE 与不同注意力实现可以共存。测试中还验证了版本行为当 PyTorch 版本 2.4.0 时训练会按预期抛出ImportError。这些测试可以直接作为“最小可用配置”的参考sequence_len: 1024、micro_batch_size: 8、optimizer: adamw_torch_fused、lr_scheduler: cosine、max_steps: 40等方便快速验证 CCE 在自己环境中的行为。常见问题与注意事项安装版本必须使用 Axolotl fork官方原版cut-cross-entropy缺少cut_cross_entropy.transformers模块和AXOLOTL_CCE_FORK标记会被插件的_check_requirements拦截并提示重装。半精度是硬性要求忘记设置bf16/fp16会在配置校验阶段直接失败训练过程中 embedding 层也会被强制转换为半精度。与其他交叉熵优化互斥chunked_cross_entropy以及 Liger 的liger_cross_entropy/liger_fused_linear_cross_entropy都不能与 CCE 同时启用。新模型架构请先确认支持列表虽然插件提供了 llama-like 通用回退补丁但该路径是实验性的官方建议以 README 中的支持列表和模型支持注册表为准。引用如果你在研究中使用了该实现请引用原始论文article{wijmans2024cut, author {Erik Wijmans and Brody Huval and Alexander Hertzberg and Vladlen Koltun and Philipp Kr\ahenb\uhl}, title {Cut Your Losses in Large-Vocabulary Language Models}, journal {arXiv}, year {2024}, url {https://arxiv.org/abs/2411.09009}, }参考资料集成插件 READMEsrc/axolotl/integrations/cut_cross_entropy/README.md插件实现src/axolotl/integrations/cut_cross_entropy/init.py参数与校验src/axolotl/integrations/cut_cross_entropy/args.py安装辅助脚本scripts/cutcrossentropy_install.py端到端测试tests/e2e/integrations/test_cut_cross_entropy.py配置校验src/axolotl/utils/schemas/validation.py模型加载集成src/axolotl/loaders/model.py示例配置examples/ministral/ministral-small-qlora.yaml【免费下载链接】axolotlGo ahead and axolotl questions项目地址: https://gitcode.com/GitHub_Trending/ax/axolotl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表