如何用 TRL 在单台 8-GPU 节点上训练百万 token 长序列(Qwen3-8B)?
2026/9/14 2:11:20 网站建设 项目流程

如何用 TRL 在单台 8-GPU 节点上训练百万 token 长序列(Qwen3-8B)?

【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl

要在单条 1,048,576 token 的序列上做微调(例如 agent 长会话场景下的继续预训练),先面对一个物理限制:一条百万 token 序列放不进单张 GPU,甚至放不进 8 张。TRL 仓库提供了官方指南 Training Beyond 1M Tokens 和配套示例 examples/sft_qwen3_8b_1m_context/sft_qwen3_8b_1m_context.py:用 4 项技术组合,把理论需要 288 GB/GPU 的一步压到约 56 GB/GPU,在单台 8×H100 节点上每步训练一整本书长度的序列。

本文的任务就是跑通这个示例。适用前提(文档明确给出的硬性要求):

  • 一台 8 卡节点,H100 或更强;
  • 依赖 TRL 和transformers 的 main 分支——示例用到的梯度检查点offload参数尚未进入任何 transformers 发布版;
  • 分布式后端为 FSDP2,通过 context parallelism(CP)把一条序列切分到 8 张卡上。

内存是怎么省下来的:四项技术各管一段

理解下面这张表,后面每一步的配置都有出处:

技术解决什么在示例中的位置
分块 loss(chunked loss)loss 需要物化序列长度 × 词表的 logits 矩阵,是长序列下第一个爆内存的点TRL 默认开启,无需任何配置
YaRN 位置重缩放(RoPE)模型没在训练分布内见过超长的位置编号,不缩放时 loss 会从正常的 4 左右恶化到约 10model_init_kwargs里的rope_parameters
激活 offload梯度检查点保留的每层sequence × hidden张量占满显存gradient_checkpointing_kwargs={"offload": True},把检查点张量放 pinned host 内存
Context parallelism单层 MLP 反向时中间张量可达 76 GB,单卡物理装不下accelerate 配置里的parallelism_config_cp_size: 8

指南给出的内存演进曲线:分块 loss 让单卡从 32k token 推到 160k;YaRN 保证这段长度上 loss 平坦(不缩放时 160k 末尾 20k token 平均 loss 7.3,缩放后 2.8);offload 让单卡推到 256k;CP 再把百万序列切给 4 张卡(每卡 262,144 token)就能落地。分块 loss 是默认行为,想改回普通 loss 才需要在SFTConfig里显式写loss_type="nll"

准备环境

脚本头部内联声明了依赖:trltrackio,以及从 transformers 主分支安装的 transformers(只为offload)。如果本地 Python 环境缺少这个能力,脚本会主动抛错而不是静默运行:

if "offload" not in inspect.signature(transformers.PreTrainedModel.gradient_checkpointing_enable).parameters: raise RuntimeError( f"This example needs gradient checkpointing `offload`, which is not in a released transformers yet. " f"Install transformers from main. Got {transformers.__version__}." )

所以第一步是确认 transformers 来自主分支(例如pip install "transformers @ git+https://github.com/huggingface/transformers.git"),并装好trltrackio

accelerate 侧使用仓库自带的 FSDP2 + CP 配置,完整内容如下:

distributed_type: FSDP mixed_precision: bf16 num_processes: 8 # one 8-GPU node fsdp_config: fsdp_version: 2 fsdp_auto_wrap_policy: TRANSFORMER_BASED_WRAP fsdp_cpu_ram_efficient_loading: true parallelism_config: parallelism_config_cp_size: 8 # the whole node forms one context-parallel group

要点:cp_size: 8表示整机 8 卡组成一个 CP 组,共享同一条序列,而不是 8 个各跑各 batch 的独立 worker。注意 CP 与后端绑定——cp_size走 FSDP2;如果换 DeepSpeed 后端,对应的是sp_size(Ulysses 序列并行),两者目前不互通。

执行训练

在仓库根目录运行(指南给出的启动命令):

accelerate launch \ --config_file examples/sft_qwen3_8b_1m_context/context_parallel_8gpu.yaml \ examples/sft_qwen3_8b_1m_context/sft_qwen3_8b_1m_context.py

