Elixir中如何从4维Nx.Tensor创建符合规则的3维Nx.Tensor
Nx 3维张量构造方案
从已生成的4维张量n_tensor提取目标值
你构造的n_tensor维度顺序为[:i, :l, :k, :j],要保留l == k的元素,直接调用对角线提取接口即可,不需要写复杂的切片逻辑:
a_mat = Nx.random_uniform({5,5}, names: [:i, :l]) b_mat = Nx.random_uniform({5,5}, names: [:k, :j]) n_tensor = Nx.dot(a_mat,[], b_mat,[]) # 提取第1、2维度(对应:l、:k轴)对角线位置元素,自动压缩重合维度 m_tensor = n_tensor |> Nx.take_diagonal(axes: [1, 2]) |> Nx.rename_axis(1, :k) # 可选,把合并后的轴名统一为:k
提取后得到的m_tensor维度顺序为[:i, :k, :j],完全满足m[i][k][j] = a[i][k] * b[k][j]的运算规则。
这个方案存在冗余计算:会先生成625个元素的4维中间张量,最终只保留125个有效元素,仅适合已经生成
n_tensor、不想重复计算的场景。
效率更高的直接构造方案
不需要先生成4维张量,直接利用广播机制做维度对齐相乘即可,全程无冗余内存占用,计算速度更快:
a_mat = Nx.random_uniform({5,5}, names: [:i, :k]) b_mat = Nx.random_uniform({5,5}, names: [:k, :j]) m_tensor = Nx.multiply( Nx.new_axis(a_mat, -1), # 把a的维度扩展为[:i, :k, 1] Nx.new_axis(b_mat, 0) # 把b的维度扩展为[1, :k, :j] )
广播机制会自动把长度为1的维度复制扩展,相乘后直接得到维度为[:i, :k, :j]的目标张量,每个位置的数值完全符合运算规则。如果你需要其他维度顺序,比如[:k, :i, :j]或[:i, :j, :k],直接调用Nx.transpose调整维度排列即可。
内容的提问来源于stack exchange,提问作者RDO
相关产品推荐
相关产品推荐

