TensorFlow 2.8中如何正确获取稀疏张量SparseTensor的形状
问题原因
- 你代码中的
A_train是tf.SparseTensor(稀疏张量)类型,它的dense_shape属性本身是一个1维张量,存储的是稀疏张量映射到稠密空间时各维度的长度,并不是Python原生的整数列表/数值。 - 你用的两种打印方式拿不到预期值的原因很明确:
- 直接调用
print(A_train.dense_shape[0]):打印的是0维张量对象本身,输出格式类似tf.Tensor(943, shape=(), dtype=int64),不会直接返回可复用的Python整数。 - 调用
tf.print(A_train.dense_shape[0]):这个方法确实能在控制台输出具体的维度数值,但它没有返回值,你不能把打印出的结果直接作为参数传入形状定义逻辑。
- 直接调用
- 你看到的代码里
A_train.dense_shape[0]的取值,就是构造A_train稀疏张量时传入的稠密形状参数的第一个值,对应评分矩阵的总行数,也就是用户总数量,这个写法本身在TensorFlow的变量初始化逻辑里是合法的——TF内部会自动解析张量类型的形状参数,不需要手动转成原生整数,所以示例代码可以正常运行。
正确获取稀疏张量形状值的方法
根据使用场景选对应方式即可:
- 动态图模式下打印、复用维度值
直接对维度张量调用.numpy()方法,把张量值转成Python原生数值即可,不需要调用session.run:# 获取用户总数,转成Python int类型 user_count = int(A_train.dense_shape[0].numpy()) print("用户总数:", user_count) # 后续可以直接把user_count作为形状参数传入初始化逻辑 U = tf.Variable(tf.random.normal( [user_count, embedding_dim], stddev=init_stddev)) - 在
tf.function装饰的静态图代码块中取值
静态图模式下不能直接调用.numpy(),可以用tf.get_static_value获取静态可推导的数值,建议最好在进入tf.function之前就把维度值取好作为Python参数传入,避免静态图推导问题:# 静态图内获取静态值 user_count = int(tf.get_static_value(A_train.dense_shape[0]))
补充说明:你看到的示例代码用了
tf.random_normal是TF1.x的旧API,TF2.8+中推荐用等价的tf.random.normal,功能完全一致。
内容的提问来源于stack exchange,提问作者sakeesh
相关产品推荐
相关产品推荐