脚本会做什么:以 streaming 方式读 PG-19 数据集,把书籍顺序拼接成长度约 600 万字符的行,共 4 行,由 trainer 截断到max_length=1_048_576。文档特别指出:这是常规继续预训练做法,不是packing=True——packing 承诺文档之间互不可见,而 CP 无法保证这一点,两者同时开启时 TRL 会直接报错。

关键配置(摘自 脚本):

training_args = SFTConfig( output_dir="Qwen3-8B-1M", max_length=SEQ_LEN, # 1_048_576 per_device_train_batch_size=1, # 一步 = 一条书长序列 logging_steps=1, save_strategy="no", # 中间 checkpoint 含优化器状态约 123 GB,只保留最终模型 pad_to_multiple_of=16, # 序列长度需整除 cp_size * 2,8 卡即 16 gradient_checkpointing_kwargs={"offload": True}, bf16=True, model_init_kwargs={ "dtype": torch.bfloat16, "rope_parameters": { "rope_type": "yarn", "rope_theta": 1_000_000, # 必须与模型自带值一致 "factor": 32.0, "original_max_position_embeddings": 32768, }, }, )

其中两处需要对照自己的模型确认:rope_theta必须与模型出厂配置一致;factor是"目标长度 ÷ 已训练位置范围",脚本对 Qwen3-8B 取 32.0(对应original_max_position_embeddings=32768)。指南正文给的是 Qwen3-4B 在 160k 长度下的对照写法(factor: 4.0original_max_position_embeddings: 40960、同样的rope_theta: 1_000_000),结构一致,数值随模型和目标长度而变。

验证运行是否成功

加载和 tokenize 约需十分钟,随后第一个 step 落地。指南展示的日志示例(文档示例,数值随机器状态会有出入):

{'loss': 4.311, 'grad_norm': 29.25, 'num_tokens': 1049000.0, 'epoch': 0.25} 8%|▊ | 1/12 [06:20<1:09:41, 380.10s/it]

按文档给出的判读方式逐项核对:

  • num_tokens约 1.049e6:一步确实是一整条百万 token 序列,而不是短序列凑出来的 batch;
  • loss起点在 4 左右:正常的起始 loss。脚本注释写明,缺少 RoPE 缩放时 1M 长度的 loss 起点约 10.6——如果日志显示 10 附近,先回头检查rope_parameters是否生效;
  • 步速与显存:脚本文档字符串给出 373 s/step、每卡 56.2 GB(指南示例为 380.10 s/it),在 80 GB 卡上留有余量。

训练结束后,模型写入Qwen3-8B-1M/trainer.save_model),这是脚本唯一的落盘产物。

限制与调整

  • OOM 时先加环境变量:设置PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True。指南说明,百万 token 下需要 63.6 GB(80 GB 卡),但不加这个设置仍会失败——所需张量又大又要求连续,分配器在已有块之间找不到空间。
  • pad_to_multiple_of必须配合cp_size:要求序列长度为cp_size * 2的倍数。8 卡用 16(即脚本的取值);CP 组缩到 4 卡就改回 8,指南中给的就是这个对应关系。
  • 不能开 packing:causal-SDPA 的要求排除了 packing 依赖的 block-diagonal mask,TRL 对两者同时开启会报错。
  • 换模型受限:CP 只表达完整的 causal attention,滑动窗口或线性 attention 层的模型会被拒绝——GPT-OSS、Gemma 3/4、Mistral、Qwen3.5 及更新版本都不行;Qwen3 和 Qwen3-MoE 全序列都是 full attention,所以示例选 Qwen3-8B。脚本注释给出同节点参考:Qwen3-0.6B 约 137 s/step。
  • DeepSpeed 分支:把配置里的parallelism_config_cp_size换成parallelism_config_sp_size并改用 DeepSpeed 后端即可走 SP 路径,要求 accelerate 1.12、DeepSpeed 0.18.1;SP 把切分维度换成 attention heads,规模上限是 KV 头数(示例模型为 8)。

进一步阅读可看指南的 Further reading:accelerate 的 context parallelism 概念文档(cp_size与其构建的 device mesh)以及 Ulysses/ring attention 两种交换方式的通信代价对比。

【免费下载链接】trlTrain transformer language models with reinforcement learning.项目地址: https://gitcode.com/GitHub_Trending/tr/trl

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询