如何在TensorFlow或Keras中提取(k,n,n)张量的k条对角线?
提取(k,n,n)张量中k个方阵的对角线方案
当然可以搞定这个需求!针对形状为(k, n, n)的张量,要提取每个(n,n)方阵的对角线并得到(k,n)的输出,这里有两种简单靠谱的方案:
方法1:直接使用tf.linalg.diag_part(推荐)
TensorFlow的tf.linalg.diag_part函数天生就支持高维张量的对角线提取——它会自动对张量的最后两个维度(也就是这里每个独立的(n,n)方阵)操作,直接返回每个方阵的对角线元素,输出形状正好是(k,n)。
举个代码例子:
import tensorflow as tf # 创建一个形状为(2, 3, 3)的测试张量 k, n = 2, 3 test_tensor = tf.random.normal(shape=(k, n, n)) # 提取对角线 diagonals = tf.linalg.diag_part(test_tensor) print("输入张量形状:", test_tensor.shape) print("输出对角线形状:", diagonals.shape) # 输出 (2, 3)
不管是TensorFlow原生代码还是Keras模型里(只要后端是TensorFlow),都可以直接用这个方法,非常简洁。
方法2:利用高级索引手动提取
如果你想更直观地控制索引逻辑,也可以通过张量索引来手动选取对角线元素:利用tf.range(n)生成对角线的位置索引,然后对每个(k,n,n)的张量取[:, idx, idx]的元素,同样能得到(k,n)的结果。
代码示例:
import tensorflow as tf k, n = 2, 3 test_tensor = tf.random.normal(shape=(k, n, n)) # 生成对角线的索引 idx = tf.range(n) # 手动提取对角线 diagonals = test_tensor[:, idx, idx] print("输出对角线形状:", diagonals.shape) # 同样输出 (2, 3)
这个方法逻辑清晰,适合需要自定义索引逻辑的场景(比如提取非主对角线的情况,只需要调整idx的生成方式即可)。
两种方法的结果完全一致,你可以根据自己的代码风格和需求选择~
内容的提问来源于stack exchange,提问作者OmarVP
相关产品推荐
相关产品推荐

