ML Training Recipes 实战:基于 Scaling Laws 的架构选择、算力预算与带宽受限训练指南
2026/9/23 18:37:25 网站建设 项目流程

ML Training Recipes 实战:基于 Scaling Laws 的架构选择、算力预算与带宽受限训练指南

【免费下载链接】AI-Research-SKILLsComprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.项目地址: https://gitcode.com/gh_mirrors/ai/AI-Research-SKILLs

本篇技术指南以 AI-Research-SKILLs 仓库中 ml-training-recipes 技能包的 scaling-and-selection.md 参考文档为核心骨架,系统讲解如何在训练神经网络前,依据数据规模、算力预算与任务类型做出正确的架构与规模决策。读完本文,你将掌握 Chinchilla 最优计算配比、按数据类型与计算预算的架构决策树、FLOPs/显存预算估算、优化器选型,以及如何在 DGX Spark 这类带宽受限硬件上最大化训练吞吐。


一、Scaling Laws:用规模定律指导训练预算

1.1 Chinchilla 最优计算配比(Hoffmann et al., 2022)

对 LLM 训练而言,最重要的规模定律是 Chinchilla 结论:在计算最优的前提下,参数量 N 与训练 token 数 D 应随计算量等比例增长,配比约为每个参数 20 个 token。这一结论直接推翻了早期 Kaplan 等人"模型越大越好、数据相对少放"的建议——训练预算固定时,把预算花在数据上而非一味放大模型,往往更划算。

其背后的 FLOPs 估算公式为:

FLOPs ≈ 6 × N × D 其中: N = 参数量 D = 训练 token 数 6 = 每个参数每个 token 的前向(2 FLOPs)+ 反向(4 FLOPs)

该公式在仓库中同样被用于 MFU(Model FLOPs Utilization)监控与 FLOPs 预估。architecture.md 给出了对应的估算实现,核心规则是"6 × N(前向 matmul 2、反向 matmul 4),并剔除 embedding 这类稀疏查找":

def estimate_flops_per_token(model): """Forward + backward FLOPs per token (approx 6 * params + attention).""" nparams_dense = sum(p.numel() for p in model.parameters()) nparams_dense -= model.wte.weight.numel() # token embedding 不计入 nparams_dense -= model.lm_head.weight.numel() # 若已 tied 则已计数 # 注意力部分: 2 * n_heads * head_dim * seq_len per layer (Q@K + attn@V) attn_flops = 0 for window in model.window_sizes: effective_seq = min(window[0], model.config.sequence_len) attn_flops += 12 * model.config.n_head * head_dim * effective_seq return 6 * nparams_dense + attn_flops

ml-training-recipes 主 SKILL 将 Chinchilla 规则直接落实为可查表:

模型规模计算最优 token 数推理最优 token 数(100×)
125M2.5B tokens12.5B tokens
1B20B tokens100B tokens
7B140B tokens700B tokens

1.2 Chinchilla 最优 vs 推理最优

策略Tokens/Param适用场景示例
Chinchilla 最优~20×研究、一次性计算开销7B 模型 → 140B tokens
推理最优100-200×生产部署7B 模型 → 700B-1.4T tokens

LLaMA 系列的哲学正是后者:部署一个"更小但喂了更多数据"的模型。因为推理是持续成本(每次请求都在花钱),而训练只是一次性成本。生产环境反复调用的小模型,其总拥有成本往往远低于训练一次更优的大模型。

1.3 Chinchilla 之后的新认知

  • Muennighoff et al. (2023):数据重复训练最多 4 个 epoch,其效果约为同量唯一数据的 85%;超过 4 epoch 后收益急剧衰减。有效数据量近似满足D_effective ≈ D × (1 - e^{-epochs})
  • 过训练(over-training)小模型现已成为生产环境标准做法(LLaMA、Mistral、Phi 系列皆如此)。
  • 数据质量 >> 数据规模(Llama 3 的发现):激进去重 + 质量过滤,比单纯堆数据量更能提升效果。

二、架构决策树:按数据与算力选型

2.1 按数据类型的主决策树

