关于TensorFlow中tf.linalg.eye参数及创建3x3单位矩阵的问询
tf.linalg.eye 中 num_rows 和 num_columns 参数说明及 3x3 单位矩阵示例
参数用法解释
num_rows:指定单位矩阵的行数,为必填参数,决定矩阵垂直方向的维度。num_columns:指定单位矩阵的列数,为可选参数。若不设置,默认与num_rows取值一致,生成方阵;若设置不同数值,会生成非方阵的类单位矩阵——主对角线上从左上角开始的min(num_rows, num_columns)个位置为1,其余位置为0。
举两个非方阵的例子:
- 当
num_rows=3,num_columns=5时,生成3行5列矩阵,主对角线前3个元素为1,其余为0; - 当
num_rows=5,num_columns=3时,生成5行3列矩阵,主对角线前3个元素为1,其余为0。
3x3 单位矩阵示例
生成3x3单位矩阵有两种写法,效果完全一致:
import tensorflow as tf # 写法1:仅指定行数,默认列数与行数相同 eye_3x3 = tf.linalg.eye(num_rows=3) print(eye_3x3.numpy()) # 写法2:显式指定行数和列数 eye_3x3_explicit = tf.linalg.eye(num_rows=3, num_columns=3) print(eye_3x3_explicit.numpy())
运行后输出结果均为:
[[1. 0. 0.] [0. 1. 0.] [0. 0. 1.]]
内容的提问来源于stack exchange,提问作者Kaddu Livingstone
相关产品推荐
相关产品推荐

