?)
如何用 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 parallelismCP把一条序列切分到 8 张卡上。内存是怎么省下来的四项技术各管一段理解下面这张表后面每一步的配置都有出处技术解决什么在示例中的位置分块 losschunked lossloss 需要物化序列长度 × 词表的 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 推到 160kYaRN 保证这段长度上 loss 平坦不缩放时 160k 末尾 20k token 平均 loss 7.3缩放后 2.8offload 让单卡推到 256kCP 再把百万序列切给 4 张卡每卡 262,144 token就能落地。分块 loss 是默认行为想改回普通 loss 才需要在SFTConfig里显式写loss_typenll。准备环境脚本头部内联声明了依赖trl、trackio以及从 transformers 主分支安装的 transformers只为offload。如果本地 Python 环境缺少这个能力脚本会主动抛错而不是静默运行if offload not in inspect.signature(transformers.PreTrainedModel.gradient_checkpointing_enable).parameters: raise RuntimeError( fThis example needs gradient checkpointing offload, which is not in a released transformers yet. fInstall transformers from main. Got {transformers.__version__}. )所以第一步是确认 transformers 来自主分支例如pip install transformers githttps://github.com/huggingface/transformers.git并装好trl、trackio。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_sizeUlysses 序列并行两者目前不互通。执行训练在仓库根目录运行指南给出的启动命令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_length1_048_576。文档特别指出这是常规继续预训练做法不是packingTrue——packing 承诺文档之间互不可见而 CP 无法保证这一点两者同时开启时 TRL 会直接报错。关键配置摘自 脚本training_args SFTConfig( output_dirQwen3-8B-1M, max_lengthSEQ_LEN, # 1_048_576 per_device_train_batch_size1, # 一步 一条书长序列 logging_steps1, save_strategyno, # 中间 checkpoint 含优化器状态约 123 GB只保留最终模型 pad_to_multiple_of16, # 序列长度需整除 cp_size * 28 卡即 16 gradient_checkpointing_kwargs{offload: True}, bf16True, 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_embeddings32768。指南正文给的是 Qwen3-4B 在 160k 长度下的对照写法factor: 4.0、original_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:201:09:41, 380.10s/it]按文档给出的判读方式逐项核对num_tokens约 1.049e6一步确实是一整条百万 token 序列而不是短序列凑出来的 batchloss起点在 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_CONFexpandable_segments:True。指南说明百万 token 下需要 63.6 GB80 GB 卡但不加这个设置仍会失败——所需张量又大又要求连续分配器在已有块之间找不到空间。pad_to_multiple_of必须配合cp_size要求序列长度为cp_size * 2的倍数。8 卡用 16即脚本的取值CP 组缩到 4 卡就改回 8指南中给的就是这个对应关系。不能开 packingcausal-SDPA 的要求排除了 packing 依赖的 block-diagonal maskTRL 对两者同时开启会报错。换模型受限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.1SP 把切分维度换成 attention heads规模上限是 KV 头数示例模型为 8。进一步阅读可看指南的 Further readingaccelerate 的 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),仅供参考