你的数据是什么类型? │ ├─ 图像 / 视频 │ ├─ 数据 < 10K → 预训练 CNN(ResNet/EfficientNet)+ 微调分类头 │ ├─ 数据 10K-1M → 预训练 ViT 微调 或 CNN 微调(两者皆可行) │ ├─ 数据 > 1M → ViT 或混合架构(ConvNeXt、CoAtNet)从零训练 │ └─ 视频 → Video Swin Transformer 或 TimeSformer(预训练) │ ├─ 文本 / NLP │ ├─ 分类/NER → 微调编码器(BERT/RoBERTa) │ ├─ 生成 → 微调解码器(GPT/LLaMA) │ ├─ Seq2seq(翻译) → 微调 T5/BART │ ├─ 数据 < 1K 样本 → 大 LLM 少样本推理(不训练) │ ├─ 序列长度 > 8K → 考虑 Mamba-hybrid 或长上下文 Transformer │ └─ 推理预算紧张 → 蒸馏模型、RWKV 或 Mamba │ ├─ 表格数据 │ ├─ 行数 < 50K → XGBoost / LightGBM(不要用深度学习) │ ├─ 行数 50K-500K → GBM 依然强劲;可尝试 FT-Transformer 对比 │ └─ 行数 > 500K → 神经网络可行;两者都 benchmark │ ├─ 时间序列 │ ├─ 单变量、短预测期 → ARIMA / Prophet / 简单 LSTM │ ├─ 多变量、中等数据 → LSTM/GRU 或 N-BEATS │ ├─ 长序列 / 大量序列 → PatchTST / Informer / Mamba │ └─ 已有基础模型 → TimesFM 或 Chronos(微调) │ ├─ 音频 / 语音 │ ├─ 语音识别 → Whisper(预训练)+ 微调 │ ├─ 音频分类 → AST 或基于频谱图的 CNN │ └─ 长音频 → Mamba / SSM 变体 │ ├─ 图数据 │ └─ GNN(GCN、GAT、GraphSAGE);大图用 Transformer-on-graphs │ └─ 多模态 └─ CLIP 风格(视觉+文本),或统一 Transformer(Gemini 风格)

主 SKILL 提供了更贴近仓库实际使用场景的版本,覆盖了生物医学等更多领域:

数据类型< 10K 样本10K-100K> 100K
图像预训练 CNN + 微调微调 ViT 或 CNNViT 从零训练
文本(生成)少样本提示微调 GPT/LLaMA(LoRA)从零预训练
表格XGBoost/LightGBM仍然是 XGBoost神经网络可行
音频预训练 Whisper微调 AST从零训练
分子预训练 GNN微调分子 LM从零训练 GNN
蛋白质ESM-2 嵌入 + 头微调 ESM-2训练蛋白质 LM
医学图像预训练 CNNnnU-Net(自动配置)Swin-UNETR / MedSAM

2.2 按计算预算的决策树

你有多少计算资源? │ ├─ 单 GPU,< 1 天 │ → 模型 < 500M 参数 │ → 微调预训练模型,不要从零训练 │ → 大模型微调用 LoRA/QLoRA │ ├─ 单 GPU,1-7 天 │ → 最多从零训练 1B 参数 │ → 或用 QLoRA 微调最多 7B │ ├─ 多 GPU(4-8),1-7 天 │ → 从零训练最多 3B │ → 或微调最多 13B │ → 使用 DDP 做数据并行 │ ├─ 集群(32+ GPU),数周 │ → 从零训练 7B+ │ → 应用 Chinchilla 缩放:至少 20 tokens/参数 │ → 使用 FSDP 或 Pipeline Parallel │ └─ 大规模集群(数百 GPU),数月 → 70B+ 模型 → 完整 5 维并行(TP + PP + DP + EP + CP) → Chinchilla 配比至关重要

与此对应,domain-specific.md 给出了分布式训练的实现模式:DDP 适用于中等规模(DistributedSampler+set_epoch保证每个 epoch 正确打乱),FSDP 适用于大模型(size_based_auto_wrap_policy自动包裹大层,配合 bf16 混合精度)。规模放大时的缩放规则包括:线性缩放(batch 放大 k 倍则 LR 放大 k 倍,有上限)、平方根缩放lr_new = lr_base * sqrt(batch_new / batch_base),更保守、往往更稳)、以及用梯度累积模拟大 batch 而无须增加 GPU。


三、数据规模阈值:各类任务的交叉点

3.1 视觉:CNN 与 ViT 的交叉点

数据集大小胜出者备注
< 5K 张图预训练 CNN无预训练时 ViT 会过拟合
5K-50K微调 ViT ≈ CNN两者皆可,ViT 需要预训练(ImageNet-21k)
50K-500KViT + 预训练略占优混合架构(CoAtNet)表现出色
> 1MViT 从零训练可行ViT-L/H 超过 CNN
> 10MViT 明显胜出原 ViT 论文在 JFT-300M 上验证

关键洞察:迁移学习抹平了差距。在大数据上预训练、再在小数据上微调的 ViT,可以击败在小数据上从零训练的 CNN。这也是 3.1 节决策树中"数据 < 10K 时用预训练 CNN"的根本原因——问题不在于架构,而在于是否携带先验知识。

3.2 NLP:模型规模阈值

任务数据量方案
< 100 条样本少样本提示(不训练)
100-1K微调小模型(BERT-base)或在大模型上做 LoRA
1K-10K全量微调中等模型
10K-100K训练领域专用模型或继续预训练
> 100K按 Chinchilla 配比同步放大模型与数据

3.3 表格数据:树模型的边界

Grinsztajn et al. (2022) 的论文标题即结论——Why do tree-based models still outperform deep learning on typical tabular data?

数据行数建议
< 10KXGBoost/LightGBM(无需讨论)
10K-50K树模型几乎总是赢,神经网络勉强有竞争力
50K-500K神经网络(FT-Transformer、TabNet)开始可行
> 500K两者皆具竞争力;高基数特征下神经网络可能胜出

这是机器学习领域最稳健的发现之一:在约 50K 行以下的典型表格数据上,神经网络极少能击败梯度提升树。选择"正确但无趣"的 GBM,而不是"时髦但低效"的深度网络,本身就是一种工程优化。

3.4 时间序列阈值

数据规模架构
< 1K 条序列经典方法(ARIMA、Prophet)或简单 LSTM
1K-100KLSTM/GRU 有竞争力,Transformer 变得可行
> 100K长预测期的 Transformer 变体或 Mamba

四、计算预算规划:FLOPs 与显存估算

4.1 按模型规模的 FLOPs 估算

模型规模token 数(Chinchilla)训练 FLOPsA100 GPU 小时(估算)
125M2.5B1.9e18~6h
350M7B1.5e19~48h
1B20B1.2e20~385h
7B140B5.9e21~19,000h
13B260B2.0e22~65,000h
70B1.4T5.9e23~1.9M h

这张表可以直接回答"我要训多大的模型":先看总预算 GPU 小时数,反查可承担的规模与 token 量,再按第二节的决策树校验是否在你的硬件拓扑内可行。

4.2 显存估算的经验法则

bf16 训练下的模型显存占用经验公式:

总显存 ≈ 18-20 × N_params(字节) 分解: 模型权重(bf16): 2 × N 字节 梯度(bf16): 2 × N 字节 优化器状态(Adam): 8 × N 字节(fp32 一阶+二阶矩) 激活值: 视情况(约 4-8 × N) 示例:1B 参数 → 至少 18-20 GB 显存

降低显存的技术手段(按优先级):

  • 梯度检查点(Gradient checkpointing):激活值内存 -50-70%,代价是 +30% 计算
  • 8-bit 优化器:优化器状态内存 -30%
  • FSDP:将状态分片到多张 GPU
  • QLoRA:4-bit 基座 + LoRA 适配器

主 SKILL 进一步给出了 OOM 的完整解决顺序,其中前几项即对应上表:

  1. 减小DEVICE_BATCH_SIZE,增大grad_accum_steps
  2. 设置PYTORCH_ALLOC_CONF=expandable_segments:True
  3. model.zero_grad(set_to_none=True)(比置零更省内存)
  4. Meta device 初始化 →to_empty(大模型零内存创建)
  5. 激活检查点:torch.utils.checkpoint.checkpoint()
  6. 8-bit 优化器(bitsandbytes):优化器状态约省 30%

其中"Meta device 初始化"是处理超大模型的利器,在主 SKILL 中有完整示例:

with torch.device("meta"): model = GPT(config) # 零内存创建 model.to_empty(device="cuda") model.init_weights()

4.3 用 MFU 校准预算

预算规划不止于"能放下",还要"跑得快"。主 SKILL 给出的 MFU 公式为:

achieved_flops = model_flops_per_token * batch_tokens / step_time mfu = achieved_flops / gpu_peak_flops # H100 SXM: 989.5 TFLOPS | A100: 312 | RTX 4090: 165

