ARTICLE DETAIL

资讯详情

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

FiD分布式训练指南:SLURM集群下多机多卡训练Reader与Retriever的完整配置

FiD分布式训练指南:SLURM集群下多机多卡训练Reader与Retriever的完整配置 FiD分布式训练指南SLURM集群下多机多卡训练Reader与Retriever的完整配置【免费下载链接】FiDFusion-in-Decoder项目地址: https://gitcode.com/gh_mirrors/fi/FiDFiDFusion-in-Decoder是一个面向开放域问答的模型框架包含多段落生成式Reader基于T5与通过知识蒸馏训练的Retriever基于BERT。本指南带你完成FiD 分布式训练的核心配置在 SLURM 集群上启动多机多卡训练覆盖环境初始化原理、作业提交脚本、Reader 与 Retriever 两套超参数设置以及断点续训等实战要点。一、FiD 分布式训练是如何工作的FiD 的训练脚本并没有把集群逻辑写死在业务代码里而是统一收敛到一个核心模块中这正是新手最容易忽略的地方src/slurm.py负责分布式初始化。init_distributed_mode()会自动检测环境变量SLURM_JOB_ID是否存在来判断是否运行在 SLURM 集群上并把 SLURM 提供的变量映射成 PyTorch 所需的通信参数。train_reader.pyReader 训练入口脚本启动时先调用分布式初始化再用DistributedDataParallel包装模型。train_retriever.pyRetriever 训练入口训练时自动使用DistributedSampler切分数据。SLURM 变量与 PyTorch 排名的映射关系见 src/slurm.py 中的init_distributed_modeSLURM 环境变量含义PyTorch 用途SLURM_JOB_NUM_NODES节点数n_nodesSLURM_NODEID当前节点编号node_idSLURM_LOCALID节点内进程号local_rank绑定哪块 GPUSLURM_PROCID全局进程号global_rank/RANKSLURM_NTASKS总进程数world_size主节点地址则通过scontrol show hostnames自动查询代码会将其与--main_port端口写入MASTER_ADDR、MASTER_PORT最终以 NCCL 后端初始化进程组。⚡ 还有一个很实用的机制init_signal_handler()会监听 SLURM 的抢占信号SIGUSR1当作业被抢占或超时主进程会自动执行scontrol requeue重新排队无需人工干预。二、环境准备一键安装依赖与下载数据1. 安装依赖FiD 依赖 Python 3、PyTorch 和transformers 3.0.2对版本敏感官方提示其他版本大概率无法运行。完整依赖清单见 requirements.txtgit clone https://gitcode.com/gh_mirrors/fi/FiD.git cd FiD pip install -r requirements.txt2. 下载数据与预训练模型数据运行bash get-data.sh可自动下载 NaturalQuestions、TriviaQA 问答数据及维基百科段落Wikipedia passages并调用src/preprocess.py完成预处理统一输出到open_domain_data/目录。预训练模型bash get-model.sh -m nq_reader_base可下载 6 个官方预训练模型如nq_reader_large、tqa_retriever等适合快速验证或作为继续训练的起点。 多机训练时务必将数据目录放在所有节点共享的存储如 NFS/Lustre上避免每个节点各自复制。三、单机多卡分布式训练前的第一步在提交多机作业前建议先跑通单机多卡。FiD 同时支持两种启动方式均由 src/slurm.py 自动识别单卡/单进程直接python train_reader.py ...代码自动回退到非分布式模式torch 启动器使用torch.distributed.launch或torchrun启动代码会读取RANK、WORLD_SIZE、NGPU环境变量完成初始化。一个最小可用的单机 Reader 训练命令python train_reader.py \ --train_data open_domain_data/NQ_open_train.json \ --eval_data open_domain_data/NQ_open_dev.json \ --model_size base \ --per_gpu_batch_size 1 \ --n_context 100 \ --name my_experiment \ --checkpoint_dir checkpoint \ --use_checkpoint⚠️ 训练 100 个段落--n_context 100非常吃显存--use_checkpoint用激活值重计算换取显存建议始终开启。四、SLURM 多机多卡训练的完整配置1. 三个关键的分布式参数提交集群作业前必须理解 src/options.py 中定义的两个分布式参数参数说明--local_rank集群上必须保持默认值 -1由 SLURM 通过SLURM_LOCALID接管代码中有断言会强制检查--main_port多机作业时主节点通信端口必须位于 10001–20000 区间且该端口在主节点上未被占用其余通用参数--name实验名决定 checkpoint 子目录、--checkpoint_dir输出根目录、--per_gpu_batch_size每卡批大小。2. sbatch 作业提交脚本以下是一个 4 节点 × 8 卡的 Reader 训练提交脚本示例#!/bin/bash #SBATCH --job-namefid_reader #SBATCH --nodes4 #SBATCH --ntasks-per-node8 #SBATCH --gresgpu:8 #SBATCH --cpus-per-task10 #SBATCH --mem150G #SBATCH --time12:00:00 #SBATCH --outputlogs/%x_%j.out #SBATCH --errorlogs/%x_%j.err srun python train_reader.py \ --main_port 12345 \ --train_data open_domain_data/NQ_open_train.json \ --eval_data open_domain_data/NQ_open_dev.json \ --model_size large \ --use_checkpoint \ --answer_maxlength 50 \ --lr 0.00005 \ --optim adamw \ --scheduler linear \ --weight_decay 0.01 \ --text_maxlength 250 \ --per_gpu_batch_size 1 \ --n_context 100 \ --total_steps 15000 \ --warmup_steps 1000 \ --accumulation_steps 4 \ --name fid_reader_large \ --checkpoint_dir /shared/checkpoint要点拆解--ntasks-per-node8srun每个 GPU 一个进程SLURM 会为每个任务设置SLURM_LOCALID/SLURM_PROCID代码据此绑定 GPU--answer_maxlength 50固定 decoder 侧张量长度消除变长张量带来的显存碎片train_reader.py 中的官方建议--accumulation_steps 4梯度累积等效把批大小扩大 4 倍而不增加显存上述lr 0.00005 / adamw / linear / 15000 steps正是官方large Reader 在 64 GPU 上使用的超参数4 节点 32 卡场景可参考沿用。3. Retriever 的多机训练配置Retriever 训练入口是 train_retriever.py分布式机制完全相同同样调用src/slurm初始化差异主要在模型与数据侧srun python train_retriever.py \ --main_port 12346 \ --train_data open_domain_data/NQ_open_train.json \ --eval_data open_domain_data/NQ_open_dev.json \ --lr 1e-4 \ --optim adamw \ --scheduler linear \ --n_context 100 \ --total_steps 20000 \ --scheduler_steps 30000 \ --per_gpu_batch_size 8 \ --name nq_retriever \ --checkpoint_dir /shared/checkpoint与 Reader 的两个实现差异新手容易困惑Retriever 训练使用DistributedSampler且每个 epoch 自动调用set_epoch()保证数据切分随机性Reader 则按global_rank / world_size手动切分数据文件Retriever 的 DDP 包装设置了find_unused_parametersTrueReader 为False因为部分参数如蒸馏分支可能不参与每次前向属正常配置不必修改。 完整工作流先用test_reader.py加--write_crossattention_scores生成 Reader 交叉注意力分数多卡下各 rank 结果会自动汇总为dataset_wscores.json见 src/util.py 的save_distributed_dataset再喂给train_retriever.py训练最后经generate_passage_embeddings.py索引知识库、passage_retrieval.py检索——整个闭环可在 SLURM 上逐环节提交作业。五、断点续训与实验监控FiD 的 checkpoint 机制对多机训练特别友好全部由主进程rank 0统一写入每隔--save_freq步保存checkpoint/step-{step}目录并维护一个指向最新检查点的latest软链接src/util.py 的save()验证指标刷新时额外保存checkpoint/best_dev续训方法不传--model_path默认none且实验目录已存在时脚本自动从checkpoint/latest恢复模型、优化器、学习率调度器与步数——被抢占后 requeue 重跑时直接复用同一--name即可无缝衔接监控主进程写 TensorBoard 事件到checkpoint_dir/name/日志输出到run.log非主进程只输出 WARN 级别避免 32 个节点刷屏。六、常见问题排查清单现象原因与解决断言报错local_rank -1集群作业中误传了--local_rank删掉该参数即可启动卡死/断言端口错误--main_port不在 10001–20000 区间或端口在主节点被占用OOM 显存溢出依次尝试开启--use_checkpoint、设置--answer_maxlength、调小--per_gpu_batch_size、减小--n_context作业被抢占后丢失进度确认--checkpoint_dir在共享存储进度会自动 requeue重跑同一--name续训依赖冲突transformers 锁定 3.0.2PyTorch 建议 1.6新建独立 conda 环境隔离总结FiD 的分布式训练把 SLURM 细节封装在 src/slurm.py 中使用者只需记住三件事集群上--local_rank保持 -1、--main_port选 10001–20000 区间、数据与 checkpoint 放共享存储。Reader 用train_reader.py起步、Retriever 用train_retriever.py进阶配合官方 64 GPU 超参数和梯度累积即可在任意规模的 SLURM 集群上稳定完成多机多卡训练。【免费下载链接】FiDFusion-in-Decoder项目地址: https://gitcode.com/gh_mirrors/fi/FiD创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表