尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyTorch与TensorFlow核心对比:从环境搭建到Transformer实战

PyTorch与TensorFlow核心对比:从环境搭建到Transformer实战 在实际的深度学习工程中PyTorch 和 TensorFlow 不是一道“二选一”的选择题而更像是一套工具箱里常用的两把扳手。学术界论文复现更多看到 PyTorch工业界 TensorFlow 的部署链路依然成熟而很多团队其实是两套框架并存研究阶段用 PyTorch上线阶段转 TensorFlow、ONNX 或者 TensorRT。这篇文章会从框架定位、环境搭建、手写数字识别实战、Transformer 注意力实现、常见安装与运行问题排查五个方面把两条技术路线同时过一遍。读完以后你可以根据需求在不同框架间切换而不是被困在“哪个更好”的争论里。需要提醒的是深度学习框架版本迭代非常快。下面示例中会出现 Python 3.10、CUDA 12.x、PyTorch 2.x、TensorFlow 2.18 等常见组合落地时请先确认你的显卡驱动、CUDA 版本和框架版本互相兼容不要机械照搬命令。1. 先理解 PyTorch 与 TensorFlow 的定位差异1.1 从动态图和静态图说起PyTorch 走的是“动态计算图”路线。你在 Python 里写一行代码计算图就同步构建一层tensor上的操作和 NumPy 很接近遇到if、for也可以直接用 Python 语法控制。这种模式非常像写普通程序调试时可以用print打印中间张量也可以用breakpoint()直接打断点所以科研人员和算法工程师上手速度很快。TensorFlow 早期使用静态计算图需要先定义完整计算图再在Session里执行。2019 年 TensorFlow 2.0 以后默认开启 Eager Execution也就是即时执行模式写起来也接近动态图。但要真正发挥 TensorFlow 在生产环境中的性能还是需要理解tf.function如何把 Python 函数编译成图以及AutoGraph会把哪些 Python 语法自动转换。这是两套框架最本质的思维差异PyTorch 是“写 Python 就是写模型”TensorFlow 是“先写 Python再考虑如何编译和优化”。1.2 生态与部署侧重点不同PyTorch 的强项在研究生态。Hugging Face Transformers、Diffusers、Ultralytics YOLO 等主流模型库默认基于 PyTorch。论文复现、Kaggle 竞赛、开源项目里PyTorch 代码占比非常高。如果你要快速验证一个新的模型结构PyTorch 的灵活性和社区样例更占优势。TensorFlow 的强项在工程化部署。TensorFlow 提供了从TFRecord数据格式、tf.data数据管道、TFX流水线到TensorFlow Serving、TensorFlow Lite、TensorFlow.js的完整链路。移动端和嵌入式端经常用 TFLite服务端有专门的 Serving 框架跨语言调用也比较方便。如果你需要把模型上线到没有 GPU 的生产容器或手机端TensorFlow 的成熟工具链仍然值得考虑。1.3 用一张表快速对照对比维度PyTorchTensorFlow计算图动态图为主也可以用torch.compile或torch.jit做图优化2.x 默认 Eager Mode通过tf.function编译成静态图编程风格更接近 NumPy对象导向调试友好高层 Keras API 简洁底层tf.tensor操作繁多典型用户学术界、算法研究员、竞赛玩家工业界、部署团队、已有 Google 生态团队模型库Hugging Face、Ultralytics、OpenMMLab 等Keras 官方模型库、TF-Hub部署工具TorchScript、ONNX、TorchServeTensorFlow Serving、TFLite、TF.js学习曲线线性容易从零开始Keras 很简单深入tf.function后曲线变陡长期维护Meta 主导社区活跃Google 主导更新节奏稳定不要因为某个框架在某一年“更流行”就否定另一个。实际选型要看团队已有代码、部署目标、模型来源和硬件环境。最稳妥的做法是两套都掌握基本流程再根据项目选择主用框架。2. 环境准备用 Anaconda 隔离 PyTorch 和 TensorFlow很多安装问题都出在“同一个 Python 环境里装了彼此冲突的依赖”。PyTorch 和 TensorFlow 对 CUDA、cuDNN、numpy 等依赖版本要求不同混装容易导致 import 报错。推荐使用 Anaconda 或者 Miniconda 创建独立的虚拟环境。2.1 安装 Anaconda 并检查显卡驱动前往 Anaconda 官网下载对应系统的安装包安装完成后在命令行执行conda --version python --version如果你要安装 GPU 版本先用命令检查显卡驱动和硬件nvidia-smi输出里会显示驱动版本、支持的 CUDA 版本。注意nvidia-smi显示的 CUDA 版本是驱动支持的最高版本不一定是当前环境里实际安装的 CUDA 工具包版本。PyTorch 和 TensorFlow 通过编译期绑定的 CUDA 运行库工作不要求你手动安装完整 CUDA 工具包但驱动必须足够新。2.2 创建 PyTorch 专属环境打开命令行或 Anaconda Prompt创建 Python 3.10 虚拟环境conda create -n pytorch-env python3.10 -y conda activate pytorch-env安装 PyTorch 时建议先到 PyTorch 官网获取当前最稳妥的安装命令。官网会根据你选择的系统、包管理器、CUDA 版本生成类似命令# CPU 版本 pip install torch torchvision torchaudio # CUDA 12.1 版本示例最新版本以官网为准 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121在某些网络环境下默认 PyPI 源下载很慢可以临时切换清华源但仍建议 PyTorch 官方源优先因为官方源的 wheel 和 CUDA 库匹配更完整pip install torch torchvision torchaudio -i https://pypi.tuna.tsinghua.edu.cn/simple2.3 创建 TensorFlow 专属环境同样创建一个独立环境conda create -n tf-env python3.10 -y conda activate tf-env安装 TensorFlow CPU 或 GPU 版本pip install tensorflowTensorFlow 2.18 是较早时期的一个版本。后续版本安装方式类似核心要求是 Python 版本和 pip 版本兼容。如果你需要 GPU 支持直接安装tensorflow通常已经包含 GPU 支持不需要单独安装tensorflow-gpu。这一点和 1.x 时代不同不要在 2.x 环境里继续安装tensorflow-gpu否则会收到“请卸载并安装 tensorflow”的提示。安装后验证python -c import tensorflow as tf; print(tf.__version__) python -c import torch; print(torch.__version__)2.4 检查 GPU 是否可用PyTorch 检查 GPUpython -c import torch; print(torch.cuda.is_available())如果输出True再打印设备名python -c import torch; print(torch.cuda.get_device_name(0))TensorFlow 检查 GPUpython -c import tensorflow as tf; print(tf.config.list_physical_devices(GPU))正常时会输出类似[PhysicalDevice(name/physical_device:GPU:0, device_typeGPU)]如果结果是空列表或者False不要急着重装整个框架。先检查驱动版本、CUDA 版本、环境变量常见排查方法放在第 5 节。注意CPU 版本同样可以学习模型结构和训练流程只是速度慢很多。如果电脑没有 NVIDIA 显卡暂时不要强行安装 GPU 版先跑熟 CPU 流程后面有 GPU 环境时再重建环境。3. 用同一个手写数字识别任务对比两套框架的建模流程这里选择 MNIST 手写数字识别作为对比任务。原因是数据小、模型结构简单、训练时间短而且能完整展示“数据加载、模型定义、训练、验证、保存”这一套核心流程。3.1 任务定义和预期结果输入是一张 28x28 的灰度图片输出是 0 到 9 共 10 个类别的概率。模型使用两层卷积加全连接层训练 3 个 epoch 后测试准确率可以达到 98% 以上。为了保持对比公平两套框架使用相同的网络结构、相同的随机种子和相同的优化器参数。3.2 PyTorch 版本的完整实现import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms # 固定随机种子便于复现 torch.manual_seed(42) # 数据预处理转 Tensor 并归一化 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_data datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) test_data datasets.MNIST( root./data, trainFalse, downloadTrue, transformtransform ) train_loader DataLoader(train_data, batch_size64, shuffleTrue) test_loader DataLoader(test_data, batch_size1000, shuffleFalse) class CNN(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 32, kernel_size3, padding1) self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(64 * 7 * 7, 128) self.fc2 nn.Linear(128, 10) self.relu nn.ReLU() def forward(self, x): x self.pool(self.relu(self.conv1(x))) x self.pool(self.relu(self.conv2(x))) x x.view(x.size(0), -1) x self.relu(self.fc1(x)) x self.fc2(x) return x device torch.device(cuda if torch.cuda.is_available() else cpu) model CNN().to(device) criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) def train_one_epoch(): model.train() total_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(train_loader) def evaluate(): model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in test_loader: images, labels images.to(device), labels.to(device) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() return correct / total for epoch in range(3): loss train_one_epoch() acc evaluate() print(fEpoch {epoch 1}, Loss: {loss:.4f}, Accuracy: {acc:.4f})这段代码展示了 PyTorch 训练循环的几个典型特征模型类继承nn.Module必须实现forward方法。手动调用optimizer.zero_grad()、loss.backward()、optimizer.step()训练三步缺一不可。使用with torch.no_grad()关闭梯度计算加快验证速度。张量通过.to(device)在 CPU 和 GPU 之间迁移。3.3 TensorFlowKeras版本的完整实现TensorFlow 2.x 的推荐写法是使用 Keras 高层 API数据和模型定义都更简洁import tensorflow as tf tf.random.set_seed(42) # 加载 MNIST 并归一化 (x_train, y_train), (x_test, y_test) tf.keras.datasets.mnist.load_data() x_train x_train.astype(float32) / 255.0 x_test x_test.astype(float32) / 255.0 # 增加通道维从 (28, 28) 变为 (28, 28, 1) x_train x_train[..., tf.newaxis] x_test x_test[..., tf.newaxis] # 标签用整数模型输出 logits损失函数计算时自动处理 one-hot model tf.keras.Sequential([ tf.keras.layers.Conv2D(32, kernel_size3, paddingsame, activationrelu), tf.keras.layers.MaxPooling2D(2, 2), tf.keras.layers.Conv2D(64, kernel_size3, paddingsame, activationrelu), tf.keras.layers.MaxPooling2D(2, 2), tf.keras.layers.Flatten(), tf.keras.layers.Dense(128, activationrelu), tf.keras.layers.Dense(10, activationsoftmax) ]) model.compile( optimizertf.keras.optimizers.Adam(0.001), losstf.keras.losses.SparseCategoricalCrossentropy(), metrics[accuracy] ) history model.fit( x_train, y_train, batch_size64, epochs3, validation_data(x_test, y_test) )Keras 的model.fit把数据加载、反向传播、参数更新、验证都封装好了代码量比 PyTorch 少很多。如果你更想理解训练细节也可以使用自定义训练循环train_ds tf.data.Dataset.from_tensor_slices((x_train, y_train)).batch(64) test_ds tf.data.Dataset.from_tensor_slices((x_test, y_test)).batch(1000) loss_fn tf.keras.losses.SparseCategoricalCrossentropy() optimizer tf.keras.optimizers.Adam(0.001) train_acc tf.keras.metrics.SparseCategoricalAccuracy() tf.function def train_step(images, labels): with tf.GradientTape() as tape: logits model(images, trainingTrue) loss loss_fn(labels, logits) grads tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(grads, model.trainable_variables)) train_acc.update_state(labels, logits) return loss for epoch in range(3): for images, labels in train_ds: loss train_step(images, labels) print(fEpoch {epoch 1}, Loss: {loss:.4f}, Accuracy: {train_acc.result():.4f}) train_acc.reset_states()这里tf.GradientTape()就是 TensorFlow 版本的自动求导上下文管理器。PyTorch 的自动求导是张量构建时自动记录的TensorFlow 则需要显式在GradientTape中记录前向计算过程这是两套框架在梯度实现上的显著差别。3.4 两套框架的核心流程对照环节PyTorchTensorFlow(Keras)数据组织torch.utils.data.DatasetDataLoadertf.data.Dataset或 NumPy 数组模型定义nn.Module子类手写forwardtf.keras.Sequential或函数式 API自动求导张量自动记录loss.backward()tf.GradientTape内记录tape.gradient参数更新手动调用optimizer.step()高层model.fit或手动apply_gradients训练模式model.train()/model.eval()参数trainingTrue/False验证阶段torch.no_grad()tf.stop_gradient或依赖trainingFalse模型保存torch.save(model.state_dict(), model.pt)model.save(model.h5)或.keras不要小看这些差异。它们会影响你阅读开源代码时的理解速度。看到model.train()要知道这是切换 dropout、batch norm 等层的行为看到GradientTape要知道这是 TensorFlow 记录前向计算、准备反向传播的入口。4. 从单机训练到 Transformer用 PyTorch 实现一个通用注意力模块Transformer 已经成为自然语言处理和视觉任务的基础结构。在你掌握两套框架的建模流程后值得用 PyTorch 细看注意力机制的实现。这不是因为 TensorFlow 不能实现 Transformer而是 PyTorch 的代码在论文复现中更常见逐行阅读更容易理解张量维度变化。4.1 自注意力的核心计算自注意力的输入是一个序列形状为(batch_size, seq_len, d_model)。每个位置的 token 都有三个向量Query、Key、Value。注意力权重通过 Query 和 Key 的点积计算再经过 softmax 归一化最后和 Value 加权求和。用公式表示Attention(Q, K, V) softmax(Q * K^T / sqrt(d_k)) * V除以sqrt(d_k)是为了防止点积数值过大梯度进入 softmax 饱和区。4.2 PyTorch 实现一个 decoder 用的通用注意力模块热搜词里出现了“a generic attention module for a decoder in seq2seq pytorch”这里给出一个简洁版本import torch import torch.nn as nn import torch.nn.functional as F class Attention(nn.Module): def __init__(self, hidden_size): super().__init__() self.hidden_size hidden_size self.query_proj nn.Linear(hidden_size, hidden_size) self.key_proj nn.Linear(hidden_size, hidden_size) self.value_proj nn.Linear(hidden_size, hidden_size) def forward(self, query, keys, values, maskNone): query: (batch_size, query_len, hidden_size) keys: (batch_size, key_len, hidden_size) values:(batch_size, key_len, hidden_size) mask: (batch_size, key_len)可选用于屏蔽 padding q self.query_proj(query) k self.key_proj(keys) v self.value_proj(values) # 注意力分数: (batch_size, query_len, key_len) scores torch.matmul(q, k.transpose(-2, -1)) / torch.sqrt( torch.tensor(self.hidden_size, dtypetorch.float32) ) if mask is not None: # mask1 表示有效位置将无效位置填充为负无穷 scores scores.masked_fill(mask.unsqueeze(1) 0, float(-inf)) weights F.softmax(scores, dim-1) context torch.matmul(weights, v) return context, weights这段代码关键点transpose(-2, -1)交换最后两个维度得到(batch_size, query_len, key_len)的二维得分矩阵。masked_fill在 padding 位置填充-infsoftmax 后概率趋近于 0避免模型关注无效位置。返回同时包含上下文向量和注意力权重方便可视化。4.3 TensorFlow 实现同样模块的方式在 TensorFlow 中可以用 Keras Layer 实现同样结构import tensorflow as tf class AttentionLayer(tf.keras.layers.Layer): def __init__(self, hidden_size): super().__init__() self.hidden_size hidden_size self.query_proj tf.keras.layers.Dense(hidden_size) self.key_proj tf.keras.layers.Dense(hidden_size) self.value_proj tf.keras.layers.Dense(hidden_size) def call(self, query, keys, values, maskNone): q self.query_proj(query) k self.key_proj(keys) v self.value_proj(values) scores tf.matmul(q, k, transpose_bTrue) / tf.sqrt( tf.cast(self.hidden_size, tf.float32) ) if mask is not None: scores -1e9 * (1.0 - tf.cast(mask, tf.float32))[:, tf.newaxis, :] weights tf.nn.softmax(scores, axis-1) context tf.matmul(weights, v) return context, weightsTensorFlow 版的关键写法是transpose_bTrue表示在matmul中自动转置第二个矩阵mask 处理方式也有差别这里给无效位置加上一个很大的负数而不是替换成-inf。逻辑等价实现习惯不同。从这段对比可以看出两个框架在张量维度变化、掩码处理、softmax 使用上几乎一一对应。熟练一个框架后再学另一个主要是在记忆 API 映射而不是重新学一遍数学原理。5. 安装与运行中的高频问题排查无论你选择哪个框架都会遇到环境问题。下面根据常见场景总结一套可执行的排查路径。5.1torch.cuda.is_available()返回 False先确认显卡驱动和 CUDA 支持nvidia-smi如果驱动看不到任何 GPU先解决驱动问题。如果驱动正常再检查 PyTorch 版本是否匹配python -c import torch; print(torch.version.cuda)如果输出的是cpu或者与驱动支持版本差距过大说明 PyTorch 安装的是 CPU 版本或错误 CUDA 版本。建议重新创建环境从 PyTorch 官网复制正确的 GPU 安装命令。常见错误是安装时没有指定--index-url导致 pip 安装了默认 PyPI 的 CPU 版本。5.2 TensorFlow 提示Could not load dynamic library cudnn64_8.dll这是 TensorFlow 找不到 cuDNN 库的典型报错。在 Windows 上可以先看环境变量PATH中是否包含 CUDA 和 cuDNN 的bin目录。更省事的做法是使用 conda 安装 CUDA、cuDNN 依赖conda install -c conda-forge cudatoolkit cudnnTensorFlow 2.x 和 PyTorch 对 CUDA 版本敏感不要混合安装。推荐使用 conda 自带的cudatoolkit避免手动下载安装后路径不匹配。5.3 PyTorch 2.6 中weights_only参数默认值变化PyTorch 2.6 以后torch.load的weights_only参数默认值改为了True。这个参数变化影响的是加载模型时的反序列化安全性。旧代码state_dict torch.load(model.pth)改为显式指定# 如果你只加载状态字典推荐这种写法 state_dict torch.load(model.pth, weights_onlyTrue) # 如果你确实加载包含自定义对象的完整 checkpoint再关闭 full_checkpoint torch.load(model.pth, weights_onlyFalse)这条规则的意义在于torch.load本质上是反序列化 Python 对象不受信任的文件可能执行恶意代码。默认开启weights_onlyTrue可以降低安全风险。对于常规模型保存推荐只保存state_dict而不要直接保存整个模型对象这样迁移部署更灵活。5.4 Jetson 等嵌入式平台安装版本选择Jetson 设备常用 JetPack 版本控制整个系统环境。例如 JetPack 6.2.2 对应的 L4T 版本、CUDA、cuDNN 都是固定组合不能直接使用桌面版的 PyTorch wheel。官方社区会发布对应的 PyTorch wheel 包安装前需要确认torch.__version__和torch.version.cuda是否和 JetPack 的 CUDA 版本一致。安装失败时优先检查cat /etc/nv_tegra_release python -c import torch; print(torch.__version__)嵌入式平台的 pip 包名和桌面环境不同不要从 PyTorch 官网复制cu121这类命令直接安装。建议从 NVIDIA 官方文档或 PyTorch 官方论坛的 “Jetson” 专区找对应版本。5.5 常见安装报错速查表问题现象常见原因检查方式处理建议pip 下载很慢或超时默认 PyPI 源在大陆连接不稳定查看 pip 源配置使用清华镜像或阿里镜像但 PyTorch GPU 版优先官方源导入 PyTorch 时提示 DLL 加载失败缺少 MSVC 运行库或依赖库路径问题查看完整报错栈末尾安装 VC redist、确认 conda 环境干净import tensorflow 时提示 protobuf 冲突其他包引入了不同版本 protobufpip list | grep protobuf重新安装pip install protobuf3.20.3或升级到兼容版本Keras 模型训练中显存不足batch size 过大或输入尺寸过大nvidia-smi查看显存占用降低 batch size或使用混合精度同一个环境同时装 torch 和 tensorflow 后 import 报错依赖冲突如 cuDNN、numpy 版本不一致分别在独享环境里验证使用独立虚拟环境不要混装torch.load报错提示 weights_only新版 PyTorch 默认值变化检查报错中的路径显式传weights_onlyTrue或只保存 state_dict秋叶启动器等第三方工具 PyTorch 安装失败工具内置的安装逻辑和自定义环境冲突查看工具日志中的具体 pip 命令不要重复手动安装在独立环境里装好后再作为外部环境接入排查原则是先确认硬件驱动再确认 Python 版本再确认 CUDA 版本最后看具体报错位置。不要一看到安装失败就卸载重装先记录完整报错信息。6. 两套框架都要学如何安排学习路径6.1 从项目需求反推主次顺序如果你是学生或研究员需要快速跑通论文代码推荐先学 PyTorch。因为 Hugging Face 生态、最新模型的transformers库、Diffusers 等默认接口都是 PyTorch。先把 PyTorch 的Dataset、DataLoader、nn.Module、训练循环、模型保存学透后面看任何开源项目都能快速定位关键代码。如果你是在工业界做模型上线或者团队已有 TensorFlow 服务先从 Keras 入手会更快。Keras 的model.fit封装度高数据集规范、模型导出、服务化流程都有官方指南。后续需要自定义训练逻辑时再补GradientTape和tf.function。学习路线建议用 NumPy 理解张量、矩阵乘法和梯度下降。从一个框架完成 MNIST 和 CIFAR-10 分类。掌握数据增强、模型保存、加载和推理。换另一个框架复现同样的模型和指标。找一个开源项目把模型结构画出来。尝试将 PyTorch 模型导出为 ONNX再用 TensorFlow 加载或部署打通互操作链路。6.2 基于同一套数学原理迁移知识不要把一个框架当作独立课程来学。两套框架共享的概念非常多损失函数、优化器、卷积核、池化、BatchNorm、Dropout、学习率调度、早停、模型检查点。你在 PyTorch 里理解了CrossEntropyLoss的内部逻辑在 TensorFlow 里看SparseCategoricalCrossentropy就只是 API 名称不同。练习迁移时可以做一个“框架对照笔记”遇到新概念就记两条概念PyTorch APITensorFlow API张量创建torch.tensor(...)tf.constant(...)卷积层nn.Conv2d(...)tf.keras.layers.Conv2D(...)丢弃层nn.Dropout(p0.5)tf.keras.layers.Dropout(0.5)BatchNormnn.BatchNorm2d(...)tf.keras.layers.BatchNormalization()学习率衰减torch.optim.lr_scheduler.StepLRtf.keras.optimizers.schedules.ExponentialDecay这种“找对照”的方式比从头再学一遍高效得多。6.3 生产环境中的常见分工一个模型从研究到落地往往会经历两个框架的转换研究阶段PyTorch 快速设计新模型、做实验、对比论文结果。转换阶段将 PyTorch 权重导出为 ONNX或重新实现 TensorFlow 版本。部署阶段TensorFlow Serving、TFLite 做服务化或端侧推理PyTorch 也提供 TorchServe但生态相对更偏向研究。是否值得花时间同时维护两套代码取决于团队规模和项目周期。如果只是个人学习至少做到能读懂两种代码能完成环境搭建和推理。如果是团队协作建议统一主框架避免模型转换过程引入精度差异。6.4 环境与工程最佳实践清单下面的清单可以直接复制到你的项目里作为检查项使用 conda 为每个框架创建独立环境环境名区分清楚比如pytorch-env、tf-env。尽量使用 Python 3.9 到 3.11 的版本除非框架官方明确支持更高版本。安装 GPU 版本前先执行nvidia-smi确认驱动版本不要太旧。PyTorch 的 GPU wheel 来自download.pytorch.org不要从普通 PyPI 下载带 CUDA 的版本。TensorFlow 2.x 不需要单独安装tensorflow-gpu直接安装tensorflow即可。不使用torch.save(model, ...)保存整个模型优先保存state_dict。加载模型时使用map_location或显式传输设备state_dict torch.load(model.pt, map_locationcpu) model.load_state_dict(state_dict)训练脚本固定随机种子至少在启动时设置import random import numpy as np import torch random.seed(42) np.random.seed(42) torch.manual_seed(42) if torch.cuda.is_available(): torch.cuda.manual_seed_all(42)频繁保存 checkpoint 时要清理旧文件避免磁盘写满导致训练中断。6.5 最终该怎么选到这里可以给出一套相对稳妥的决策原则看开源生态新论文、Hugging Face、Diffusers 直接抄 PyTorch。看部署链路TensorFlow Serving、TFLite、移动端兼容性更成熟。看团队惯例团队已有代码用什么优先沿用不要为了“个人偏好”重建。看学习目标只要入门选 PyTorch 更顺要做全工程链路补学 TensorFlow 是有价值的。大多数情况下先深度掌握 PyTorch再按需学习 TensorFlow是性价比最高的路径。两套框架之间不是零和竞争它们解决的是同一类问题只是在不同环节各有优势。理解这一点后你真正需要关注的就不再是“哪个框架封神”而是你的模型、数据和业务目标适合哪条技术链路。把两套框架的基础能力都掌握牢固才能在面对不同项目时快速切换到最合适的工具。
返回列表