动手跑通attorch的4个实战示例:MNIST、GPT与ImageNet训练,基准测试对标PyTorch 动手跑通attorch的4个实战示例MNIST、GPT与ImageNet训练基准测试对标PyTorch【免费下载链接】attorchA subset of PyTorchs neural network modules, written in Python using OpenAIs Triton.项目地址: https://gitcode.com/gh_mirrors/at/attorchattorch 是一个用 Python 和 OpenAI Triton 编写的 PyTorch 神经网络模块子集支持训练与推理双向计算。本文带你动手跑通它自带的 4 个实战示例——MNIST 分类、GPT 语言模型、ImageNet 图像分类和回归任务并自动完成与 PyTorch 的基准测试对比。为什么值得关注 attorch对于想深入自定义深度学习算子、又不想从零写 CUDA 内核的开发者来说attorch 是一个理想的起点✅纯 Python 实现基于 Triton 内核代码自包含、单文件、可读性强✅训练推理全支持完整支持前向与反向传播支持自动混合精度AMP✅覆盖视觉与 NLP不同于只聚焦 Transformer 的同类库attorch 还提供卷积、池化、BatchNorm 等视觉层✅无缝对接 PyTorchattorch/nn.py 提供attorch.nn接口缺失的层会自动回退到 PyTorch 实现可像torch.nn一样直接替换使用需要说明的是attorch 自带的卷积和池化层性能远慢于 PyTorch因此官方在示例中会通过attorch.nn优先调用 PyTorch 的卷积实现这也是实际使用时的推荐做法。环境安装3 步搞定attorch 的依赖非常轻量只需固定版本的两个库安装依赖pip install torch2.4.0 triton3.0.0克隆仓库git clone https://gitcode.com/gh_mirrors/at/attorch按需安装示例额外依赖示例额外依赖MNIST / Imagenettetorchvision0.19.0WikiText-2datasets2.18.0、transformers4.39.0Regression无示例一MNIST 手写数字分类入门首选这是最适合新手的第一个 attorch 实战示例用多层感知机MLP在经典 MNIST 数据集上做 10 分类。python -m examples.mnist.main常用参数--hidden_dim隐藏层特征数默认 128--depth隐藏层数量默认 1--epochs/--batch_size训练轮数与批大小模型定义见 examples/mnist/mlp.py核心技巧只有一个开关use_attorchTrue时nn.Linear nn.ReLU会被替换成 attorch 的融合算子attorch.Linear(dim, hidden_dim, act_funcrelu)——线性变换和激活在同一个 Triton 内核里完成这就是 attorch 的内核融合用法也是它与 PyTorch 写法的主要区别。示例二WikiText-2 上的 GPT 语言模型NLP 实战第二个示例训练一个 GPT-2 结构的语言模型是体验 attorch 在 NLP 领域MultiheadAttention、RMSNorm 等算子的最佳场景。python -m examples.wikitext-2.main常用参数--downsize将 GPT-2 原始深度和宽度除以该因子小显卡建议调大以省显存--scheduler学习率策略one-cycle或cosine--seq_len/--batch_size序列长度与批大小该示例使用了梯度缩放GradScaler 混合精度训练完整训练流程参考 examples/wikitext-2/main.py模型结构在 examples/wikitext-2/gpt.py。示例三ImageNet 图像分类视觉实战第三个示例在 Imagenette 数据集ImageNet 的 10 类小版本上训练视觉模型可选模型阵容豪华ResNet 系列resnet18 ~ resnet152及 resnet14、resnet26 等精简版ConvNeXt 系列从 convnext_atto 到 convnext_xlargeViT 系列从 vit_tiny_patch16 到 vit_large_patch14python -m examples.imagenette.main --model resnet18模型定义分别位于 examples/imagenette/resnet.py、examples/imagenette/convnext.py 和 examples/imagenette/vit.py。建议新手先用resnet14或convnext_atto这类小模型快速验证流程。示例四合成数据回归零下载、零门槛回归示例使用随机生成的合成数据训练一个简单的 MLP无需下载任何数据集是冒烟测试环境是否装好的最佳选择python -m examples.regression.main它直观演示了attorch.nn的 drop-in 替换能力同一套模型代码只需在attorch.nn和torch.nn之间切换后端其余训练逻辑完全不变。基准测试自动对标 PyTorch每个示例的运行脚本都内置了双重实验设计先完整跑一遍 attorch 后端再跑一遍 PyTorch 后端并各自输出三项指标由 examples/utils.py 中的benchmark_fw_and_bw实现指标含义Forward pass mean execution time前向传播平均耗时Backward pass mean execution time反向传播平均耗时Forward plus backward单步训练总耗时这样你无需自己写任何计时代码一条命令就能得到 attorch 相对 PyTorch 的性能对比数据。常见问题与小贴士版本必须固定attorch 依赖torch2.4.0和triton3.0.0版本不匹配会直接报错️需要 NVIDIA GPU所有示例都会将模型放到cuda设备精度差异部分单元测试可能因浮点精度问题失败实际训练场景中通常不影响结果验证正确性每个模块都有与 PyTorch 对照的测试位于tests/目录可用pytest运行加--subset参数可快速跑子集写在最后从 MNIST 到 GPT-2从 ResNet 到 ViTattorch 用 4 个开箱即用的示例覆盖了分类、语言建模、视觉与回归四大主流任务并且每次运行都自带与 PyTorch 的性能基准对比。无论你是想学习 Triton 算子怎么写还是想为自己的项目寻找一个可快速二开的深度学习模块集合这个仓库都值得花一个下午动手跑一遍。【免费下载链接】attorchA subset of PyTorchs neural network modules, written in Python using OpenAIs Triton.项目地址: https://gitcode.com/gh_mirrors/at/attorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考