如何从两个Ragged Tensor中随机选取元素并实现批处理?
解决Ragged Tensor中随机选取行内元素并支持批处理的问题
我明白你的需求了——原来的代码是随机抽取整个行,现在要进一步在选中的行里随机挑单个元素,还要支持批处理(一次处理多个样本)。下面我给你一步步拆解实现思路,附上完整的可运行代码:
实现步骤说明
- 第一步:生成批处理规模的行索引,比如你想要一次处理3个样本,就生成3个0到行数-1之间的随机整数
- 第二步:根据行索引,获取对应行的长度,然后在每个行的长度范围内生成随机的元素索引
- 第三步:把行索引和元素索引组合成
tf.gather_nd需要的二维索引格式(每个元素是[行号, 元素位置]) - 第四步:用组合后的索引对两个Ragged Tensor执行
gather_nd,就能批量得到对应位置的元素
完整示例代码
import tensorflow as tf # 定义两个同结构的Ragged Tensor c = tf.ragged.constant([[1, 2, 3], [4, 5]]) v = tf.ragged.constant([[10., 20., 30.], [40., 50.]]) # 设置批处理大小,比如一次处理3个样本 batch_size = 3 # 1. 生成批处理用的随机行索引(范围0到行数-1) row_indices = tf.random.uniform([batch_size], maxval=tf.shape(c)[0], dtype=tf.int32) # 2. 获取每个选中行的长度,然后生成对应行内的随机元素索引 # 先获取所有行的长度 row_lengths = c.row_lengths() # 根据行索引拿到对应行的长度 selected_row_lengths = tf.gather(row_lengths, row_indices) # 生成每个行内的随机元素索引(范围0到对应行长度-1) elem_indices = tf.map_fn( lambda x: tf.random.uniform([], maxval=x, dtype=tf.int32), selected_row_lengths, dtype=tf.int32 ) # 3. 组合成gather_nd需要的索引格式:shape为[batch_size, 2] gather_indices = tf.stack([row_indices, elem_indices], axis=1) # 4. 对两个Ragged Tensor执行gather_nd,得到批处理结果 c_results = tf.gather_nd(c, gather_indices) v_results = tf.gather_nd(v, gather_indices) # 运行验证 with tf.Session() as sess: rows, elems, c_out, v_out = sess.run([row_indices, elem_indices, c_results, v_results]) print(f"选中的行索引: {rows}") print(f"行内元素索引: {elems}") print(f"c的结果: {c_out}") print(f"v的结果: {v_out}")
代码解释
row_lengths = c.row_lengths():获取Ragged Tensor每一行的长度,这里会得到[3, 2]selected_row_lengths = tf.gather(row_lengths, row_indices):根据随机选中的行索引,拿到对应行的长度,比如如果行索引是[0,1,0],这里就会得到[3,2,3]tf.map_fn:遍历每个选中行的长度,生成该行内的随机元素索引,确保不会超出该行的长度范围tf.stack([row_indices, elem_indices], axis=1):把行索引和元素索引拼成[[行号, 元素位], ...]的格式,这是tf.gather_nd要求的索引格式
运行这段代码,你会得到类似这样的输出(每次随机结果不同):
选中的行索引: [0 1 0] 行内元素索引: [2 0 1] c的结果: [3 4 2] v的结果: [30. 40. 20.]
这样就实现了先随机选行,再在该行内随机选元素的批处理逻辑,完全匹配你的需求~
内容的提问来源于stack exchange,提问作者JEM_Mosig
相关产品推荐
相关产品推荐

