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

TensorFlow中如何非原地计算对数或无需log直接进行类别采样

问题解决说明

  1. 首先澄清认知误区:TensorFlow的原生数学操作默认均为非原地操作,直接调用tf.math.log(move)只会生成新的张量,不会修改你外部传入的原始输入move,你担心的原地修改问题实际不存在。如果确实需要手动复制张量,不要用tf.constant(该接口用于从Python数值/ numpy数组生成静态常量,不支持传入运行时张量),改用tf.identity(move)即可,完全兼容tf.function和32位GPU环境。
  2. 解决0概率取log出NaN的问题:对输入概率加一个极小的截断值再取对数即可,32位浮点数场景下用1e-8就足够,不会影响概率的相对大小,也不会出现NaN。
  3. 修复原代码的索引越界bug:你输入的张量形状为(N,8,3),轴2的索引最大值为2,原代码中p3取索引3会直接报错,p2的索引也和前面定义的m2不匹配,已同步修正。

修正后可运行代码

@tf.function
def convert_to_move(move):
    # 加极小epsilon避免0概率取log出NaN
    log_move = tf.math.log(move + 1e-8)
    m1 = log_move[:, :, 0] 
    m2 = log_move[:, :, 1]
    m3 = log_move[:, :, 2]
    x1 = tf.random.categorical(m1, num_samples=1)
    x2 = tf.random.categorical(m2, num_samples=1)
    x3 = tf.random.categorical(m3, num_samples=1)

    r1 = tf.squeeze(tf.one_hot(x1, move.shape[1]))
    r2 = tf.squeeze(tf.one_hot(x2, move.shape[1]))
    r3 = tf.squeeze(tf.one_hot(x3, move.shape[1]))

    p1 = tf.boolean_mask(move[:, :, 0], r1)
    p2 = tf.boolean_mask(move[:, :, 1], r2)
    p3 = tf.boolean_mask(move[:, :, 2], r3)

    k1 = tf.stack([x1, x2, x3], axis=-1)
    k2 = tf.stack([p1, p2, p3], axis=-1)
    return tf.squeeze(k1), tf.squeeze(k2)

无对数计算的备选采样实现

如果完全不想使用对数操作,可以自己基于逆变换采样实现分类采样,直接输入概率值即可:

def sample_categorical(probs, num_samples=1):
    # probs形状要求为(N, num_classes),最后一维和为1
    u = tf.random.uniform(shape=(tf.shape(probs)[0], num_samples, 1))
    cumsum = tf.cumsum(probs, axis=-1)[:, tf.newaxis, :]
    samples = tf.argmax(tf.cast(cumsum > u, tf.int32), axis=-1)
    return samples

把原代码中tf.random.categorical替换为上述自定义函数即可,不需要额外计算对数。


内容的提问来源于stack exchange,提问作者user5178150

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 03:27:03