单卡目标参考:>30% 合格、>40% 良好、>50% 优秀。低 MFU 的排查顺序(来自主 SKILL 调试清单):确认torch.compile生效 → 检查torch.set_float32_matmul_precision("high")→ 锁页内存 +non_blocking传输 → 用torch.profiler分析 → 处理 GC 停顿 → 核对 Tensor Core 对齐(维度为 8/64 的倍数)。

4.4 时间预算:让实验可比较

experiment-loop.md 提出了一个与算力预算直接相关的实践:用固定时间预算而非固定步数/epoch 来定义一次实验。墙钟时间天然包含了吞吐差异,因此用 5 分钟固定预算跑实验,结果可直接横向比较(约 12 个实验/小时,一晚可跑上百个)。这也是第二节"1-7 天""数周"等预算表述在工程上的落地方式。


五、优化器选择指南

5.1 优化器对比总表

优化器最适合显存备注
AdamW一切场景的默认2× 参数β1=0.9, β2=0.95(LLM)
8-bit Adam(bitsandbytes)显存受限~1.3× 参数质量几乎无损
Adafactor超大模型~1× 参数分解二阶矩
SGD+momentum视觉 CNN1× 参数需要更多 LR 调参
MuonTransformer 矩阵~2× 参数正交更新,前沿方向
LAMB/LARS超大 batch(>32K)2× 参数逐层缩放 LR 保证稳定
Lion(Google)值得一试1× 参数基于符号,比 Adam 省显存
Schedule-Free Adam追求简单2× 参数无需 LR 调度
SOAPLLM 训练~2× 参数Shampoo 类但更实用

5.2 何时用哪个

  • 默认选择:AdamW。永远有效、理解透彻、文献海量。
  • 显存压力:8-bit Adam 或 Adafactor。
  • 超大 batch:LAMB/LARS(否则线性缩放规则失效)。
  • 前沿 LLM:Muon 处理矩阵参数 + AdamW 处理 embedding(autoresearch 模式)。
  • 追求简单:Schedule-Free Adam——彻底去掉 LR 调度。

5.3 仓库中的优化器源码级细节

optimizers.md 对 AdamW 给出了与默认 PyTorch 参数的关键差异:

  • β2 = 0.95(非默认 0.999):默认值约有 ~1000 步的记忆窗口,对 LLM 训练快速变化的损失面太慢;β1 可选 0.8-0.9 以获得更快的动量。
  • eps = 1e-10(非默认 1e-8):bf16 训练中梯度可能非常小,1e-8 会导致二阶矩极小时更新停滞。autoresearch 使用 1e-10。
