tf.random.categorical结果异常,求实现TensorFlow版np.random.choice
问题分析与解决
嘿,我找到你代码里的问题啦!核心原因是TensorFlow 1.x中,每次调用eval()或sess.run()时,所有依赖的随机操作都会重新执行一次。
看你的代码流程:
- 当你执行
sample_selected.eval()时,tf.random.categorical生成了索引3; - 但当你接着执行
sess.run(op)时,op依赖的sample_selected会再次触发tf.random.categorical,这次生成了新的随机索引1,所以p被赋值成了k[1]也就是1,和你之前打印的索引3完全不一致。
另外还有个小细节:你定义p的类型是tf.int32,但k默认是tf.int64(tf.convert_to_tensor处理整数列表时的默认类型),虽然TensorFlow会隐式转换,但最好保持类型一致,避免潜在的类型兼容问题。
修正后的代码方案
方案1:固定随机索引,避免重复生成
通过tf.identity保存生成的随机索引,确保后续操作复用同一个值:
import numpy as np import tensorflow as tf selection_sample = [i for i in range(10)] # 待选择的样本列表 k = tf.convert_to_tensor(selection_sample, dtype=tf.int64) # 生成随机索引并固定下来 samples = tf.random.categorical(tf.math.log([[1, 0.5, 0.3, 0.6]]), 1) sample_selected = tf.cast(samples[0][0], tf.int64) fixed_sample_idx = tf.identity(sample_selected) # 固定索引值,避免重复生成随机数 # 保持变量类型与k一致 p = tf.Variable(0, tf.int64) assign_op = tf.assign(p, k[fixed_sample_idx]) init = tf.global_variables_initializer() with tf.Session() as sess: sess.run(init) # 先获取一次固定的索引 selected_idx = sess.run(fixed_sample_idx) print(f"选中的索引:{selected_idx}") print(f"k的取值:{k.eval()}") # 执行赋值操作,复用之前的索引 sess.run(assign_op) print(f"p的最终值:{p.eval()}")
方案2:一次性获取所有结果
通过一次sess.run()调用获取所有需要的张量,确保随机操作只执行一次:
import numpy as np import tensorflow as tf selection_sample = [i for i in range(10)] k = tf.convert_to_tensor(selection_sample, dtype=tf.int64) samples = tf.random.categorical(tf.math.log([[1, 0.5, 0.3, 0.6]]), 1) sample_selected = tf.cast(samples[0][0], tf.int64) p = tf.Variable(0, tf.int64) assign_op = tf.assign(p, k[sample_selected]) init = tf.global_variables_initializer() with tf.Session() as sess: sess.run(init) # 一次性运行所有需要的张量,随机索引只生成一次 selected_idx, k_vals, _, p_val = sess.run([sample_selected, k, assign_op, p]) print(f"选中的索引:{selected_idx}") print(f"k的取值:{k_vals}") print(f"p的最终值:{p_val}")
额外提示:实现np.random.choice的更简洁方式
如果你的目标是完全模拟np.random.choice的功能,其实可以用tf.random.categorical结合权重,或者直接用tf.random.shuffle+切片,比如:
# 模拟np.random.choice(selection_sample, size=1, p=[0.25, 0.25, 0.25, 0.25]) weights = tf.constant([0.25, 0.25, 0.25, 0.25]) samples = tf.random.categorical(tf.math.log([weights]), num_samples=1) selected = tf.gather(selection_sample, samples[0][0])
内容的提问来源于stack exchange,提问作者onexpeters
相关产品推荐
相关产品推荐

