资讯动态

flame 训练指南:flash-linear-attention 中线性注意力语言模型的数据预处理、从零训练与持续预训练实战

发布时间:2026/9/17 10:03:46 来源:尧图企业网站定制
flame 训练指南flash-linear-attention 中线性注意力语言模型的数据预处理、从零训练与持续预训练实战【免费下载链接】flash-linear-attention Efficient implementations for emerging model architectures项目地址: https://gitcode.com/GitHub_Trending/fl/flash-linear-attention本文基于 flash-linear-attention 仓库中legacy/training目录下的 Flame 训练框架文档README完整讲解如何用极少代码训练 GLAGated Linear Attention语言模型涵盖环境安装、数据集预分词、train.sh关键超参配置与global_batch_size的计算方法、断点续训以及从 Mistral-7B 权重迁移到 GLA-7B 进行持续预训练continual pretraining的完整流程。读完后你将能够独立复现仓库提供的 340M/1B/7B 规模的 GLA 训练实验并理解 run.py、train.sh 与 flame 包背后的实际实现逻辑。需要首先说明的是仓库 README 顶部已明确标注flame项目已迁移到基于 torchtitan 构建的新项目本目录代码作为legacy 存档不再同步后续更新。本文描述的安装、命令与参数均以当前仓库中legacy/training的实际代码为准适用于复现历史实验和理解训练框架设计若追求最新能力应参考新项目其地址见 README 顶部说明。整体架构datasets transformers accelerate 三件套Flame 的设计目标是“几行代码就能训练大语言模型”它完全构建在 Hugging Face 生态之上datasets负责数据加载与处理训练前通过预分词缓存save_to_disk避免在线 tokenize 的开销transformers提供AutoConfig/AutoModelForCausalLM/AutoTokenizer/Trainer等模型定义与训练循环基础设施accelerate负责分布式训练启动官方说明中默认使用 DeepSpeedREADME 脚注说明也支持 megatron 等框架。从源码结构看训练入口是 run.py其核心流程为get_train_args()解析参数flame/parser.py 中通过HfArgumentParser解析一个扩展的TrainingArgumentsdataclass根据from_config参数决定是随机初始化AutoModelForCausalLM.from_config用于从零训练还是加载预训练权重AutoModelForCausalLM.from_pretrained用于持续预训练——这正是 parser.py 中from_config字段默认True的两种取值含义load_from_disk(args.cache_dir)直接加载预分词好的数据集并按seed打乱用 DataCollatorForLanguageModeling 组装 batch支持varlen变长打包模式调度器自动配置cosine_with_min_lr会附加min_lr_rate0.1warmup_stable_decay即 WSDwarmup-stable-decay会自动设置num_stable_steps 0.9 × max_steps - warmup_steps、num_decay_steps 0.1 × max_steps见 run.pyTrainer.train(resume_from_checkpoint...)启动训练结束后保存模型、tokenizer、指标与 state。环境准备SetupFlame 与fla的依赖都很少。按 README 的说明克隆仓库并安装git clone https://github.com/sustcsonglin/flash-linear-attention.git pip install . pip install accelerate注意accelerate是分布式训练的必要依赖train.sh最终通过accelerate launch启动而datasets、transformers会随pip install .的依赖链引入。README 中特别提醒CAUTIONHuggingFace 的tokenizers在处理超长文档时存在内存泄漏问题务必安装tokenizers0.20.4。这一限制在实际预处理阶段影响很大因为preprocess.py会对整个文档语料做批量 tokenize。数据预处理preprocess.py训练前需要先下载并**预分词pre-tokenize**数据集。仓库提供了 preprocess.py 脚本。以 tokenizefineweb-edu的 10B 样本为例python preprocess.py \ --dataset HuggingFaceFW/fineweb-edu \ --name sample-10BT \ --split train \ --context_length 2048处理结果会缓存到data/HuggingFaceFW/fineweb-edu/sample-10BT/train供训练时load_from_disk直接读取。参数说明以源码为准README 示例中使用的--context_length对应 preprocess.py 命令行解析器中的--seq_len参数默认 2048即每个训练样本的总序列长度。完整的命令行参数及默认值如下摘自源码argparse定义参数默认值说明--datasetHuggingFaceFW/fineweb-edu数据集名称或本地路径--nameNone数据集配置名如sample-10BT--splittrain处理的 split--seed42打乱数据集的随机种子--outputdata输出根目录--tokenizerfla-hub/gla-1.3B-100B分词器与 GLA 训练一致的 32k 词表--num_proc64并行处理进程数--batch_size2048tokenize 时的批处理大小--seq_len2048每个训练样本的总序列长度README 中写作--context_length--ctx_lenNone保留的最大连续长度超过则切分文档再拼接--return_offsetsFalse是否输出拼接偏移量配合 varlen 训练分词与切块的实际逻辑preprocess.py 中的tokenize函数揭示了缓存数据的组织方式对每批examples[text]调用 tokenizer 得到input_ids若指定了ctx_len先把每条序列按ctx_len切成连续片段避免在文档中间硬切断将所有片段扁平拼接itertools.chain总长度取整到seq_len的倍数再按seq_len逐段切出训练样本——即典型的“打包式packing”语料处理--ctx_len与--seq_len存在约束源码中明确校验ctx_len不得超过seq_len否则抛出ValueError。输出路径规则同样值得注意指定--name时为{output}/{dataset}/{name}/{split}否则为{output}/{dataset}/{split}见 preprocess.py。这与 README 中data/HuggingFaceFW/fineweb-edu/sample-10BT/train的缓存路径完全一致。SlimPajama 的处理方式GLA 论文中预训练使用的是 SlimPajama 的子集。由于数据集体量大README 建议使用git lfs快速下载git lfs install git clone https://huggingface.co/datasets/cerebras/SlimPajama-627B --depth 1 python preprocess.py \ --dataset SlimPajama-627B \ --split train \ --context_length 2048注意此处不传--name因此缓存路径为data/SlimPajama-627B/train——这正是后文 7B 持续预训练命令中cachedata/SlimPajama-627B/train的来源。从零训练train.sh 与关键超参训练 340M 模型的完整命令README 原文示例bash train.sh \ typegla \ lr3e-4 \ schedulercosine_with_min_lr \ batch32 \ update1 \ warmup1024 \ steps20480 \ context2048 \ gpus8 \ nodes1 \ pathexp/gla-340M-10B \ projectfla \ modelconfigs/gla_340M.json \ dataHuggingFaceFW/fineweb-edu \ namesample-10BT \ cachedata/HuggingFaceFW/fineweb-edu/sample-10BT/train参数速查表参数对应的 Trainer 参数默认值lrlearning_rate3e-4schedulerlr_scheduler_typecosine_with_min_lrbatchbatch_size即per_device_train_batch_size32updategradient_accumulation_steps1contextcontext_length2048gpusnum_gpus_per_node8nodesnum_nodes1warmupwarmup_steps1024stepsmax_steps20480其中model参数指向 legacy/training/configs 下的模型配置。以 gla_340M.json 为例340M 规模的关键配置为hidden_size1024、num_heads4、num_hidden_layers24、hidden_ratio4、expand_k0.5即 key 投影维度为 head 维度的 0.5 倍、attn_modechunk、vocab_size32000、fuse_normtrue与fuse_cross_entropytrue启用融合算子。同目录还准备了 gla_1B.jsonhidden_size2048、24 层、gla_7B.jsonhidden_size4096、32 层、num_kv_heads8以及 transformer_340M.json标准注意力对照配置。global_batch_size的计算每个 batch 实际处理的 token 总数global_batch_size按如下公式计算global_batch_size batch_size × gradient_accumulation_steps × context_length × num_gpus_per_node × num_nodes以 340M 示例为例32 × 1 × 2048 × 8 × 1 524,2880.5M tokens/step。由于每个 step 处理global_batch_size个 tokenmax_steps20480对应处理约 10B tokens——这也是实验目录命名为gla-340M-10B的由来。相应地warmup_steps1024即为学习率预热阶段的 step 数。⚠️ README 特别强调修改任何超参数时务必仔细核对global_batch_size、warmup_steps、max_steps三者之间的比例关系否则会改变预训练的有效数据量与调度行为。学习率调度器默认学习率3e-4配合 cosine 调度器cosine_with_min_lr最终衰减到初始学习率的 10%。除此之外run.py 还支持 WSD 调度warmup_stable_decay它自动将训练划分为“warmup → 90% 稳定段 → 10% 衰减段”无需手动指定stable/decay的步数边界。train.sh 底层做了什么阅读 train.sh 源码可以看到它不仅是启动器还承担了大量工作完整超参集除 README 示例中的参数外还有seed默认 42、save保存间隔默认 2048、limitsave_total_limit默认 1、optim默认adamw_torch_fused、decayweight_decay默认 0.01、beta1/beta20.9/0.95、normmax_grad_norm默认 1.0、workers/prefetchdataloader 配置、logging默认 32 步记录一次以及训练精度固定为bf16见 train.sh 的params拼装model的默认值fla-hub/gla-1.3B-100B即默认按 1.3B 预训练权重启动配合from_config决定随机初始化或加载权重分布式配置自动生成当config名包含deepspeed默认configs/deepspeed.yaml时脚本会动态生成 ZeRO-2 的ds_config.jsonallgather_bucket_size5e8、reduce_scattertrue等与 accelerate 配置若config名包含fsdp则生成 FSDP 配置HYBRID_SHARD_ZERO2、SHARDED_STATE_DICT、TRANSFORMER_BASED_WRAP见 train.sh多机参数设置rank、nodes、ip、port后会向accelerate launch追加--machine_rank、--num_processesnodes × gpus、--main_process_ip/port等参数实验归档启动前会把脚本、configs、flame乃至fla包整体拷贝到path实验目录并设置WANDB_NAME/WANDB_PROJECT/WANDB_RUN_ID与WANDB_RESUMEallow离线模式export TRANSFORMERS_OFFLINE1与HF_DATASETS_OFFLINE1说明训练假定数据与 tokenizer 已预先就位。断点续训Resumeflame通过指定 checkpoint 路径恢复中断的训练。与从零训练相比命令只需追加checkpoint参数README 原文示例bash train.sh \ typegla \ lr3e-4 \ steps20480 \ batch32 \ update1 \ warmup1024 \ context2048 \ gpus8 \ nodes1 \ pathexp/gla-340M-10B \ projectfla \ modelconfigs/gla_340M.json \ dataHuggingFaceFW/fineweb-edu \ namesample-10BT \ cachedata/HuggingFaceFW/fineweb-edu/sample-10BT/train \ checkpointexp/gla-340M-10B/checkpoint-8192从源码链路看train.sh 把checkpoint转成--resume_from_checkpointrun.py 将其直接传给Trainer.train()而 checkpoint 的产生则由--save_steps $save默认 2048 步与--save_total_limit 1控制。此外WANDB_RESUMEallow的设置也保证了 wandb 指标在续训后能接上同一曲线。训练过程中的监控通过 wandb 完成train.sh中当WANDB_DISABLED ! true时自动附加--report_to wandbrun_name形如gla.gla-340M-10B。持续预训练从 Mistral-7B 到 GLA-7Bflame支持从预训练 checkpoint 继续训练。README 给出一个代表性案例把 Mistral-7B 的微调转化为 GLA-7BGSA 论文实验的复现路径。流程分两步第一步权重迁移按 GLA-7B 的配置文件全新初始化模型然后把 Mistral-7B 中形状匹配的权重拷贝过来README 原文示例在legacy/training目录下执行../utils即仓库根目录的 utils/convert_from_llama.pycd ../utils python convert_from_llama.py \ --model mistralai/Mistral-7B-v0.1 \ --config ../training/configs/gla_7B.json \ --output ../training/converted/gla-7B cd -convert_from_llama.py 的实际行为值得展开先保存 tokenizer再以precision默认float32可选float16/bfloat16加载 Llama 权重用AutoModelForCausalLM.from_config(config)初始化目标 GLA 模型——注意此处的 GLA 模型保留了与 Llama 相同的q_proj/k_proj/v_proj/o_proj命名这是gla_7B.json配置下模型的权重命名约定因此逐层直接拷贝embed_tokens → embeddings、input_layernorm → attn_norm并同步variance_epsilon、self_attn.q/k/v/o_proj、post_attention_layernorm → mlp_norm、mlp.gate/up/down_proj、最终norm若tie_word_embeddings为 false 则额外拷贝lm_headgla_7B.json中该字段为false每完成一次拷贝都调用torch.testing.assert_close校验一致性保证转换无损词表大小不一致时会告警并截断/随机扩展 embedding——Mistral-7B 与 GLA-7B 均为 32k 词表正好对齐。GLA 中真正“新”的参数门控g相关权重等保持随机初始化这正是持续预训练而非纯微调的语义模型在保留 Llama 主干语义的前提下学习线性注意力的新机制。第二步从转换后的 checkpoint 启动训练README 原文示例bash train.sh \ typegla \ lr3e-5 \ steps10240 \ batch4 \ update8 \ warmup512 \ context2048 \ pathexp/gla-7B-20B \ projectfla \ modelconverted/gla-7B \ dataSlimPajama-627B \ cachedata/SlimPajama-627B/train几个值得注意的点学习率降一个数量级3e-5vs 从零训练的3e-4这是持续预训练的典型做法等效 batch 保持一致batch4 × update8 × 2048 × 8 × 1 524,288与 340M 实验的global_batch_size相同——用小 micro batch 梯度累积换取 7B 模型的显存可行性10240 × 0.5M ≈ 20Btokens对应目录名gla-7B-20Bmodel指向转换后的本地目录converted/gla-7B此时run.py走from_pretrained分支加载权重对应 parser.py 中model_name_or_path的含义模型权重路径或 Hub 标识多机建议README 明确提示单节点微调 7B 模型未必高效条件允许时应使用多机 GPUtrain.sh已内置多机启动逻辑传入rank、nodes、gpus、ip、port即可见 train.sh更大规模的调度方式可参考 accelerate 的多机教程。小结与延伸阅读本文以legacy/training/README.md为主线串起了 flame 的完整训练链路阶段核心文件关键动作安装READMEpip install .acceleratetokenizers0.20.4预处理preprocess.py分词、打包成seq_len样本、save_to_disk缓存训练入口run.pyfrom_config/from_pretrained双模式、Trainer 调度器启动与分布式train.sh超参拼装、DeepSpeed/FSDP 配置生成、accelerate launch参数解析flame/parser.py扩展TrainingArgumentscache_dir、context_length、varlen等数据组装flame/data.pyDataCollatorForLanguageModeling支持 varlen offsets 打包权重迁移utils/convert_from_llama.pyLlama → FLA 权重拷贝与逐层校验模型配置configs340M / 1B / 7B / transformer 对照再次提醒legacy/training为存档代码README 顶部 IMPORTANT 声明新特性开发已转移至基于 torchtitan 的 flame 新项目。但在理解“如何把线性注意力模型从零或从既有权重训起来”这一问题上这套最小化实现——global_batch_size的推导、WSD 调度器配置、DeepSpeed ZeRO-2 自动生成、Llama→GLA 权重迁移校验——依然是极具参考价值的工程范本。【免费下载链接】flash-linear-attention Efficient implementations for emerging model architectures项目地址: https://gitcode.com/GitHub_Trending/fl/flash-linear-attention创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价