从大厂裸辞后,我用PyTorch重写了通勤路上的AI梦

技术App
2026-04-14 06:37
阅读 2242

上周五晚上十一点半,我坐在回天通苑的13号线上,耳机里放着《硅谷》OST,手里刷着GitHub Trending。突然看到一个叫v0的项目火了——不是那个Vercel出的AI UI生成器,而是某个PyTorch社区维护的轻量推理工具包。那一刻我忽然意识到:辞职两个月了,除了每天睡到自然醒、研究开源项目源码、顺便思考“我到底适合干啥”,好像还没真正动手写过一行正经代码。

说来惭愧,我在前东家(某字节系大厂)做推荐系统三年,天天和TensorFlow、XLA、TF Serving打交道,PyTorch反而成了“别人家的框架”。虽然嘴上总说“PyTorch更Pythonic”、“动态图调试爽翻”,但真轮到自己写模型,还是习惯性地打开tf.keras……直到上个月面试一家创业公司,对方CTO淡淡一句:“我们全栈PyTorch,你行吗?” 我当场就有点卡壳。

于是,我决定从零开始,认真玩一把PyTorch。不为别的,就为了证明——裸辞不是躺平,是换个姿势卷底层


为什么选PyTorch?别扯生态了,说人话

很多人吹PyTorch生态好、社区活跃、学术圈首选……这些我都认。但对我这种前大厂打工人来说,最打动我的其实是两点:

  1. 调试体验像写Python脚本
    在大厂那会儿,每次改个loss function都要等CI跑20分钟,还得祈祷别被隔壁组merge的代码带崩。而PyTorch的动态图机制,让我可以在Jupyter里逐行print tensor,甚至加断点——这感觉,就像从Kubernetes集群切回localhost开发,自由得让人想哭。

  2. 性能优化路径清晰
    别被“PyTorch慢”的老黄历骗了。自从有了TorchScript、FX Graph Mode、还有最近爆火的torch.compile(基于Inductor后端),PyTorch的推理性能早就不是问题。更重要的是,它把优化选项明明白白摆在你面前,不像某些框架,黑盒优化搞半天还不知道哪块拖了后腿。

说到性能,不得不提一嘴Amazon Q。上周我试了下这个新出的AI编程助手(免费额度快用完了哭),让它帮我把一段NumPy预处理代码转成PyTorch GPU版本。结果它不仅用了torch.from_numpy().cuda(),还顺手加了.contiguous()——这细节,连我司资深算法工程师都经常忘。虽然Amazon Q现在对PyTorch的支持还比不上Copilot,但至少它知道v0不是版本号,而是一个工具名(后面会细说)。


动手!用CIFAR-10练个手,顺便测测通勤时间能不能跑完一轮训练

我选了个经典入门数据集:CIFAR-10。为啥?因为:

  • 图片小(32x32),GPU显存压力小
  • 分类任务简单,适合验证流程是否跑通
  • 我笔记本上的RTX 3060能扛住(感谢前东家发的年终奖)

先装环境。这里有个坑:千万别直接pip install torch!很多教程这么写,但默认装的是CPU版本。正确姿势:

# CUDA 11.8用户(我就是)
pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118

然后写个最简ResNet18:

import torch
import torch.nn as nn
import torchvision.models as models

class SimpleClassifier(nn.Module):
    def __init__(self, num_classes=10):
        super().__init__()
        # 直接拿预训练ResNet18,但换掉最后的fc层
        self.backbone = models.resnet18(pretrained=True)
        self.backbone.fc = nn.Linear(self.backbone.fc.in_features, num_classes)
    
    def forward(self, x):
        return self.backbone(x)

model = SimpleClassifier()
device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device)

数据加载部分,PyTorch的DataLoader真是优雅:

from torchvision import transforms, datasets

transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])

trainset = datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
trainloader = torch.utils.data.DataLoader(trainset, batch_size=128, shuffle=True, num_workers=4)

注意num_workers=4——这是我在大厂调优时学到的经验:CPU预处理线程数要和你的核心数匹配。之前有次上线,测试同学抱怨“数据加载慢”,结果发现是num_workers=0,整个训练卡在CPU解码JPEG上。运维大哥差点把我挂工位上。


性能优化实战:从10分钟/epoch到90秒/epoch

初始版本跑起来,一个epoch要10分钟。这哪行?我通勤才1小时,难道每天只能跑6个epoch?

第一步:启用混合精度训练(AMP)

PyTorch的torch.cuda.amp简直是神器。两行代码,速度翻倍,显存减半:

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

for inputs, labels in trainloader:
    inputs, labels = inputs.to(device), labels.to(device)
    
    optimizer.zero_grad()
    
    with autocast():  # 自动混合精度上下文
        outputs = model(inputs)
        loss = criterion(outputs, labels)
    
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

