Axolotl 集成 Cut Cross Entropy 指南:用低显存交叉熵损失优化大词表模型微调
2026/9/15 20:44:01 网站建设 项目流程

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

【免费下载链接】axolotlGo ahead and axolotl questions项目地址: https://gitcode.com/GitHub_Trending/ax/axolotl

本文是 Axolotl 中 Cut Cross Entropy(CCE)集成插件的完整技术指南。CCE 是 Apple 提出的交叉熵损失优化实现,通过对 loss 计算阶段的优化显著降低训练显存占用,尤其适合词表规模巨大的 LLM 微调场景。读完本文,你将掌握 CCE 的安装方式、Axolotl 配置方法、可用模型范围、底层插件机制与常见约束,并能在自己的训练配置中直接落地。

什么是 Cut Cross Entropy

Cut Cross Entropy 的核心思想是在损失(cross-entropy)计算阶段做优化,从而减少训练时的显存(VRAM)占用。在标准的因果语言模型训练中,模型的lm_head会将隐藏状态映射到词表大小的 logits(例如 128K 词表),这一巨大张量是显存的主要消耗来源之一。CCE 通过特殊的 kernel 实现,避免了在训练过程中物化完整的 logits 张量,从而在 forward/backward 阶段都节省显存。相关论文为Cut Your Losses in Large-Vocabulary Language Models(Wijmans 等,2024),仓库提供了 BibTeX 引用条目,见 集成 README。

Axolotl 将其封装为插件CutCrossEntropyPlugin,并维护了带 transformers 支持的 fork 版本(axolotl-ai-cloud/ml-cross-entropy),以适配 Axolotl 所支持的众多模型架构。

环境要求

  • PyTorch ≥ 2.4.0:这是 CCE 的最低版本要求。插件在运行时也会再次校验,见下文“插件机制与校验逻辑”。
  • 需要fp16/bf16 混合精度训练:CCE 的 backward pass 依赖半精度,因此在配置中必须开启bf16fp16,否则配置校验会直接报错。

安装

安装分为两种场景,两者都需要安装带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] @ git+https://github.com/axolotl-ai-cloud/ml-cross-entropy.git@4dfa522"
  • 若在脚本后追加--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] @ git+https://github.com/axolotl-ai-cloud/ml-cross-entropy.git@4dfa522"

说明: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 pass"
chunked_cross_entropy必须为false(或不设置)CCE 与 chunked cross entropy 互斥,同时开启会报错:"Cut Cross Entropy does not support chunked cross entropy"

此外,在更上层的配置校验 validation.py 中,cut_cross_entropychunked_cross_entropyliger_cross_entropyliger_fused_linear_cross_entropy这四种交叉熵优化选项同一时间只能启用一个,同时启用多个会抛出校验错误。这意味着如果你同时注册了 Liger 插件并使用其融合交叉熵 kernel,需要先关闭其中一个。

一个完整的参考配置见 examples/ministral/ministral-small-qlora.yaml:

base_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能力)。完整列表见 集成 README:

afmoe、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_scale);
  • muse_glimmer 注册表:说明 CCE 直接打补丁到MuseGlimmerForConditionalGeneration,fork 会把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,其关键流程如下:

  1. 参数注入get_input_args()返回axolotl.integrations.cut_cross_entropy.CutCrossEntropyArgs,即 args.py 中定义的CutCrossEntropyArgs(包含cut_cross_entropy: bool = True,并带有前文所述的两条 pydantic 校验规则)。

  2. 加载前校验(_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标志,若为假或导入失败则报错,提示使用官方推荐的安装命令。
  3. 模型能力检查(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传入。

  4. 通用补丁回退(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."),因此新架构建议先确认官方支持列表。

  5. 加载流程集成:在模型加载器 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: auto,40 步训练,断言模型输出存在且 TensorBoard loss 下降;
  • Qwen2 + CCE:基于axolotl-ai-co/tiny-qwen2-129m,50 步训练,同样断言 loss 下降;
  • Llama + CCE + 不同注意力实现:参数化测试flash_attentionsdp_attention两种注意力后端,验证 CCE 与不同注意力实现可以共存。

测试中还验证了版本行为:当 PyTorch 版本 < 2.4.0 时,训练会按预期抛出ImportError。这些测试可以直接作为“最小可用配置”的参考(sequence_len: 1024micro_batch_size: 8optimizer: adamw_torch_fusedlr_scheduler: cosinemax_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}, }

参考资料

  • 集成插件 README:src/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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询