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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 05:06:05