纯TensorFlow实现C51:替换np.add.at以加速训练
纯TensorFlow重写C51的add_to函数(替代np.add.at)
我有一个可运行的TensorFlow版C51实现,但因为使用tf.numpy_function调用依赖np.add.at的add_to函数,导致训练循环速度大幅下降。尝试用tf.tensor_scatter_nd_add重写np.add.at但未成功,现寻求不依赖np.add.at的替代实现,或用纯TensorFlow函数重写add_to函数。
原代码如下:
batch_size = 512 v_max = 5 v_min = -5 n_atoms = 51 support = tf.cast(tf.expand_dims(tf.linspace(v_min, v_max, 51), axis=0), tf.float32) delta_z = (v_max - v_min) / float(n_atoms - 1) # 需要重写的函数 def add_to(l, u, l_dist, u_dist): m = np.zeros_like(l, dtype=np.float32) for i in range(batch_size): np.add.at(m[i], l[i], l_dist[i]) np.add.at(m[i], u[i], u_dist[i]) return m # 完整流程函数 def projected_dist(t_dist, rewards, actions): Tz = tf.broadcast_to(support, [batch_size, support.shape[1]]) Tz = (0.99 ** 5) * Tz Tz += tf.expand_dims(rewards, axis=-1) Tz = tf.clip_by_value(Tz, v_min, v_max) b = (Tz - v_min) / delta_z l, u = tf.math.floor(b), tf.math.ceil(b) l_dist = t_dist * (u - b) u_dist = t_dist * (b - l) l, u = tf.cast(l, tf.int32), tf.cast(u, tf.int32) m = tf.numpy_function(add_to, [l, u, l_dist, u_dist], tf.float32, stateful=False) perjected_dist = tf.clip_by_value(m, 0.0, 1.0) null_dist = tf.zeros((batch_size, 3, n_atoms), tf.float32) indices = tf.stack([tf.range(0, batch_size, dtype=tf.int32), tf.cast(actions, tf.int32)], axis=-1) return tf.tensor_scatter_nd_add(null_dist, indices, perjected_dist)
解决方案:纯TensorFlow重写add_to函数
np.add.at的核心是对指定索引位置进行累加操作,用tf.tensor_scatter_nd_add可以实现完全等价的功能,关键是构造正确的二维索引矩阵:
- 生成每个样本的索引序列,与原子维度索引拼接,得到所有需要更新的位置坐标
- 将
l_dist和u_dist的对应值合并,作为累加的数值 - 初始化全零矩阵后执行批量累加
重写后的add_to函数及完整修改代码如下:
batch_size = 512 v_max = 5 v_min = -5 n_atoms = 51 support = tf.cast(tf.expand_dims(tf.linspace(v_min, v_max, 51), axis=0), tf.float32) delta_z = (v_max - v_min) / float(n_atoms - 1) # 纯TensorFlow实现的add_to函数 def add_to(l, u, l_dist, u_dist): # 初始化全零矩阵 m = tf.zeros_like(l_dist, dtype=tf.float32) # 生成样本维度的索引:(batch_size, n_atoms) batch_indices = tf.tile(tf.range(batch_size)[:, tf.newaxis], [1, n_atoms]) # 构造l对应的二维索引:(batch_size * n_atoms, 2) l_indices = tf.stack([tf.reshape(batch_indices, [-1]), tf.reshape(l, [-1])], axis=-1) # 构造u对应的二维索引:(batch_size * n_atoms, 2) u_indices = tf.stack([tf.reshape(batch_indices, [-1]), tf.reshape(u, [-1])], axis=-1) # 合并两部分索引 all_indices = tf.concat([l_indices, u_indices], axis=0) # 合并对应的累加值 all_updates = tf.concat([tf.reshape(l_dist, [-1]), tf.reshape(u_dist, [-1])], axis=0) # 执行批量累加 m = tf.tensor_scatter_nd_add(m, all_indices, all_updates) return m def projected_dist(t_dist, rewards, actions): Tz = tf.broadcast_to(support, [batch_size, support.shape[1]]) Tz = (0.99 ** 5) * Tz Tz += tf.expand_dims(rewards, axis=-1) Tz = tf.clip_by_value(Tz, v_min, v_max) b = (Tz - v_min) / delta_z l, u = tf.math.floor(b), tf.math.ceil(b) l_dist = t_dist * (u - b) u_dist = t_dist * (b - l) l, u = tf.cast(l, tf.int32), tf.cast(u, tf.int32) # 直接调用纯TensorFlow的add_to,无需tf.numpy_function m = add_to(l, u, l_dist, u_dist) perjected_dist = tf.clip_by_value(m, 0.0, 1.0) null_dist = tf.zeros((batch_size, 3, n_atoms), tf.float32) indices = tf.stack([tf.range(0, batch_size, dtype=tf.int32), tf.cast(actions, tf.int32)], axis=-1) return tf.tensor_scatter_nd_add(null_dist, indices, perjected_dist)
说明
- 去掉了
tf.numpy_function的跨语言调用开销,训练速度会显著提升 - 新的
add_to完全基于TensorFlow原生操作,支持自动微分和图模式优化 - 逻辑与原函数完全一致:对每个样本的每个原子位置,将对应分布值累加到
l和u索引处
内容的提问来源于stack exchange,提问作者user8075709
相关产品推荐
相关产品推荐

