TensorFlow中3D曲面生成的向量化优化方案问询
用向量化方式优化TensorFlow中3D曲面生成效率
问题背景
现有形状为[sampling_size* sampling_size, 2]的2D网格,当前通过Python循环遍历网格点生成3D曲面(如立方体、棱柱),但循环效率极低,生成单个立方体耗时达10秒,需改用向量化操作提升性能。
原实现代码(低效循环版)
网格生成代码
import tensorflow as tf import math sampling_size = 100 limit = math.pi def generate_grid(_from, _to, _step): range_ = tf.range(_from, _to, _step, dtype=float) x, y = tf.meshgrid(range_, range_) _x = tf.reshape(x, (-1,1)) _y = tf.reshape(y, (-1,1)) return tf.squeeze(tf.stack([_x, _y], axis=-1)), x, y grid, X, Y = generate_grid(-limit, limit, 2*limit / sampling_size)
立方体曲面生成(循环版)
def cube(G): res = [] for (X, Y) in G: if X >= -1 and X < 1 and Y >= -1 and Y < 1: res.append(1.) else: res.append(0.) return tf.convert_to_tensor(res) Z_cube = cube(grid) cube_2d = tf.reshape(Z_cube, [sampling_size, sampling_size])
棱柱曲面生成(循环版)
def prism(G): res = [] for (X, Y) in G: if X >= -1 and X < 1 and Y >= -1 and Y < 1: res.append(X + 1.) else: res.append(0.) return tf.convert_to_tensor(res) Z_prism = prism(grid) prism_2d = tf.reshape(Z_prism, [sampling_size, sampling_size])
绘图代码
import matplotlib.pyplot as plt from matplotlib import cm def plot_surface(X, Y, Z, a = 30, b = 15): fig = plt.figure() ax = plt.axes(projection='3d') ax.plot_surface(X, Y, Z, rstride=3, cstride=3, linewidth=1, antialiased=True, cmap=cm.viridis) ax.view_init(a, b) ax.set_xlabel('X') ax.set_ylabel('Y') ax.set_zlabel('Z') plt.show()
向量化优化方案
TensorFlow的核心优势是张量的批量操作,完全可以避免Python循环,直接对整个网格张量进行条件判断和计算,大幅提升效率。
优化后的立方体曲面生成
def cube_vectorized(G): # 提取X和Y分量 X = G[:, 0] Y = G[:, 1] # 生成布尔掩码:判断每个点是否在[-1,1)×[-1,1)范围内 mask = (X >= -1) & (X < 1) & (Y >= -1) & (Y < 1) # 将掩码转换为浮点型,符合原逻辑(范围内为1,否则为0) return tf.cast(mask, tf.float32) Z_cube = cube_vectorized(grid) cube_2d = tf.reshape(Z_cube, [sampling_size, sampling_size]) plot_surface(X, Y, cube_2d)
优化后的棱柱曲面生成
def prism_vectorized(G): X = G[:, 0] Y = G[:, 1] mask = (X >= -1) & (X < 1) & (Y >= -1) & (Y < 1) # 范围内的点取X+1,否则为0 Z = tf.where(mask, X + 1., 0.) return Z Z_prism = prism_vectorized(grid) prism_2d = tf.reshape(Z_prism, [sampling_size, sampling_size]) plot_surface(X, Y, prism_2d)
优化原理
- 规避Python循环的逐点处理开销,直接利用TensorFlow底层并行计算能力处理整个张量
- 布尔掩码和
tf.where均为向量化批量操作,执行效率远高于Python循环 - 无需手动维护列表并转换为张量,减少数据拷贝与转换的额外消耗
内容的提问来源于stack exchange,提问作者Luiz Doleron
相关产品推荐
相关产品推荐

