如何在PyTorch中实现类似scipy.linalg.circulant的循环矩阵生成功能
PyTorch循环矩阵生成函数实现
功能说明
实现与scipy.linalg.circulant逻辑完全一致的循环矩阵生成能力,输入为1维torch张量,输出对应2维循环矩阵,可直接嵌入深度学习模型中,用于降低全连接层的过参数化问题。
实现代码
import torch def torch_circulant(c: torch.Tensor) -> torch.Tensor: """ 生成循环矩阵 Args: c: 1维torch张量,作为循环矩阵的第一列 Returns: 2维循环矩阵,形状为 [n, n],n为输入c的长度 """ n = c.size(0) # 生成与输入张量同设备的行、列索引 row_idx = torch.arange(n, device=c.device)[:, None] col_idx = torch.arange(n, device=c.device)[None, :] # 计算循环偏移索引,匹配循环矩阵的元素取值规则 idx = (row_idx - col_idx) % n return c[idx]
使用示例
# 输入1维张量,开启自动微分适配模型训练需求 c = torch.tensor([1, 2, 3, 4], dtype=torch.float32, requires_grad=True) circulant_mat = torch_circulant(c) print(circulant_mat)
输出结果:
tensor([[1., 4., 3., 2.], [2., 1., 4., 3.], [3., 2., 1., 4.], [4., 3., 2., 1.]], grad_fn=<IndexBackward0>)
注意事项
- 该实现原生支持PyTorch自动微分、GPU加速,可直接作为模型组件使用
- 用于全连接层参数压缩时,仅需存储长度为n的1维向量即可替代原本n²规模的全连接权重,参数压缩比例为1/n,符合低参数化设计需求
- 可通过和scipy的circulant输出做数值对比验证一致性,无精度损失
内容的提问来源于stack exchange,提问作者Mevan Ekanayake
相关产品推荐
相关产品推荐

