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 依赖半精度,因此在配置中必须开启
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] @ 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_entropy | true(默认开启) | 插件自带该参数,默认即为true,用于控制是否应用 CCE 补丁 |
bf16/fp16 | 至少一个为true | CCE 的 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_entropy、chunked_cross_entropy、liger_cross_entropy、liger_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,其关键流程如下:
参数注入:
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标志,若为假或导入失败则报错,提示使用官方推荐的安装命令。
- 校验 PyTorch 版本 ≥ 2.4.0,否则抛出
模型能力检查(
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-135M,bf16: auto,40 步训练,断言模型输出存在且 TensorBoard loss 下降; - Qwen2 + CCE:基于
axolotl-ai-co/tiny-qwen2-129m,50 步训练,同样断言 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}, }参考资料
- 集成插件 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),仅供参考