You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

纯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可以实现完全等价的功能,关键是构造正确的二维索引矩阵:

  1. 生成每个样本的索引序列,与原子维度索引拼接,得到所有需要更新的位置坐标
  2. 将l_dist和u_dist的对应值合并,作为累加的数值
  3. 初始化全零矩阵后执行批量累加

重写后的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.24 16:22:47