optimizer = torch.optim.AdamW( params, lr=3e-4, betas=(0.9, 0.95), # β1=0.9, β2=0.95(不是默认 0.999) eps=1e-8, # LLM 训练常改 1e-10 weight_decay=0.1, )

Muon是专为 2D 矩阵(权重)参数设计的优化器:Nesterov 动量 + "Polar Express" 正交化(用快速 Newton-Schulz 迭代近似矩阵极分解,找到距离梯度最近的正交矩阵)。正交化的动机是:普通梯度下降长期更新后可能让权重矩阵秩亏,正交化更新方向能鼓励特征多样化、防止权值空间模式坍缩。其关键超参数为:

参数典型值备注
lr0.02-0.04非方阵按max(1, rows/cols)^0.5缩放
momentum0.95前 300 步从 0.85 预热
ns_steps5Newton-Schulz 迭代次数(越多越准、越慢)
beta20.95NorMuon 二阶矩跟踪
weight_decay0.1-0.2Cautious(仅在梯度与参数同号处生效)

混合 MuonAdamW是现代 LLM 训练的核心模式——不同参数类型用不同优化器:

参数类型优化器原因
2D 权重矩阵(attention、MLP)Muon受益于正交化
Token embeddingsAdamW稀疏更新,不是矩阵变换
Unembedding(lm_head)AdamW需要更低 LR 保持稳定
逐层标量AdamW太小,不适合矩阵方法
Value embeddingsAdamW同 token embeddings

完整的分组配置示例(来自 optimizers.md):

def setup_optimizer(model, d_model=768): lr_scale = (d_model / 768) ** -0.5 param_groups = [ # Unembedding: 低 LR、无 weight decay { 'kind': 'adamw', 'params': list(model.lm_head.parameters()), 'lr': 0.004 * lr_scale, 'betas': (0.8, 0.95), 'eps': 1e-10, 'weight_decay': 0.0, }, # Token embeddings: 更高 LR(稀疏更新需要更大步长) { 'kind': 'adamw', 'params': list(model.wte.parameters()), 'lr': 0.6 * lr_scale, 'betas': (0.8, 0.95), 'eps': 1e-10, 'weight_decay': 0.0, }, # Transformer 矩阵: Muon { 'kind': 'muon', 'params': list(model.transformer.h.parameters()), 'lr': 0.04, 'momentum': 0.95, 'ns_steps': 5, 'beta2': 0.95, 'weight_decay': 0.2, }, # 逐层标量: 单独 AdamW { 'kind': 'adamw', 'params': [model.resid_lambdas], 'lr': 0.005 * lr_scale, 'betas': (0.8, 0.95), 'eps': 1e-10, 'weight_decay': 0.0, }, ] optimizer = MuonAdamW(param_groups) return optimizer

其中LR 随维度缩放规则为lr_effective = lr_base * (d_model / d_reference)^(-0.5):模型越宽,单参数学习率应越低,因为大矩阵会放大梯度范数,按 1/√d 缩放可让有效步长跨模型尺寸保持恒定。Weight decay 策略上,LLM 训练中应只对 Transformer 权重矩阵做 decay,embeddings、bias、LayerNorm 参数与逐层标量不做;更激进的方案是线性衰减 WD 到 0(训练后期让模型完全承诺于已学特征),Muon 则使用 cautious decay(仅在与梯度同号处衰减,避免 WD 与梯度"打架")。


六、大规模训练的不稳定性与对策

6.1 常见失败模式(OPT-175B、BLOOM、PaLM、Llama 实测)

失败症状修复
Loss spikes损失突然跳升,可能恢复也可能不恢复降 LR、跳过该 batch、回滚到更早 checkpoint(PaLM 策略)
Slow divergence损失逐渐上升数据质量问题或 LR 过高
Embedding collapse所有 embedding 收敛到相近值加 embedding LayerNorm、降低 embedding LR
Attention entropy collapse注意力均匀化或 one-hotz-loss 正则化、QK-norm
fp16 下 NaN训练崩溃换 bf16,或在 matmul 前调整归一化顺序

6.2 PaLM loss spike 策略

检测到 loss spike 时:

  1. 回滚到 spike 之前的最后一个 checkpoint
  2. 跳过导致 spike 的数据 batch
  3. 可选:临时降低 LR,再逐步回升
  4. 恢复训练

这套流程如今已成为多数大型训练实验室的标准操作。experiment-loop.md 将其纳入了自主实验循环的决策纪律:指标变差即 discard(git reset --hard HEAD~1)、崩溃分三类处理(trivial 修复重试、fundamental 记录后跳过、超时按 timeout 记录),保证每步只做一个改动、diff 可审、回滚干净。

6.3 稳定性技术(现已成标准)

  • Pre-norm(在 attention/FFN 之前归一化,而非之后)
  • QK-norm(点积前对 Q 和 K 归一化)
  • 线性层不加 bias(除最终输出层外)
  • 梯度裁剪max_norm=1.0
  • Embedding LayerNorm(大规模训练尤其重要)
  • bf16 优先于 fp16(无需 loss scaling)

architecture.md 给出了这些技术的实现级细节:pre-norm 的标准写法是x = x + self.attn(norm(x))x = x + self.mlp(norm(x)),最终输出到 lm_head 前再做一次x = norm(x);QK-norm 在 RoPE 应用之后、进入注意力之前执行q, k = norm(q), norm(k)Logit soft cappingsoftcap * torch.tanh(logits / softcap)(如 softcap=15)平滑钳制极端 logit;零初始化输出投影让残差流从干净状态起步;此外还有可学习的逐层残差缩放(resid_lambdas初始为 1.0、x0_lambdas初始为 0.1)帮助梯度直达 embedding 层。主 SKILL 的"2025 默认配方"将这些收敛为可直接照抄的参数组合:AdamW(β1=0.9, β2=0.95, eps=1e-10)、weight decay 0.1、Cosine 或 WSD 调度、峰值 LR 3e-4(大模型按比例下调)、bf16、max_norm=1.0、RMSNorm(pre-norm)、SwiGLU、RoPE、Flash Attention(可选 GQA)。

Loss exploding / NaN 的完整排查链(主 SKILL):先降 LR(3-10×)→ 加梯度裁剪 → 检查输入中的 inf/nan → 加 logit soft capping → 加 QK-norm → 核对权重初始化 → 检查梯度累积下的 loss 归约(loss / grad_accum_steps)。


七、DGX Spark / 带宽受限 GPU 训练专题

7.1 GB10 Grace Blackwell 规格与瓶颈

规格数值对比 H100 SXM
GPU 显存128 GB LPDDR5X(CPU+GPU 统一)80 GB HBM3
显存带宽~273 GB/s~3,350 GB/s(少 12×
CPU-GPU 互连NVLink C2C(~900 GB/s)N/A(分立)
FP4 Tensor Core支持(Blackwell)不支持
FP8 Tensor Core支持支持
bf16 峰值 TFLOPS~TBD(Blackwell 架构)989.5
功耗~300W 整机单 GPU 700W
形态桌面工作站数据中心

DGX Spark 最大的约束是显存带宽——比 H100 少 12×,由此推断:

  • 计算受限算子(大 matmul):表现正常,每 FLOP 效率相近
  • 内存受限算子(逐元素、归约、attention):严重受限
  • 相同模型的有效 MFU 会低于 HBM GPU

经验法则:当算子的算术强度(FLOPs/byte)< 50 时,它在 DGX Spark 上就是带宽受限的。增大 batch、加宽模型能提高算术强度。

7.2 带宽受限训练的七大优化策略

策略 1:最大化计算/内存比
# 用更大 batch 提高 matmul 的算术强度 # 更大 batch → 每次权重加载对应更多 FLOPs → 更好利用带宽 # 用梯度累积模拟大 batch 而不 OOM grad_accum_steps = 16 # 等效 16 倍 batch
策略 2:量化训练(FP8 / FP4)

DGX Spark 的 Blackwell 核心原生支持 FP4 与 FP8——它们按比例减少内存流量:

# 使用 transformer engine 做 FP8 训练 import transformer_engine.pytorch as te # 用 FP8 版本替换 nn.Linear linear = te.Linear(in_features, out_features, bias=False) # FP8 autocast with te.fp8_autocast(enabled=True): output = model(input)

FP8 相比 bf16 将带宽需求削减约 2×,FP4(可用时)削减约 4×。既然带宽是瓶颈,这直接转化为速度提升。

策略 3:算子融合
# torch.compile 在带宽受限硬件上至关重要 # 它把逐元素算子(norm、激活、残差加)融合进单一 kernel model = torch.compile(model, dynamic=False, fullgraph=True) # 手动融合示例:融合 RMSNorm + linear # 而非:norm(x) → 写回内存 → linear(normed_x) # 融合:norm + linear 单 kernel 完成,x 永不写回内存
策略 4:梯度检查点(在这里反而更划算)

在 HBM GPU 上,梯度检查点是用计算换内存;在 DGX Spark 上权衡不同——重算激活值可能比从内存加载更快

from torch.utils.checkpoint import checkpoint class Block(nn.Module): def forward(self, x): # 重算 attention 激活值而非存储 x = x + checkpoint(self.attn, x, use_reentrant=False) x = x + checkpoint(self.mlp, x, use_reentrant=False) return x
策略 5:统一内存优势

CPU 与 GPU 间的 NVLink C2C(~900 GB/s)意味着:

  • 无需显式 CPU↔GPU 拷贝——统一地址空间
  • 可以训练大于 GPU VRAM的模型,而无 offloading 开销
  • torch.cuda.mem_get_info()检查可用统一内存
  • 128GB 池是共享的——要监控的是整个系统内存,而不只是"GPU 显存"
策略 6:推理侧的 KV-cache 优化

对 DGX Spark 上的 LLM 推理,KV-cache 是带宽瓶颈:

  • GQA/MQA:更少的 KV 头 = 更小的 cache = 更少带宽
  • KV-cache 量化:INT8 或 FP8 KV cache 将带宽降低 2-4×
  • 滑窗注意力:无论序列多长,cache 大小有界
  • PagedAttention(vLLM):变长序列的高效内存管理

GQA 的模型侧实现见 architecture.md:n_kv_head = n_head为 MHA(满质量、最费内存),n_kv_head = n_head / 4为常见 GQA 取舍,n_kv_head = 1为 MQA(最省内存、轻微质量损失),且要求n_head % n_kv_head == 0

策略 7:DGX Spark 模型选型
模型规模可行性备注
< 1B优秀从零训练,快速迭代
1-7B良好可从零训练;微调无压力
7-13B可行QLoRA 微调;从零训练较慢
13-30B仅微调QLoRA;统一内存有助于放下模型
30-70B仅推理需量化(GPTQ/AWQ 4-bit)
> 70B不推荐连推理都可能太慢

7.3 DGX Spark 训练检查清单

  • 启用 FP8 训练(transformer_engine)——最大单项收益
  • 使用torch.compilefullgraph=True做算子融合
  • 在显存允许范围内尽量增大 batch(提高算术强度)
  • 启用梯度检查点(在带宽受限硬件上是"免费"性能)
  • 注意力重的模型使用 GQA/MQA
  • 监控torch.cuda.max_memory_allocated()——统一内存意味着不同的上限
  • torch.profiler找出带宽受限的 kernel
  • 若 Blackwell kernel 支持,推理考虑 FP4

八、关键参考文献

以下为规模定律与架构选择领域的核心文献(编号对应 arxiv,可按需检索原文):

Scaling Laws

  • Kaplan et al. (2020):Scaling Laws for Neural Language Models(arxiv:2001.08361)
  • Hoffmann et al. (2022):Training Compute-Optimal Large Language Models(Chinchilla,arxiv:2203.15556)
  • Muennighoff et al. (2023):Scaling Data-Constrained Language Models(arxiv:2305.16264)

架构选择

  • Dosovitskiy et al. (2020):An Image is Worth 16x16 Words(ViT,arxiv:2010.11929)
  • Liu et al. (2022):A ConvNet for the 2020s(ConvNeXt,arxiv:2201.03545)
  • Grinsztajn et al. (2022):Why do tree-based models still outperform deep learning on tabular data?(arxiv:2207.08815)

替代架构

  • Gu & Dao (2023):Mamba: Linear-Time Sequence Modeling(arxiv:2312.00752)
  • Peng et al. (2023):RWKV: Reinventing RNNs for the Transformer Era(arxiv:2305.13048)
  • Sun et al. (2023):Retentive Network(RetNet,arxiv:2307.08621)

训练配方与方法论

  • Karpathy (2019):A Recipe for Training Neural Networks(博客)
  • Wightman et al. (2021):ResNet Strikes Back(arxiv:2110.00476)
  • Yang et al. (2022):Tensor Programs V(µP,arxiv:2203.03466)
  • Google Research:Deep Learning Tuning Playbook
  • Stas Bekman:ML Engineering
  • Geiping & Goldstein (2022):Cramming: Training a Language Model on a Single GPU in One Day(arxiv:2212.14034)

大规模训练

  • Zhang et al. (2022):OPT: Open Pre-trained Transformer Language Models(arxiv:2205.01068)
  • Chowdhery et al. (2022):PaLM: Scaling Language Modeling with Pathways(arxiv:2204.02311)
  • Touvron et al. (2023):LLaMA(arxiv:2302.13971)

九、总结:从决策到执行的完整闭环

将本指南串联成一条可执行路径:先用 Chinchilla 定律与 FLOPs/显存估算表确定规模 → 用数据/算力决策树选定架构 → 按优化器选型表与源码级参数组配置训练 → 用稳定性技术与 PaLM 回滚策略兜底 → 在带宽受限硬件上套用七大优化策略。仓库中配套的 ml-training-recipes 主 SKILL 提供训练循环、LR 调度、混合精度、调试清单等落地代码,optimizers.md 提供 Muon/AdamW 混合配置与编译优化器实现,architecture.md 提供 GQA、滑窗注意力、残差缩放等稳定性组件,domain-specific.md 覆盖视觉/扩散/分布式场景,experiment-loop.md 则给出把上述决策快速验证为实验结果的自主循环——四者合起来,就是从"选什么模型"到"如何把实验跑完"的完整工程闭环。

【免费下载链接】AI-Research-SKILLsComprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.项目地址: https://gitcode.com/gh_mirrors/ai/AI-Research-SKILLs

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

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

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

立即咨询