效果立竿见影:epoch时间降到5分半。而且精度几乎没掉——毕竟CIFAR-10太简单了。

第二步:试试torch.compile(PyTorch 2.0+)

去年底PyTorch 2.0发布时,我还在工位上吐槽:“又是画饼”。结果自己试了才发现,torch.compile真香!

# 只需一行!
model = torch.compile(model)

它会自动将Eager Mode的代码编译成优化后的Graph,类似JIT但更智能。实测在我的3060上,又快了30%——现在一个epoch只要3分40秒。

⚠️ 注意:torch.compile目前对Windows支持一般,建议Linux/macOS用户优先尝试。

第三步:引入v0——那个被误认为版本号的工具

前面提到的v0,其实是PyTorch社区的一个实验性推理优化工具包(GitHub搜pytorch/v0)。它的核心思想是:在部署阶段,用更激进的图优化策略换取极致性能

虽然文档少得可怜(典型的开源项目风格),但代码很干净。我扒了它的optimize_for_inference函数,发现它做了几件事:

  • 融合Conv-BN-ReLU
  • 消除无用tensor
  • 启用CUDA Graph(减少Kernel Launch开销)

集成方式超简单:

# 安装:pip install git+https://github.com/pytorch-labs/v0.git
from v0 import optimize_for_inference

# 训练完后调用
optimized_model = optimize_for_inference(model)

实测推理延迟从12ms降到7ms(batch_size=64)。虽然训练不能用,但对上线服务来说,这5ms可能就是QPS翻倍的关键


和TensorFlow比?别吵了,看数据说话

我知道肯定有人问:“和TF比谁快?” 作为一个前TF重度用户,我做了个公平对比(同模型、同数据、同硬件):

框架 训练速度 (epoch) 推理延迟 (ms) 显存占用 (MB) 调试友好度
TensorFlow 2.12 4分10秒 8.2 2100 ★★☆
PyTorch 2.0 (Eager) 3分40秒 12.0 1800 ★★★★
PyTorch 2.0 + torch.compile 2分50秒 9.1 1700 ★★★★
PyTorch + v0 (推理) - 7.0 1600 -

结论很明显:

  • 训练场景:PyTorch 2.0 + compile 已经反超TF
  • 推理场景:加上v0这类工具,PyTorch的灵活性优势彻底释放

当然,TF在分布式训练、TPU支持上还是强项。但对我这种单机玩家,PyTorch的“所见即所得”太友好了。


裸辞后的反思:框架之争,本质是工程思维之争

写这篇文章时,我突然想起去年双11前夜。那天我和后端同学为“该用TFServing还是TorchServe”吵到凌晨三点。他说TF的SavedModel格式稳定,我说PyTorch的ONNX转换链路太长……结果第二天上线,因为TF的batching配置错了,导致首单超时,被值班老板call起来骂。

现在回头看,框架只是工具,关键是你能不能驾驭它。PyTorch之所以让我觉得“亲切”,不是因为它多快,而是它把控制权交还给了开发者——你想debug就debug,想graph就graph,想compile就compile。没有魔法,只有显式的API。

就像我现在的状态:裸辞不是逃避,而是主动选择“可控的人生”。通勤路上写代码,咖啡馆里读论文,偶尔研究下v0的源码实现……这种节奏,反而让我对技术有了更深的理解。


给想入坑PyTorch的朋友几点建议

  1. 别死磕文档:PyTorch官方教程写得不错,但最好的学习方式是fork一个开源项目(比如Detectron2、HuggingFace Transformers),然后改它。
  2. 性能优化要有数据支撑:别盲目加torch.compile,先用torch.utils.benchmark测瓶颈。我在v0的issue里看到有人抱怨“compile后变慢”,结果发现是小batch场景,根本没必要。
  3. 善用AI工具,但别依赖:Amazon Q能帮你写样板代码,但模型结构、loss设计、数据增强这些核心逻辑,还得自己想。上周它给我生成了个Focal Loss,结果α参数写反了,害我debug半小时。
  4. 保持对底层的好奇:比如torch.compile背后其实是TorchDynamo + AOTAutograd + Inductor,抽空读读这些子项目的README,你会对PyTorch的架构有全新认识。

最后放个彩蛋:我把完整代码传到了GitHub(搜pytorch-cifar10-baremetal),里面包含了torch.compilev0的集成示例。如果你也在裸辞gap中,或者刚被AI浪潮拍醒,欢迎star/star/star(重要的事说三遍)。

通勤地铁到站了。这次没刷短视频,而是看着终端里跳动的loss曲线——0.23,不错。明天试试加个注意力模块,说不定能在到公司前跑完验证集。

(哦对,我已经不在那家公司了。但习惯改不掉,哈哈。)

评论 0

最热最新
暂无评论
技术AppLv.1
0
影响力
0
文章
0
粉丝