如何在TensorFlow2中用ragged tensor实现指定索引采集逻辑降低显存占用
TensorFlow2多切片gather_nd合并实现方案
你设想的直接将ind1/ind2/ind3封装为ragged tensor后调用tf.gather_nd的写法不能直接得到目标结果,核心原因是三个索引对应的g的切片存在偏移,索引的第一维度是相对切片的局部索引,没有对齐g的全局索引。
不需要扩展索引为和g同形状的大张量,也不需要依赖ragged tensor,仅需要给三组索引分别加上对应切片的第一维偏移量,拼接后单次调用tf.gather_nd即可实现完全等价的功能,显存开销极低:
import tensorflow as tf import numpy as np # 原有变量定义保持不变 g = tf.random.normal([1000,627,16,1]) ind1 = tf.constant(np.random.randint(100,size=(251,128,128,3))) ind2 = tf.constant(np.random.randint(100,size=(150,128,128,3))) ind3 = tf.constant(np.random.randint(100,size=(100,128,128,3))) # 给索引添加对应切片的全局偏移:仅修改第一维,其余维度偏移为0 ind1_global = ind1 + tf.constant([0,0,0], dtype=ind1.dtype) ind2_global = ind2 + tf.constant([30,0,0], dtype=ind2.dtype) ind3_global = ind3 + tf.constant([120,0,0], dtype=ind3.dtype) # 拼接所有全局索引后单次调用gather_nd all_ind = tf.concat([ind1_global, ind2_global, ind3_global], axis=0) f = tf.gather_nd(g, all_ind) print('debug')
该方案的优势:
- 额外显存开销仅为三个偏移常量和索引拼接的临时内存,远低于扩展索引到g同形状的方案
- 不需要三次拆分调用gather_nd,运行效率更高
- 结果和原有三次调用后concat的结果完全一致
如果你一定要用ragged tensor实现,也需要先给每个子张量加上对应偏移再调用tf.gather_nd,实际运行效率和显存开销和上述普通张量方案没有差异,没有额外收益。
内容的提问来源于stack exchange,提问作者wangwei
相关产品推荐
相关产品推荐

