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

如何在PyTorch张量中为对角线掩码并填充指定值?

嘿,刚好知道怎么解决这个PyTorch里填充对角线的问题!你提到在NumPy里用np.fill_diagonal很方便,其实PyTorch里也有对应的方案,还有几种灵活的实现方式,我来给你一一讲清楚:

方法1:最简洁的原地填充——fill_diagonal_()

PyTorch的张量对象自带了fill_diagonal_()方法,和NumPy的fill_diagonal用法几乎一致,而且是原地修改张量,非常高效:

import torch
# 创建3x3的零张量
a = torch.zeros((3, 3), dtype=torch.int)
# 填充对角线为5
a.fill_diagonal_(5)
print(a)

输出结果:

tensor([[5, 0, 0],
        [0, 5, 0],
        [0, 0, 5]], dtype=torch.int32)

如果不想原地修改,可以先克隆张量再操作:a.clone().fill_diagonal_(5)。

方法2:用布尔掩码赋值

如果你想通过掩码的方式实现(就像你提到的用对角线作为掩码),可以用torch.eye生成单位矩阵转成布尔掩码,再通过掩码索引赋值:

import torch
a = torch.zeros((3, 3), dtype=torch.int)
# 创建和a同尺寸的布尔掩码,对角线为True,其余为False
mask = torch.eye(a.shape[0], dtype=torch.bool, device=a.device)
# 给掩码对应的位置赋值5
a[mask] = 5
print(a)

这种方法的好处是可以灵活调整掩码,比如填充偏移对角线(比如torch.eye(3, offset=1)就能取右上偏移1的对角线)。

方法3:直接通过索引定位对角线

对角线元素的索引规律是(i, i),我们可以用torch.arange生成索引序列,直接定位对角线位置赋值:

import torch
a = torch.zeros((3, 3), dtype=torch.int)
# 生成0到2的索引序列
idx = torch.arange(a.shape[0])
# 给(i,i)位置赋值
a[idx, idx] = 5
print(a)

这种方式逻辑直观,适合需要自定义索引的场景,比如只填充部分对角线元素。

以上几种方法都能实现你想要的效果,其中fill_diagonal_()是最推荐的,因为它是PyTorch官方优化过的方法,简洁又高效。

内容的提问来源于stack exchange,提问作者NicolaiF

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 04:21:45