Skip to content

【训练营】小模型训练支持 - #228

Open
linglingsansan0907 wants to merge 5 commits into
InfiniTensor:masterfrom
linglingsansan0907:feature/mnist-cnn
Open

linglingsansan0907 wants to merge 5 commits into
InfiniTensor:masterfrom
linglingsansan0907:feature/mnist-cnn

Conversation

@linglingsansan0907

Copy link
Copy Markdown

概述

为 InfiniTrain 补齐小模型(CNN)训练能力:新增 Conv2d / ReLU 算子(CPU + CUDA 前向与反向),基于 MNIST 完成端到端 CNN 手写数字识别 Demo,与 PyTorch 参考实现完成数值对齐,并进一步支持 DDP 多卡分布式训练(单进程多线程 + NCCL AllReduce)。

功能清单

  • autograd::Conv2d 节点:前向 + input / weight / bias 三路反向
  • nn::Conv2d 模块:KaimingUniform(a=√5) 初始化,与 PyTorch 对齐
  • ReLU:autograd 节点、nn 模块、CPU / CUDA elementwise kernel
  • CPU / CUDA Conv2d 四个 kernel(forward / backward_input / backward_weight / backward_bias)
  • example/mnist:MNISTCNN(Conv2d(1,16,3)→ReLU→Conv2d(16,32,3)→ReLU→Flatten→Linear),支持单卡与 DDP 多卡训练
  • example/mnist 既有问题修复:模块 shared_ptr 管理(bad_weak_ptr)、Dataset operator[] 的 float32 字节数
  • 单元测试:test_autograd_conv2d.cctest_autograd_relu.cc,共 14 个测试全部通过

数值对齐(PyTorch)

固定初始权重与输入,与 PyTorch 参考实现逐张量比对(22 个张量,含 forward logits、loss、参数梯度、单步 optimizer 更新):

  • CPU 最大误差 2.384e-07,CUDA 最大误差 1.453e-07
  • 判定阈值 1e-3,实际误差低 4 个数量级

端到端训练结果(MNIST,3 epochs)

配置 最终 test acc 吞吐
单卡(bs=64) 0.9827 ~8558 samples/s
DDP 1 卡 与单卡逐位一致(验证 DDP 封装数值无损) 同单卡
DDP 2 卡(bs=64) 0.9783 ~16934 samples/s(≈1.98×)
DDP 4 卡(bs=60) 0.9715 ~31070 samples/s(≈3.63×)

说明:DDP 实现为单进程多线程(gpt2 示例风格),梯度经 NCCL AllReduce(kAvg)聚合。训练批数需能被卡数整除,故 4 卡使用 bs=60(1000 % 4 = 0)。DDP2/4 与单卡曲线不逐位一致,原因是全局等效 batch 变大且各 rank 数据分片不同,数值正确性由 DDP1 与单卡逐位一致证明。

框架级问题说明

  • LinearBackwardBias CUDA kernel 的转置读取 bug:开发过程中独立发现 bias 反向归约按转置坐标读取输入,并基于 row_stride / col_stride 完成修复;rebase 到最新 master 后发现上游已包含等价修复(ReduceRowsKernel 重写),故本 PR 采用上游实现,不再包含该文件改动。
  • ExpBackward 上游既有 segfault(本 PR 不处理):仅在测试与报告中记录说明。

构建说明

DDP 需要 NCCL。根 CMakeLists.txt 修复了 NCCL include/link 传播问题(原逻辑 find 到 NCCL 但未向 CUDA kernels 与主框架目标传递 include 路径与库)。无 sudo 环境下可用 pip install nvidia-nccl-cu12 抽取到本地前缀后通过 NCCL_ROOT 指定,详见示例文档。

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant