TensorFlow按随机列索引将[3136,410]张量更新到[3136,512]零张量的实现方法
实现方案
可以通过带axis参数的tf.scatter_nd实现该需求,该接口从TensorFlow 2.4版本开始支持,无需手动拼接行索引,代码逻辑非常简洁。
核心逻辑说明
你需要的按列赋值需求,刚好匹配axis=1的散射更新逻辑:
- 指定
axis=1表示更新操作作用在张量的列维度 - 你生成的
indices_col长度为410,和initial_weight的列数一一对应,第k个索引位置对应initial_weight的第k列 - 输出形状直接指定为目标张量形状
[3136, 512]即可,不需要额外reshape操作
完整可运行代码
import tensorflow as tf import random if __name__ == '__main__': # 张量形状定义 shape_for_layer = [3136, 512] subshape = [3136, 410] # 生成待赋值的初始权重张量 initial_weight = tf.random.uniform(subshape, minval=0, maxval=1, dtype=tf.float32) # 生成随机列索引 layer_456_col = random.sample(range(512), 410) indices_col = tf.convert_to_tensor(layer_456_col) # 核心更新逻辑:带axis参数的scatter_nd result = tf.scatter_nd( indices=indices_col, updates=initial_weight, shape=shape_for_layer, axis=1 ) # 验证输出 print("输出张量形状:", tf.shape(result)) # 验证数值正确性:取第一个索引对应列和初始权重第一列对比 print("数值验证结果:", tf.reduce_all(result[:, indices_col[0]] == initial_weight[:, 0]))
注意事项
如果你的TensorFlow版本低于2.4不支持axis参数,也可以通过手动构造坐标的方式实现,但更推荐升级到2.4及以上版本使用更简洁的写法。
内容的提问来源于stack exchange,提问作者Fanto
相关产品推荐
相关产品推荐

