ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

TensorFlow张量操作性能优化实战:从底层原理到XLA与混合精度

TensorFlow张量操作性能优化实战:从底层原理到XLA与混合精度 开头有人问我搞了这么多年TensorFlow最大的感触是什么。我的答案不是深度学习真神奇而是张量操作决定天花板——模型结构再漂亮数据管道和计算过程中的张量操作写得很烂训练速度能差出好几倍甚至直接决定你能不能跑起来。这个结论不是拍脑袋。去年帮团队优化一个序列模型同样的硬件、同样的模型结构仅仅是把数据预处理和模型内部的一些张量操作从能跑就行改成符合TensorFlow设计习惯训练吞吐直接翻了接近3倍。那次之后我意识到很多人对张量的理解停留在多维数组这个层面对TensorFlow内部如何处理张量、如何调度算子、如何与GPU驱动和CUDA库协作几乎一无所知。这篇文章打算把我在实际项目中积累的张量操作经验完整梳理一遍。从张量的底层本质说起覆盖创建、变形、广播、索引这些基础操作的隐含开销再进入性能优化层面讲清楚算子融合、XLA编译、混合精度、显存管理的真实工作方式。你有一定的TensorFlow基础最好完全零基础也问题不大前面的部分足够细致。重点是读完以后你应该能看懂别人代码里张量操作为什么快、为什么慢并且能自己下手优化。1. 张量的真实面目别把它当成多维数组就完事了1.1 从NumPy数组到TensorFlow张量底层存储与调度逻辑的差异我见过太多从NumPy转过来的开发者习惯性认为tf.Tensor就是np.ndarray换了个名字。这个类比在直觉层面成立但到了性能层面会害死人。NumPy数组的数据确实存在一块连续内存里但每个操作都是立即执行的结果是同步返回的。TensorFlow则完全不是这样——tf.Tensor更像是一个计算图节点的句柄。当你写下c a b在Eager模式下它确实会立即计算但这个计算最终会被dispatch到具体的设备CPU或GPU上执行而且TensorFlow有自己的内存分配器、算子调度器、甚至可以在后台把相邻的Eager操作重新编排成更高效的执行序列。底层存储上TensorFlow张量默认使用行优先row-major内存布局这点和NumPy一样。但区别在于TensorFlow张量对设备有强绑定一个在GPU上的张量你访问它的numpy()方法时TensorFlow必须先把数据从显存拷贝回内存这个操作在性能敏感路径上要尽量避免。还有tf.Tensor是只读的你不能像操作NumPy数组那样原地给某个元素赋值。如果你需要做原地更新得用tf.Variable而tf.Variable在分布式训练和多设备场景下的读写同步策略又比NumPy复杂得多。1.2 shape、dtype、device三要素性能优化的起点我看过太多跑不动就先加GPU的案例但实际上有相当一部分性能问题出在张量的三个基本属性没设计好逻辑形状shape、数据类型dtype、设备位置device。先说shape。TensorFlow的shape有三种状态完全已知的静态shape、部分已知比如batch维为None、完全未知的动态shape。性能差异非常明显——静态shape已知的张量TensorFlow在编译和执行时可以提前分配内存、确定算子kernel、甚至做算子融合动态shape则意味着每次执行都可能触发重新编译或者内存重分配。我实测过一个简单的tf.matmul在静态shape下比动态shape下通常快10%~20%。养成好习惯能用tf.TensorShape明确标注的地方别偷懒写None。dtype的影响更直观。默认的tf.float32在GPU上占据4字节tf.float16只占2字节对应带宽占用减半、计算吞吐可以翻倍。但很多人不敢用float16是因为数值溢出——这恰恰是混合精度要解决的问题后面专门讲。device决定计算发生在哪里。一个需要注意的坑TensorFlow不会自动帮你把张量从GPU拷回CPU再拷回GPU。如果你在GPU上做计算中间为了调用某个NumPy函数把张量转成numpy数组再转回TensorFlow张量这里面的PCIe拷贝开销足以吞掉你辛辛苦苦优化出来的计算时间。1.3 tf.Tensor的只读属性与视图机制理解数据流的钥匙我第一次踩tf.Tensor只读的坑是在做一个数据增强模块时想复用缓冲区减少内存分配。结果发现tf.Tensor根本没有__setitem__方法原地修改是不可能的。后来我意识到TensorFlow这种设计是有意的——只读张量意味着在执行流中可以被安全地共享和引用不需要担心数据被意外修改这让计算图的执行模型可以放心地对算子做并行和重排。而视图view机制则隐藏在reshape、transpose、slice这些操作背后。TensorFlow底层存储是连续的reshape和slicing通常只改变元数据、不拷贝数据所以非常快。但transpose就没这么好运——它改变了内存访问顺序在TensorFlow里很多情况下会触发数据拷贝把非连续的内存重新排成连续布局。这也是为什么同一段代码里的transpose用多了性能会肉眼可见地变差。理解这一点你就明白为什么很多优化指南反复强调减少transpose、优先用reshape和broadcast。2. 环境配置里的坑CUDA/CUDNN/驱动版本的三角恋情2.1 TensorFlow 2.5.0的版本匹配表为什么不能随便装搜TensorFlow相关热词时tensorflow 2.5.0 cuda cudnn nvidia驱动 driver version: 550.144.03这个长尾词出现频率很高说明这是大家共同的痛点。先说结论TensorFlow 2.5.0这个版本官方在编译时绑定的是CUDA 11.2和CUDNN 8.1。如果你机器上装的是更新的CUDA 12.x直接装TensorFlow 2.5.0跑GPU十有八九会报Could not load dynamic library libcudnn.so.8这类错误。原因在于TensorFlow的pip包是预编译的它内部的二进制是在特定CUDA/CUDNN版本下编译出来的运行时动态加载这些库版本不匹配就直接拒绝工作。所以不要问我能不能装CUDA 12.5然后配TensorFlow 2.5.0——答案是可以但你需要自己处理兼容层。实际上最简单可靠的做法是要么升级TensorFlow版本2.10及以下支持GPU的最后一个版本是2.10因为2.11之后Windows上的GPU支持转成了WSL2方案要么老老实实按官方匹配表装CUDA 11.2 CUDNN 8.1。2.2 驱动550.144.03到底意味着什么驱动版本550.144.03这个数字是NVIDIA Linux驱动的版本号。NVIDIA驱动和CUDA Toolkit之间有兼容关系CUDA 11.2需要的最低驱动版本是460.x只要驱动版本高于这个要求就可以正常运行CUDA 11.2的程序。所以驱动550.144.03完全能向下兼容CUDA 11.2这点放心。但有个容易忽略的点nvidia-smi显示的CUDA版本是驱动支持的最高CUDA版本不是系统里实际安装的CUDA Toolkit版本。很多人看到nvidia-smi显示CUDA 12.5就以为系统里装了CUDA 12.5其实驱动只是支持而已实际开发还需要额外安装Toolkit或者依赖TensorFlow pip包自带的那套CUDA/CUDNN库很多官方预编译包通过pip安装时会顺带拉入依赖库在Python环境的nvidia目录下。我的经验是不要用nvidia-smi里的CUDA版本来判断TensorFlow能不能跑。网上传来传去的版本匹配表只是参考。真正的验证方式是装完后直接在Python里跑一段tf.config.list_physical_devices(GPU)能看到GPU设备就说明驱动和CUDA层的连接已经打通。再跑一个简单的tf.matmul对比CPU和GPU时间就能确认CUDNN和底层算子库也正常工作。2.3 手把手环境自检用张量操作验证整套GPU链路环境配好后我强烈建议做一次完整的链路自检别急着开始写模型。有一段核心自检脚本我每次在新机器上都会跑一遍包含import tensorflow as tf # 检查GPU设备是否可见 print(GPU available:, tf.config.list_physical_devices(GPU)) # 检查TensorFlow版本和编译信息 print(TF version:, tf.__version__) print(Built with CUDA:, tf.test.is_built_with_cuda()) # 真实计算验证 with tf.device(/GPU:0): a tf.random.normal([4096, 4096]) b tf.random.normal([4096, 4096]) c tf.matmul(a, b) # 强制同步确保计算完成 _ c.numpy() print(GPU matmul result shape:, c.shape) with tf.device(/CPU:0): a_cpu tf.random.normal([4096, 4096]) b_cpu tf.random.normal([4096, 4096]) c_cpu tf.matmul(a_cpu, b_cpu) _ c_cpu.numpy() print(CPU matmul result shape:, c_cpu.shape)这个过程如果哪个环节报错定位思路是驱动没装好就去找驱动工具链TensorFlow加载不了CUDA库就检查pip包里的依赖或者直接用官方推荐的镜像源重新安装CUDNN库报错就检查版本匹配。环境通了之后性能优化才有讨论的前提不然一切都是在空中楼阁。3. 张量创建与基础操作性能的隐性成本往往藏在这里3.1 创建张量的姿势对比常量、变量、占位符与变量的坑在TensorFlow 2.x里tf.constant和tf.Variable是最常用的两种张量创建方式。它们底层的行为差别需要每个TensorFlow使用者都清楚。tf.constant创建的是不可变张量TensorFlow可以对它做常量折叠——如果在计算图中多次使用同一个常量编译器会复用内存甚至把一个常量直接嵌入到kernel里。所以能用tf.constant的地方别用tf.Variable特别是在定义一些不变的索引表、mask、滤波核时。tf.Variable则更像传统编程中的可变状态。它的创建需要注意在GPU上初始化一个大变量比如几GB的Embedding表是很贵的因为需要分配显存并执行初始化kernel。一个常见优化是tf.Variable配上tf.float16或者更小的dtype可以显著降低显存占用和初始化时间。另外tf.placeholder已经随1.x退出历史舞台了。但很多老项目代码里还能看到运行时容易直接报错。TensorFlow 2.x的哲学是直接传张量——你不需要定义占位符计算时用Python变量直接传入tf.function即可。这不仅API更简洁底层的执行性能也因为计算图重构而更高效。3.2 重塑与转置的本质区别为什么reshape被推荐、transpose要慎用先说结论tf.reshape和tf.transpose的性能开销完全不是一个量级。reshape是元数据操作通常不涉及数据移动时间复杂度可以认为是O(1)级别的——它只是改变了张量的逻辑解释方式而transpose改变的是张量各维度的存储顺序在底层存储连续的前提下transpose几乎必然需要将数据重新排列是一个实打实的数据拷贝操作。这个差异在高维张量上尤其明显。比如一个[8, 128, 128, 32]的特征图在做transpose([0, 3, 1, 2])时要搬运的数据总量是8128128*32个元素相当于全量拷贝一遍。如果你在计算热的路径上反复做transpose那这部分开销很快就会成为瓶颈。有些场景下transpose无法避免比如需要把[batch, height, width, channel]转成[batch, channel, height, width]才能对接某些算子。我的建议是在GPU计算中把transpose操作合并进下一次需要的算子中。例如使用tf.nn.conv2d配合data_formatNCHW有时可以避免显式的维度交换或者在构建模型时一开始就选好layout全程保持唯一的维度排列方式减少从NCHW到NHWC的来回跳动。3.3 广播机制的甜蜜与陷阱广播broadcasting是张量操作最方便的机制之一但也是一把双刃剑。a tf.random.normal([1024, 512]) b tf.random.normal([1, 512]) c a b # 广播b被隐式扩展为[1024, 512]这个写法在代码层面非常优雅。但底层的实际算力开销取决于kernel实现现代GPU的cuDNN和Eigen内核CPU后端处理广播时很多情况并不是真正复制数据而是通过步长参数在计算时复用同一个内存区域。所以广播本身并不一定导致额外内存开销反而可能比显式tf.tile更快。但陷阱在于广播到超大shape时即使kernel不显式复制数据显存带宽的计算量还在如果广播结果还要跟另一个张量做逐元素的复杂运算中间数据范围变大可能导致缓存命中率下降。我遇过一种情况a b里b的shape是[512]而a是[100000, 512]带宽开销比预期高很多。优化时可以考虑在预处理阶段把b扩展成完整shape配合在tf.function里做JIT编译后性能反而提升。具体哪种好建议profile说话自己tf.profiler跑一遍才知道。4. 高性能计算的核心机制并行、编译、融合与数值精度4.1 并行化的关键算子粒度比Python循环重要一百倍许多从算法转过来的开发者有个惯性思维要并行就用Python写多线程、多进程。但到了TensorFlow这层真正的并行粒度在算子和kernel层面而不是Python循环层面。TensorFlow的GPU并行效果取决于算子operation是否调用了底层的并行kernel。典型的tf.matmul、tf.conv2d、tf.nn.embedding_lookup都已经被NVIDIA的cuBLAS、cuDNN优化到极致它们内部会占用整个GPU的SM单元做大规模并行。而你在Python层写一个for循环逐个处理序列元素每次调用一个小算子GPU的并行能力基本发挥不出来——因为每个小算子之间的调制、调度、内存分配都要走一遍框架的路由开销远大于计算本身。一个经验法则把逐元素操作的Python循环改成张量级操作性能提升通常在几个数量级。比如计算一组向量的L2范数不要写for x in xs: np.linalg.norm(x)而是把堆成[batch, dim]的张量一次tf.norm算完。GPU喜欢一次性吞下大批数据讨厌小粒度零碎调用。4.2 tf.function与AutoGraph让Python代码变成高速计算图的核心手段TensorFlow 2.x里最容易提升性能的工具就是tf.function。它通过AutoGraph机制把Python函数的控制流如if、for转换成TensorFlow的图操作进而可以整图编译优化。我第一次用tf.function跑一段数据生成代码时速度提升了近4倍当时最大的感受是原来Python解释器的开销可以这么轻易地被图执行消除。但tf.function有它自己的脾气。它默认会把Python参数int、str、bool等当作常量处理所以如果你的同一个函数被不同的Python参数调用TensorFlow会重新trace、生成多个专化计算图这就是所谓的retrace问题。我在实际项目里见过有人在一个for循环里动态增大多个int参数来调同一个tf.function结果Volatile GPU Util直接掉到接近0性能比不用tf.function还差。正确做法是输入数据尽量用tf.Tensor或tf.TensorSpec明确指定避免不必要的重编译模型内部如果有动态维度的RNN/Transformer类结构尽量把输入静态化到固定长度或者使用tf.shape里的动态维度时保持谨慎。4.3 XLA编译把多个小算子熔成一个大kernelXLAAccelerated Linear Algebra是TensorFlow自带的高性能编译器能把多个小算子优化融合成单个kernel减少kernel启动开销和中间张量的内存读写。这对GPU特别有效GPU kernel启动是有明显开销的如果你写一个上百个算子的大计算图仅kernel launch的时间都可能占很大比例。启用XLA不复杂tf.function(jit_compileTrue)就能开启或者在全局配置TF_XLA_FLAGS--tf_xla_auto_jit2。实测下来BERT类模型的Transformer层在开启XLA后训练速度能提升15%~30%而在一些纯张量计算密集型任务如矩阵链乘上甚至能提升数倍。不过XLA并非万灵药。它对动态shape支持有限某些自定义算子无法融合还可能导致编译时间显著变长。我有一次把一个包含大量mask操作的模型开启XLA后显存暴涨——因为XLA为了融合算子倾向于在显存里保留多个中间结果。这种情况需要评估后再决定是否全局启用XLA比较稳妥的是只对热区hot region用jit_compileTrue。4.4 混合精度内存减半、速度翻倍的数学原理与实操混合精度训练是当前提升GPU利用效率最直接的手段之一。基本原理很简单用float16存参数和中间激活值用float32做梯度累积必要时用tf.keras.mixed_precision.set_global_policy(mixed_float16)开启。为什么能提速GPU的FP16算力通常比FP32高出一倍甚至更多比如V100上FP16是112 TFLOPSFP32是14 TFLOPS差了8倍。同时张量数据量减半显存带宽压力也随之减半。这意味着同样的模型在支持FP16加速的GPU上用混合精度训练不仅省显存计算速度还会明显提升。但FP16的动态范围很小最大约65504最小正常数约6e-5。如果直接拿FP16做计算梯度很容易下溢成0。混合精度的关键是损失缩放loss scaling——把loss乘以一个缩放因子如1024让梯度进入FP16可表示的范围梯度计算后再缩放回来更新FP32权重。TensorFlow的混合精度API会把loss scaling自动处理好但你得知道它存在否则遇到loss突然变Nan时排查思路会走弯路。关于精度损失绝大多数模型的混合精度训练与全精度几乎没差别。个别稳定性差、对数值极其敏感的模型可以只对部分算子使用FP16。例如在TensorFlow里手动控制dtype某些层保持float32其余用float16也能获得大部分收益。5. 张量操作实战优化从一份能跑的代码到跑得快的完整记录5.1 基线版本一个典型的低效写法我拿一个我们实际项目中的案例来走一遍优化链路从一组样本计算加权滑动窗口的均值特征然后做归一化。原始代码里第一版是Python循环加NumPy混用TensorFlow部分反而只用了最外围的tf.convert_to_tensor。def compute_feature_tf_slow(samples, weights): # samples: [num_samples, window, dim] results [] for i in range(samples.shape[0]): window_data samples[i] weighted_sum tf.zeros([samples.shape[2]], dtypetf.float32) for j in range(samples.shape[1]): weighted_sum weighted_sum window_data[j] * weights[j] results.append(weighted_sum) return tf.stack(results)这种写法的问题非常典型内层循环逐个timestep做乘法加法外层循环逐个样本处理。在GPU上几乎没有任何并行计算可言Python层循环和tf算子调用开销占大头。实测在[64, 500, 128]规模的样本上单次推理耗时接近800ms慢到没法用。5.2 向量化与张量化用矩阵操作替代循环第一步优化很直接把滑动窗口加权这个操作变成矩阵乘法。权重长度为500窗口数据是[64, 500, 128]每个样本的窗口加权和就是从[500, 128]的矩阵和[500]的权重向量做加权求和这本质上就是一个批量矩阵向量乘。def compute_feature_tf_vec(samples, weights): # samples: [num_samples, window, dim] # weights: [window] weights tf.reshape(weights, [-1, 1]) # [window, 1] weighted samples * weights # 广播乘法 [num_samples, window, dim] return tf.reduce_sum(weighted, axis1) # [num_samples, dim]这一步的结果是质变原本嵌套循环变成三个张量级算子reshape、broadcast multiply、reduce_sumGPU并行效率一下子提了上来。同一个数据规模上这个版本耗时从800ms降到约25ms。为什么能降这么多因为GPU终于开始大批量并行计算了——窗口维度上的500个元素、样本维度上的64个样例都在SM上同时跑不再是一个一个来。但是25ms还不够好。观察一下samples * weights这一步生成了中间张量[64, 500, 128]约4MB的临时数据它需要被读取-写入-读取。这个读写开销是可以在更底层消除的。5.3 算子融合与XLA把多个算子合成一个kernel于是第二步优化把这个函数包进tf.function并开启XLA编译。tf.function(jit_compileTrue) def compute_feature_tf_xla(samples, weights): weights tf.reshape(weights, [-1, 1]) weighted samples * weights return tf.reduce_sum(weighted, axis1)XLA会把reshape、multiply、reduce_sum这三个算子融合成一个kernel中间张量weighted不再真正落到显存而是直接在寄存器或L2缓存里完成整个计算。实测这个版本耗时从25ms进一步降到约12ms。这里提升的主要来源就是减少了中间张量的内存写入/读取同时kernel启动次数从3次变成了1次。注意一点这个例子中的XLA之所以这么有效是因为算子链路非常干净——没有动态shape、没有不规则控制流、算子类型也完全在XLA支持范围内。如果你的优化路径里存在Python动态分支或某些tf自定义算子XLA的效果可能会打折扣需要自己权衡。5.4 减少数据拷贝与内存对齐最后一公里第三步优化是看主机的数据链路。原始流程是从NumPy读数据转成tf.Tensor送进GPU计算。如果每次推理都调用.numpy()把结果拷贝回主机这一步的开销在实时性要求高时同样不能忽视。在实际项目中我们优化成了预处理阶段把数据一次性打成TFRecord格式训练/推理时用tf.data直接加载到GPU显存全程避免Python层和GPU显存之间的来回拷贝。而且用tf.data的prefetch(tf.data.AUTOTUNE)让数据加载和GPU计算重叠执行无需人工同步等待。这一步之后整个读取样本-特征计算链路从原来的数据拷贝开销无法容忍变成了GPU利用率保持在80%以上推理吞吐稳定提升到初始版的约25倍从800ms降到约12~14ms后期整体链路优化后几乎不再有可感知的瓶颈。5.5 性能数据复盘与踩坑记录把这个案例的优化进程和对应时间整理一下优化阶段主要手段耗时说明基线Python循环 算子零散调用约800msGPU利用率极低大部分时间在调度张量化广播乘法 reduce_sum约25ms并行度大幅提升但仍有多余中间张量XLA融合tf.function(jit_compileTrue)约12mskernel启动次数减少、中间读写消除数据链路tf.data prefetch 避免numpy()端到端吞吐提升明显重点在于消除主机/设备拷贝与等待踩过的坑也很典型第一次开XLA因为数据shape里有None维词表动态填充编译频繁重trace整体比不开还慢。后来统一了输入shape用固定长度补齐XLA才稳定下来。还有一次把reduce_sum在axis1上的操作误写成了axis2导致广播维度不匹配梯度链整个乱了。这种张量维度语义的错误报错往往非常隐晦需要仔细print每个中间张量的shape逐层排查。6. 与PyTorch的张量哲学对比为什么有的操作快、有的操作慢6.1 动态图与静态图的编排差异对张量操作的影响TensorFlow 2.x的Eager模式和PyTorch一样都是动态图但TensorFlow的tf.function静态图执行模式在性能优化上走得更远。PyTorch也有torch.compile来追赶但两者的设计哲学仍然不同。TensorFlow的推荐路径是Eager模式下调试tf.function模式下交付。静态图的最大价值在于编译器可以看到整个计算的全貌能做跨算子优化和内存复用。比如上面案例中XLA的算子融合静态图下编译器知道reshape - mul - reduce_sum是一个完整的无分支链路就可以放心融合。动态图下每次执行都是独立解释框架不敢做激进的跨算子重排。PyTorch的张量操作更PythonicAPI直接、调试方便但性能的上限取决于你是否把整个模型的关键部分包进torch.compile或手写CUDA kernel。从张量操作的角度看TensorFlow在框架层面帮你优化这件事上做得更早也更系统Pytorch则更偏向把灵活性完全交给你。这没有绝对好坏取决于团队更看重调试体验还是部署性能。6.2 API风格差异同样一个reshape两边体验完全不同以一个常见的需求为例把[batch, seq, hidden]变换成[batch * seq, hidden]。TensorFlow的写法是tf.reshape(tensor, [-1, hidden])PyTorch是tensor.view(-1, hidden)或tensor.reshape(-1, hidden)。表面看几乎一样但TensorFlow的tf.reshape对内存连续性更加宽容必要时自动拷贝PyTorch的view则要求内存连续连续才不拷贝不连续了要用contiguous()。这个差异会导致初学者在从PyTorch往TensorFlow迁移时少写很多contiguous()但同时也更容易忽略内存布局问题。在做高性能计算时TensorFlow这边反而更容易写出跨设备的统一代码——你在tf.function里写tf.reshape在CPU和GPU上的行为一致PyTorch的view和contiguous组合在CPU/GPU之间偶尔会有不同的性能表现。6.3 2024年流行趋势下张量操作能力如何影响框架选型从这两年社区活跃度和开源项目数量看PyTorch在研究领域已经占据生态主导地位很多新论文默认PyTorch实现TensorFlow则在工业界有深厚积累特别是TensorFlow Serving、TFX流水线、以及TPU生态依然是大规模生产环境里不可忽视的力量。如果你看重的是张量操作和静态图优化能力TensorFlow的tf.functionXLA组合在部署性能一致性上依然有不可替代的优势——尤其当你需要在不同硬件平台CPU/GPU/TPU上保持相同计算语义时。PyTorch的灵活性更适合快速迭代的科研探索场景但在一些极端规模化、低延迟推理的方案上TensorFlow的图优化思路往往能帮你省下更多工程成本。我个人的建议不要盲目跟风框架之争。如果你是做研究选你团队和审稿人最熟悉的生态如果你在搭建长生命周期的生产系统认真评估两者的张量操作API能否支撑你的性能需求。框架只是工具张量操作的理解深度才是真正通用的能力——换框架只是换语法底层的数据布局、算子调度、并行计算逻辑是相通的。7. 写在最后的一点实在经验张量操作优化这件事说到底是搞清楚你写的每一行TF代码在底层发生了什么。很多时候性能瓶颈不在模型结构而在不起眼的数据预处理、维度变换、广播操作里。我见过太多工程师把大量时间花在超参调优上却忽略了一个多余的transpose或一个unnecessary的.numpy()拷贝带来的耗时——这些往往是几十倍性能差距的来源。最后分享一个我一直在用的习惯每当写一个相对复杂的张量操作函数我会花两分钟用TensorFlow Profiler看一眼算子的执行时间分布。不要凭感觉猜测瓶颈在哪profile数据会告诉你答案。优化完一个算子立刻重新profile确认收益别忙着优化下一个——很多人就是栽在感觉应该优化这里上。如果你刚开始学TensorFlow先把这篇里的基础概念弄扎实如果你已经有经验试着把你之前写过的训练循环或数据管道拿出来看看能不能用tf.function、XLA、混合精度和tf.data再优化一轮。这套思路在CPU、单卡GPU、多卡分布式上都是通用的真正掌握以后你会发现自己对TensorFlow的理解上升一个台阶。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进