Skip to content

TensorFlow - 构建图

在 TensorFlow 2.x 中,默认启用 Eager Execution (即时执行),这意味着操作会立即执行。然而,TensorFlow 仍然支持图执行(Graph Execution),它提供了性能优化和可移植性等优势。在 TF 2.x 中创建图的主要方式是使用 tf.function 装饰器 (decorator)。当你用 tf.function 装饰一个 Python 函数时,TensorFlow 可以追踪它以创建一个可调用的 TensorFlow 图。

我们将通过模拟一个受偏微分方程 (PDE) 控制的简单物理系统——具体来说,是一个二维波动方程,比如池塘中的涟漪——来演示这一点。我们将在一个 Python 函数中定义模拟步骤,然后使用 tf.function 将其转换为高性能的图操作。

假设我们的池塘是一个大小为 N x N 的二维网格(例如,N=500)。

步骤 1:导入必要的库。

import tensorflow as tf
import numpy as np
import matplotlib.pyplot as plt
import time # 用于比较执行速度(可选)
# 确保如果启用了 TF1 兼容性,Eager Execution 已关闭
# 对于纯 TF2,默认启用 Eager,这并非严格必要
# tf.compat.v1.enable_eager_execution() # 确保 TF2 行为启用 Eager Execution
print(f"TensorFlow version: {tf.__version__}")
print(f"Eager execution enabled: {tf.executing_eagerly()}")

步骤 2:定义用于卷积(计算 Laplacian 算子)的辅助函数。

# 从 2D 数组创建卷积核的函数
def make_kernel(a):
a = np.asarray(a, dtype=np.float32)
# 将形状重塑为 [filter_height, filter_width, in_channels, out_channels]
# 对于 depthwise_conv2d,形状是 [filter_height, filter_width, in_channels, channel_multiplier]
# 这里,in_channels=1, channel_multiplier=1
return tf.constant(a.reshape(list(a.shape) + [1, 1]))
# 简化的 2D 卷积。注意:tf.nn.conv2d 更通用。
# 这里使用 tf.nn.depthwise_conv2d 是为了简化以匹配旧示例逻辑。
def simple_conv(x, k):
# x 形状:[height, width]
# k 形状:[kernel_height, kernel_width, 1, 1]
# 给 x 添加批量和通道维度:[1, height, width, 1]
x_expanded = tf.expand_dims(tf.expand_dims(x, 0), -1)
# 深度可分离卷积:每个输入通道与其自己的一组滤波器进行卷积。
# 这里,in_channels=1,因此类似于具有 1 个输入通道的标准卷积。
y = tf.nn.depthwise_conv2d(x_expanded, k, strides=[1, 1, 1, 1], padding='SAME')
# 移除批量和通道维度:[height, width]
return y[0, :, :, 0]
# 计算数组的二维 Laplacian 算子的函数
def laplace(x):
# Laplacian 核:近似空间二阶导数
laplace_k = make_kernel([
[0.5, 1.0, 0.5],
[1.0, -6., 1.0], # 中心权重是其他权重总和的负值(对于某些模板)
[0.5, 1.0, 0.5]
])
return simple_conv(x, laplace_k)

步骤 3:设置初始条件和模拟参数。

N = 100 # 减小网格尺寸以便在教程中更快执行
# 初始条件:模拟一些雨滴落入池塘
u_init = np.zeros([N, N], dtype=np.float32)
ut_init = np.zeros([N, N], dtype=np.float32) # 初始速度
# 添加一些扰动(雨滴)
for _ in range(20): # 减小雨滴数量
a, b = np.random.randint(0, N, 2)
u_init[a, b] = np.random.uniform()
# 显示池塘的初始状态
plt.figure(figsize=(6,6))
plt.imshow(u_init)
plt.title("池塘的初始状态 (u_init)")
plt.colorbar()
plt.show()
# 模拟参数(将作为 tf.function 的参数)
# eps: 时间分辨率 (delta_t)
# damping: 波的阻尼系数
# 创建 TensorFlow Variable 用于存储模拟状态
U = tf.Variable(u_init, name='U') # 当前位移
Ut = tf.Variable(ut_init, name='Ut') # 当前速度

步骤 4:在 Python 函数中定义 PDE 更新规则,并用 tf.function 装饰它。

@tf.function # 启用 tf.function 装饰器以生成可调用的 TensorFlow 图
def pde_step(current_U, current_Ut, eps, damping):
"""执行 PDE 模拟的一步。"""
# 离散化的 PDE 更新规则(波动方程的欧拉积分)
# U_t+1 = U_t + eps * Ut_t
# Ut_t+1 = Ut_t + eps * (laplacian(U_t) - damping * Ut_t)
new_U = current_U + eps * current_Ut
new_Ut = current_Ut + eps * (laplace(current_U) - damping * current_Ut)
# 更新状态变量(赋值新值)
current_U.assign(new_U)
current_Ut.assign(new_Ut)
# 在 TF2 eager/tf.function 中不需要 tf.group,因为赋值是按顺序执行的
return current_U, current_Ut
# 立即执行一步(可选,在多次迭代前先看是否正常工作)
# pde_step(U, Ut, tf.constant(0.03, dtype=tf.float32), tf.constant(0.04, dtype=tf.float32))
# print("一次 eager 步骤后的 U(最大值):", tf.reduce_max(U).numpy())

步骤 5:运行模拟。

num_iterations = 1000 # 迭代次数
eps_val = tf.constant(0.03, dtype=tf.float32) # eps 值(时间分辨率)
damping_val = tf.constant(0.04, dtype=tf.float32) # damping 值
print("使用 tf.function 启动模拟...") # 使用 tf.function 启动模拟...
start_time = time.time() # 记录开始时间
for i in range(num_iterations):
# 调用 tf.function 装饰的步骤函数
U, Ut = pde_step(U, Ut, eps_val, damping_val)
# 偶尔可视化进度
if i % 200 == 0 or i == num_iterations -1: # 降低频率以加快总体运行速度
print(f"Iteration {i}") # Iteration {i}
plt.figure(figsize=(6,6))
plt.imshow(U.numpy()) # 使用 .numpy() 获取数组以用于 matplotlib
plt.title(f"池塘状态 (Iteration {i})") # 池塘状态(第 {i} 次迭代)
plt.colorbar()
plt.show()
end_time = time.time() # 记录结束时间
print(f"Simulation finished in {end_time - start_time:.2f} seconds.") # 模拟在 {end_time - start_time:.2f} 秒内完成。

输出将包含显示“池塘”初始状态以及模拟过程中不同迭代状态的图。你应该会观察到类似波的模式传播和消散。tf.function 装饰器将 pde_step 转换为 TensorFlow 图,以实现潜在更快的执行速度,尤其是在涉及许多小型 TensorFlow 操作的计算中,通过减少 Python 开销和启用图优化。

tf.function 的要点:

  • 它弥合了 Eager Execution (即时执行) 的易用性与图执行的性能之间的差距。
  • 它最适合 TensorFlow 操作。tf.function 内的 Python 原生操作或 NumPy 操作可能会频繁触发重新追踪,或导致调用 tf.py_function,这会影响性能。
  • 如果作为 tf.function 输入参数的 Python 标量或 NumPy 数组的值发生变化,可能会触发重新追踪。通常建议使用 tf.Tensor 参数以保持稳定性。
  • tf.Variable 可以在 tf.function 中使用其 .assign() 方法进行更新。

有关 tf.function 的更多详细信息,请参阅官方 TensorFlow 指南:https://www.tensorflow.org/guide/function