十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

TensorFlow与PyTorch对比:深度学习框架选型指南

TensorFlow与PyTorch对比:深度学习框架选型指南 过去两年深度学习框架的“选型之争”一直是新手入行时绕不开的难题。每次打开技术社区总能看到“TensorFlow 和 PyTorch 到底选哪个”的讨论评论区也经常吵得不可开交。作为一个经历过从 TensorFlow 迁移到 PyTorch又因为项目部署需求重新拾起 TensorFlow 的开发者我深知这种选择困难背后的真实痛点不是某个框架不好而是新手往往不清楚自己的应用场景也没有建立一套选型标准。这篇文章我想把两大框架的差异、适用场景、学习成本这件事讲透。内容包括环境搭建、核心概念拆解、同一模型的代码对比、常见坑点以及选型建议。不管你是刚接触深度学习的学生还是准备在工业界落地 AI 项目的工程师都可以参考这篇文章建立自己的判断逻辑。1. 两大框架到底解决了什么问题1.1 TensorFlow工业界的老牌选手TensorFlow 由 Google Brain 团队于 2015 年开源是目前工业界应用最广泛的深度学习框架之一。它最早采用静态计算图机制先定义完整的计算流程再放入会话Session中执行。这种设计牺牲了一定的灵活性但换来了更好的性能优化空间和部署能力。在 TensorFlow 2.x 版本发布后框架默认开启了动态图模式Eager Execution同时保留了tf.function和 Keras 高层 API让入门门槛大幅降低。目前 TensorFlow 的核心优势集中在生产部署链路TensorFlow Serving、TensorFlow Lite、TensorFlow.js 等工具覆盖了服务端、移动端、浏览器端多种部署场景。需要特别说明的是网上经常出现 TensorFlow 已过时的说法这并不客观。在实际企业中尤其是涉及推荐系统、搜索排序、传统图像识别任务TensorFlow 的存量项目和部署基建仍然非常庞大。学习 TensorFlow 不等于落后它更像是掌握一套工业级工具链。1.2 PyTorch学术界迅速崛起的研究利器PyTorch 由 Facebook AI ResearchFAIR团队于 2016 年开源底层基于 Torch 框架。它的核心特点是采用动态计算图Define-by-Run模型结构在每次前向传播时动态构建这让调试和代码编写更符合 Python 程序员的直觉。过去几年顶级学术会议NeurIPS、CVPR、ICML 等的论文代码绝大多数都优先发布 PyTorch 版本。Hugging Face Transformers 库刚发布时以 PyTorch 为主力实现进一步放大了它在自然语言处理和大模型领域的影响力。如果你关注 AI 前沿研究或者需要快速复现论文、跑通开源项目PyTorch 的学习曲线确实更平缓。同时 PyTorch 也在积极补齐部署短板通过 TorchScript、ONNX 导出、TorchServe 等方式进入生产环境。虽然其部署生态相比 TensorFlow 仍有差距但差距在逐渐缩小。1.3 新手必须掌握的基础概念在直接对比之前先厘清三个核心概念张量、计算图、自动微分。张量Tensor是深度学习的基本数据单位可以理解成多维数组。标量是 0 维张量向量是 1 维张量矩阵是 2 维张量图像数据通常是 3 维或 4 维张量高度、宽度、通道数、批量大小。计算图Computational Graph是框架内部记录运算流程的结构。静态计算图一旦构建就不能修改适合部署优化动态计算图在运行时边执行边构建调试方便、代码可读性高。自动微分Automatic Differentiation是反向传播算法的工程实现。框架自动记录前向传播的每个运算节点反向传播时自动计算梯度开发者不需要手动推导导数公式。理解这三者的区别是理解 TensorFlow 和 PyTorch 差异的基础。尤其是计算图机制直接决定了两个框架的代码风格和调试体验。2. 环境准备与版本说明2.1 安装前的硬件与系统判断无论选择哪个框架安装前都要先确认自己的硬件环境。如果只是入门学习使用 CPU 版本即可跑通 MNIST 等经典数据集但训练速度会明显受限。如果需要训练稍大的模型建议使用支持 CUDA 的 NVIDIA 显卡并安装 GPU 版本。macOS 用户可以关注 M 系列芯片的兼容性问题Windows 用户则需要特别注意 Python 版本与框架版本的匹配关系。总体建议是使用 Anaconda 管理环境避免多个 Python 项目之间发生依赖冲突。2.2 Anaconda 创建虚拟环境Anaconda 是数据科学领域最常用的环境管理工具。创建虚拟环境的主要目的是让 TensorFlow 和 PyTorch 的依赖互不干扰。建议不要直接在 base 环境中安装深度学习框架环境隔离后即使某个环境被装坏也不会影响全局。# 创建 Python 3.10 的虚拟环境 conda create -n dl_study python3.10 # 激活环境 conda activate dl_study这里选择 Python 3.10 是保守做法兼容性较好。如果你需要特定框架版本可以调整为其他版本。创建完成后后续框架安装都发生在这个虚拟环境中出现问题也能一键删除重建。2.3 TensorFlow 安装TensorFlow 的安装命令相对简单PyPI 上默认会安装最新稳定版。文章发布时 TensorFlow 2.18 是较新版本但版本迭代速度快建议安装前访问 PyPI 或官方文档确认。# CPU 版本 pip install tensorflow # GPU 版本需要本机有 NVIDIA 显卡和对应 CUDA 驱动 pip install tensorflowTensorFlow 2.18 起默认安装包已经包含 GPU 支持不再需要单独区分tensorflow-gpu。但实际是否启用 GPU取决于本机 CUDA、cuDNN 版本是否满足要求。如果使用老版本或特殊需求仍可按官方文档单独安装。安装完成后用以下命令验证import tensorflow as tf print(tf.__version__) print(tf.config.list_physical_devices(GPU))tf.config.list_physical_devices(GPU)输出为空说明没有识别到 GPU需要安装 CUDA 工具包和 cuDNN。2.4 PyTorch 安装PyTorch 的安装方式更讲究一点官方会根据你的系统、包管理器、CUDA 版本生成不同的 pip 命令。建议打开 PyTorch 官网首页选择对应配置后复制安装命令避免使用错误的 CUDA 版本。# CPU 版本使用官网生成的命令 pip install torch torchvision torchaudio # GPU 版本示例需根据官方页面选择具体 CUDA 版本 pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121验证 PyTorch 是否可用import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU) # 使用 CUDA 进行张量计算 x torch.tensor([1.0, 2.0]).to(cuda) print(x * 2)需要提醒的是torch.cuda.is_available()返回True只是第一步还要确认实际计算设备是 GPU 而不是 CPU。许多新手在训练时忘记把模型和数据放到 GPU 上导致训练速度毫无提升。2.5 版本差异带来的坑搜索关键词中出现了一个非常有代表性的报错信息(1) in pytorch 2.6, we changed the default value of the weights_only argument这是 PyTorch 2.6 版本后出现的兼容性提示。torch.load()默认weights_only参数从False改为True意味着加载模型权重时不再反序列化任意 Python 对象降低安全风险。但同时也意味着一些老代码在没有手动设置weights_onlyTrue时会出现加载失败。应对方案是在加载模型时明确指定参数checkpoint torch.load(model.pth, map_locationcpu, weights_onlyTrue) # 如果旧模型包含额外状态可设为 False但需确认模型来源可信这类版本变化很难从入门教程中提前得知建议日常关注官方 release note。3. TensorFlow 与 PyTorch 核心差异拆解3.1 静态图 vs 动态图这是 TensorFlow 和 PyTorch 最根本的架构差异。TensorFlow 1.x 时代最明显的特征是“先建图后执行”。你需要先把整个计算流程定义成静态图然后通过tf.Session运行。这种模式有利于分布式训练和部署优化因为图结构在运行前是完整的编译器可以进行整体优化。TensorFlow 2.x 虽然默认开启 Eager Execution动态图但tf.function仍可以将 Python 函数转换为静态图。也就是说TensorFlow 目前同时支持两种模式偏底层推理时仍会用到静态图。PyTorch 从设计之初就采用动态图机制。每次前向传播都实时构建计算图你可以像写普通 Python 一样打印中间变量、断点调试甚至使用if、for控制流。对于研究和快速验证场景这种灵活性极大提升了开发效率。用一个简单比喻理解TensorFlow 像先画好完整的建筑设计图再施工PyTorch 则像边设计边施工随时可以调整墙体和窗户。3.2 API 设计风格TensorFlow 2.x 主推 Keras API用tf.keras.Sequential可以快速堆叠网络层model tf.keras.Sequential([ tf.keras.layers.Dense(128, activationrelu), tf.keras.layers.Dense(10, activationsoftmax) ])PyTorch 更倾向于 Python 原生风格通过继承nn.Module来定义模型class MLP(nn.Module): def __init__(self): super(MLP, self).__init__() self.fc1 nn.Linear(784, 128) self.fc2 nn.Linear(128, 10) def forward(self, x): x torch.relu(self.fc1(x)) x self.fc2(x) return x从代码风格来看PyTorch 更“Pythonic”TensorFlow 更“工程化”。没有绝对优劣取决于你的编程习惯和团队技术栈。3.3 调试体验对比调试是新手学习过程中最重要的部分。PyTorch 因为动态图机制可以直接在forward函数中打断点查看每个中间张量的形状和数值也可以直接使用print输出。TensorFlow 2.x 在 Eager 模式下调试体验有所改善但仍有一些暗坑。比如数据管道tf.data.Dataset在写复杂数据处理逻辑时报错不易定位模型内部张量形状不匹配时的报错信息在部分场景下不够直观。在调试方面PyTorch 有肉眼可见的优势。这也是论文复现和算法开发场景中 PyTorch 更受欢迎的原因之一。3.4 部署生态与生产环境TensorFlow 的部署生态非常完整TensorFlow Serving用于服务端的高性能模型服务。TensorFlow Lite用于移动端和嵌入式设备。TensorFlow.js用于浏览器端。TensorFlow ExtendedTFX用于生产级机器学习流水线。PyTorch 在部署侧也给出了对应方案TorchScript将模型序列化为可部署的脚本支持在 C 环境中运行。TorchServe官方提供的模型服务框架。ONNX 导出能够导出到其他推理引擎如 ONNX Runtime、TensorRT。整体来看如果你所在公司已经有成熟的 TensorFlow 运维体系选择 TensorFlow 在生产接入时更顺畅。但如果你专注算法模型研发PyTorch 模型通过 ONNX 或 TensorRT 也能完成多数部署场景。3.5 社区与学习资源社区生态决定了新手遇到问题后找答案的容易程度。PyTorch 在学术圈占据统治地位CVPR、ICCV、NeurIPS 等顶会论文的开源代码大量使用 PyTorch。GitHub 上很多知名模型仓库如 Hugging Face Transformers、Ultralytics YOLOv5/v8都同时兼容 PyTorch 和 TensorFlow但 PyTorch 版本往往是最先更新、资料最全的。TensorFlow 则在产业界积累了庞大的存量案例很多企业级项目仍然运行在 TensorFlow 栈上。Google 官方提供了大量系统化的学习文档适合动手能力偏弱的初学者跟着走。热门搜索词中同样包含“tensorflow与pytorch的流行趋势 2024年”这说明两大框架的趋势变化是社区持续关注的话题。从当前主流开源社区的活跃度来看PyTorch 在 AI 研究和模型发布端更活跃TensorFlow 在传统工业落地端依然稳固。4. 用同一个模型对比两个框架4.1 对比任务与思路说明为了更直观地比较两个框架的差异这里使用 MNIST 手写数字识别作为统一任务实现一个结构相同的两层卷积神经网络CNN。任务目标包括数据加载、模型定义、训练循环、模型评估四个环节。MNIST 是 28x28 的灰度图像共 10 个类别是深度学习的经典入门数据集。通过同一任务的两个框架实现方式可以清楚看到代码组织风格的差异。4.2 TensorFlow 实现TensorFlow 使用 Keras 高层 API 时训练代码非常简洁。数据加载直接使用内置方法模型定义通过Sequential堆叠网络层训练过程通过compile与fit两步完成。import tensorflow as tf from tensorflow.keras import layers, models # 1. 加载 MNIST 数据集 (x_train, y_train), (x_test, y_test) tf.keras.datasets.mnist.load_data() # 2. 数据预处理归一化 增加通道维度 x_train x_train.astype(float32) / 255.0 x_test x_test.astype(float32) / 255.0 x_train x_train[..., tf.newaxis] # (60000, 28, 28) - (60000, 28, 28, 1) x_test x_test[..., tf.newaxis] # 3. 构建 CNN 模型 model models.Sequential([ layers.Conv2D(32, (3, 3), activationrelu, input_shape(28, 28, 1)), layers.MaxPooling2D((2, 2)), layers.Conv2D(64, (3, 3), activationrelu), layers.MaxPooling2D((2, 2)), layers.Flatten(), layers.Dense(128, activationrelu), layers.Dense(10, activationsoftmax) ]) # 4. 编译模型 model.compile(optimizeradam, losssparse_categorical_crossentropy, metrics[accuracy]) # 5. 训练模型 model.fit(x_train, y_train, epochs3, batch_size64, validation_data(x_test, y_test))这段代码体现了 TensorFlow 的高层封装风格。losssparse_categorical_crossentropy适用于整数标签metrics[accuracy]会在训练过程实时返回准确率fit方法自动完成数据分批、前向传播、反向传播等流程新手甚至不需要了解梯度计算细节。需要说明的是输入数据的形状变化很关键。MNIST 原始数据是 (60000, 28, 28)卷积层需要通道维所以要增加一个维度变为 (60000, 28, 28, 1)。4.3 PyTorch 实现PyTorch 的代码会显式区分模型定义和训练循环。模型通过继承nn.Module实现训练过程需要手动遍历数据加载器并调用损失函数、优化器。import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms # 1. 数据预处理转为 Tensor 并归一化 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) test_dataset datasets.MNIST(root./data, trainFalse, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse) # 2. 定义 CNN 模型 class CNN(nn.Module): def __init__(self): super(CNN, self).__init__() self.conv1 nn.Conv2d(1, 32, 3, 1) self.conv2 nn.Conv2d(32, 64, 3, 1) self.pool nn.MaxPool2d(2, 2) self.fc1 nn.Linear(64 * 5 * 5, 128) self.fc2 nn.Linear(128, 10) def forward(self, x): x self.pool(torch.relu(self.conv1(x))) x self.pool(torch.relu(self.conv2(x))) x torch.flatten(x, 1) x torch.relu(self.fc1(x)) x self.fc2(x) return x model CNN() # 3. 定义损失函数和优化器 criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr0.001) # 4. 训练循环 for epoch in range(3): running_loss 0.0 for images, labels in train_loader: # 清空梯度 optimizer.zero_grad() # 前向传播 outputs model(images) # 计算损失 loss criterion(outputs, labels) # 反向传播 loss.backward() # 更新参数 optimizer.step() running_loss loss.item() print(fEpoch {epoch1}, Loss: {running_loss/len(train_loader):.4f})PyTorch 的nn.Linear(64 * 5 * 5, 128)这一行需要计算卷积池化后的特征图尺寸。输入 28x28经过一次卷积3x3变为 26x26再池化为 13x13第二次卷积后变为 11x11池化后为 5x5所以全连接层输入维度是 64x5x5。训练循环中optimizer.zero_grad()容易被新手遗忘。PyTorch 默认会累积梯度每轮迭代前必须把上一步的梯度清零否则会导致梯度累加错误。4.4 运行结果对比两个框架在相同任务上训练 3 个 epoch最终准确率差异不大通常都在 99% 左右因为 MNIST 本身相对简单模型和轮次足够达到较高精度。关键差异在于代码组织维度TensorFlow (Keras)PyTorch模型定义Sequential 堆叠继承 nn.Module 自定义训练流程compile fit自动完成手动循环逐步执行数据加载tf.keras.datasets 内置torchvision.datasets调试灵活度适合成体系训练流程适合自定义训练逻辑代码控制感低适合快速上手高适合研究场景初学者如果习惯“一键训练”TensorFlow 的 Keras 模式更容易上手如果希望理解训练底层逻辑PyTorch 的显式循环反而能帮助建立完整认知。4.5 训练后的模型保存与加载模型保存也是实际开发中的高频操作。两个框架的 API 差异贯穿了整个使用链路。TensorFlow 推荐使用 SavedModel 格式它同时保存模型结构和权重方便后续部署到 TensorFlow Serving。# 保存模型 model.save(mnist_model.keras) # 加载模型 loaded_model tf.keras.models.load_model(mnist_model.keras) # 推理 predictions loaded_model.predict(x_test[:10])PyTorch 通常只保存参数字典state_dict因为模型结构定义在 Python 代码中。推荐加载方式如下# 保存模型参数 torch.save(model.state_dict(), mnist_model.pth) # 加载模型参数需要重新实例化模型 model CNN() model.load_state_dict(torch.load(mnist_model.pth, weights_onlyTrue)) model.eval()PyTorch 的model.eval()很重要因为模型训练和推理时某些层的行为如 Dropout、BatchNorm不同。切换为 eval 模式可以确保推理结果正确。5. 常见问题与排查思路5.1 框架安装类问题问题现象常见原因解决思路TensorFlow 导入报错 DLL load failed缺少 VC 运行库或 CUDA 版本不匹配Windows 安装 Visual C Redistributable确认 CUDA/cuDNN 版本PyTorch 安装时下载速度慢默认源来自国外服务器使用清华源或阿里源加速但需检查命令格式Anaconda 创建环境后 pip 仍指向全局未激活环境或 PATH 配置错误执行conda activate后检查which python同一项目 TensorFlow 和 PyTorch 冲突依赖版本互相覆盖使用独立 conda 环境分别安装5.2 CUDA 与 GPU 加速问题GPU 安装是新手最容易卡住的地方。核心原则是先确认本机驱动力支持的 CUDA 版本再选择对应框架版本不能倒过来装。推荐方式是通过 NVIDIA 官方工具或命令行确认驱动版本再对照支持的 CUDA 版本表。如果你用的是 PyTorch直接访问 PyTorch 官网首页生成安装命令即可官网上每一步都有交互选项基本不会出错。需要注意的是TensorFlow 对 CUDA 版本的要求更严格官方文档会列出不同 TensorFlow 版本对应的 CUDA 和 cuDNN 版本。安装前务必核对这些表格。5.3 与版本相关的兼容性坑前面提到的 PyTorchweights_only默认值变化就是版本升级带来的典型问题。这类问题在新框架版本发布后尤其常见解决思路主要有两个升级代码适配新 API 改名例如weights_onlyTrue显式传参。锁定项目依赖版本在requirements.txt中固定版本号保证可复现。在其他框架项目中比如搜索热词中出现的若依框架、pytest 框架也存在类似的版本兼容问题。框架类工具的升级都要谨慎生产环境升级前必须在测试环境提前验证。5.4 训练过程常见错误错误现象可能原因解决方法Loss 为 NaN学习率过大、数据未归一化调低学习率、检查输入数据训练速度很慢数据没有放到 GPU 上检查tensor.to(cuda)和model.to(cuda)模型不收敛标签类别错误、损失函数选错确认分类任务使用 CrossEntropy 还是 MSE显存不足 OOM批大小过大调低 batch_size 或使用梯度累积6. 新手到底该选哪个6.1 先想清楚自己的目标场景“TensorFlow 和 PyTorch 哪个好”本质上是个伪命题。选型应该取决于你的实际目标而非他人评价。这里给出三个典型场景场景一入门深度学习、跑通经典模型、理解神经网络原理。推荐 PyTorch。动态图调试友好代码风格接近 Python 原生习惯遇到问题更容易定位。场景二进入互联网公司做算法工程师参与推荐、搜索、广告等场景。需要了解目标公司技术栈。搜索和广告方向 TensorFlow 存量设施更多推荐系统方向 PyTorch 后来居上。场景三希望在移动端、嵌入式设备部署模型。TensorFlow Lite 生态更成熟TensorFlow 是较好的选择。PyTorch 的 TorchScript 也能做但工程化工具链相对薄弱。6.2 不同岗位的选型建议身份/岗位建议优先学习理由学生、科研人员PyTorch论文复现方便社区最新资源多传统企业后端工程师TensorFlow部署工具链成熟Keras API 上手快算法工程师搜索/推荐方向取决于公司栈建议两者都懂存量系统与前沿模型可能并存移动端开发工程师TensorFlow Lite端侧部署资源和案例更丰富6.3 如何避免“反复横跳”我的经验是选定一个框架先深入学透不要频繁切换。很多新手学了两周 TensorFlow看到 PyTorch 火就换 PyTorch换来换去结果两边都没学扎实基础概念也没理清。深度学习框架的学习价值并不绑定在某个具体 API 上。张量运算、自动求导、反向传播、卷积、循环神经网络这些核心知识在两个框架中都是相通的。真正让你成为优秀工程师的不是你会哪个框架而是你能快速理解框架设计思想在需要的时候迁移到另一套生态。等到具备一定基础后建议再花时间了解第二个框架。两个框架都掌握后你就能在技术选型时更理性地做决策而不是听别人说哪个好就用哪个。7. 最佳实践与工程建议7.1 代码层面的工程规范注释和命名规范在深度学习中容易被忽视但同样重要。模型类的命名应体现网络结构功能如ResNetEncoder、TransformerDecoder训练配置文件与代码分离超参数不要散落在各个函数中。建议创建一个config.py统一管理学习率、批大小、训练轮次、数据路径等参数。数据处理是另一个容易踩坑的点。TensorFlow 的tf.data.Dataset和 PyTorch 的Dataset/DataLoader都支持复杂的数据预处理流水线但两者差异较大。建议尽早掌握所在框架的数据加载最佳实践避免每一步都用for循环手动处理。7.2 GPU 训练的可复现性深度学习训练涉及随机性对可复现性要求高的场景需要设置随机种子# PyTorch 设置随机种子 import random import numpy as np import torch def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)TensorFlow 对应写法import tensorflow as tf import numpy as np def set_seed(seed42): np.random.seed(seed) tf.random.set_seed(seed)需要注意的是即使设置了随机种子GPU 并行计算仍可能出现微小差异。在正式实验对比中建议多次重复实验取均值。7.3 生产环境部署注意事项生产环境的约束和学习环境完全不同。模型推理不仅要关注精度还要关注延迟、吞吐量、显存占用、稳定性。如果使用 TensorFlow建议将模型导出为 SavedModel 格式使用 TensorFlow Serving 进行服务化部署。Triton Inference Server 也是当前企业中常用的部署方案它同时支持 TensorFlow、PyTorch 和 ONNX Runtime。如果使用 PyTorch建议先用torch.jit.trace或torch.jit.script将模型转换为 TorchScript再通过 LibTorchPyTorch C API进行高性能部署。ONNX 导出是另一种常用方式。注意动态模型包含数据依赖的 if 分支或循环在 trace 时可能出错改用 script 方式更稳定。7.4 安全与权限意识在服务器上安装驱动、更新 CUDA、修改系统环境变量时需要谨慎操作。生产服务器尤其建议先在测试环境验证。权限方面始终使用最小权限原则避免用 root 账号执行不明确的安装脚本。加载他人模型权重文件时要留意安全风险。PyTorch 的weights_onlyTrue设计就是为了防止恶意代码在反序列化时执行。即使你的框架版本还不需要添加这个参数也建议养成显式传参的习惯。8. 总结与学习路线规划写到这里全文的核心结论已经明确了TensorFlow 强在工业部署、工具链完整、Keras API 入门友好。PyTorch 强在动态图调试、学术生态繁荣、模型创新速度最快。两者都支持 GPU 训练、模型导出、生产部署选型更多取决于你的应用场景。如果你还是刚入门的新手我建议先选择 PyTorch 作为主攻方向因为它的代码更直观能让你把注意力集中在神经网络原理本身。学完基础后再回头了解 TensorFlow 的 Keras 模式你会发现很多概念是互通的。如果你是为了入职传统企业做运维或工程平台开发TensorFlow 则更贴近已有技术体系项目落地更顺畅。动手实践是唯一的捷径。建议先去跑通 MNIST 和 CIFAR-10 两个经典数据集用两个框架分别实现一遍。当你亲手写好第一个 CNN 模型并且在 GPU 上完成训练后你就不会再纠结框架的选择问题了。以后再看新的框架不管叫 JAX 还是 MindSpore你都能以同样的方式快速上手。
返回列表