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

如何在PyTorch模型中将张量分割为重叠块?解决梯度兼容问题

兼容梯度的PyTorch重叠帧块生成方案

需要将形状为(batch, c, h, w)的图像批量张量,转换为(-1, depth, c, h, w)的张量,其中第i个块包含第i到i+depth帧(重叠块)。原自定义函数因使用numpy()操作导致梯度传播报错(RuntimeError: Can't call numpy() on Tensor that requires grad),以下是纯PyTorch的兼容梯度实现方案:

方法1:使用torch.unfold(推荐,原生滑动窗口API)

torch.unfold是PyTorch专门用于生成滑动窗口的原生函数,无需手动处理索引,且完全支持梯度传播。

代码示例

import torch

# 示例输入:batch=12, c=1, h=1, w=1的张量,对应[1,2,...,12]
x = torch.arange(1, 13, dtype=torch.float32).view(12, 1, 1, 1).requires_grad_(True)
depth = 2

# 在batch维度(dim=0)上创建滑动窗口,窗口大小depth,步长1
x_unfolded = x.unfold(dimension=0, size=depth, step=1)
# 调整维度顺序至目标形状:(-1, depth, c, h, w)
result = x_unfolded.permute(0, 4, 1, 2, 3)

# 验证形状
print(result.shape)  # 输出: torch.Size([11, 2, 1, 1, 1])

原理说明

  • unfold(dim=0, size=depth, step=1)会在batch维度上滑动,每次取连续depth个元素,步长为1(即重叠滑动),输出形状为(batch - depth + 1, c, h, w, depth)
  • permute(0,4,1,2,3)将最后一个维度(窗口内的帧)移到第二个位置,得到目标形状(-1, depth, c, h, w)

方法2:手动生成张量索引(直观易懂)

通过PyTorch原生的张量运算生成重叠块的索引,直接提取对应帧,同样支持梯度传播。

代码示例

import torch

x = torch.arange(1, 13, dtype=torch.float32).view(12, 1, 1, 1).requires_grad_(True)
depth = 2
batch_size = x.shape[0]

# 生成每个重叠块的索引矩阵:形状为(batch - depth +1, depth)
# 每行对应一个块的帧索引,比如[0,1], [1,2], ..., [10,11]
indices = torch.arange(batch_size - depth + 1).unsqueeze(1) + torch.arange(depth)

# 按索引提取帧,直接得到目标形状
result = x[indices]

# 验证形状
print(result.shape)  # 输出: torch.Size([11, 2, 1, 1, 1])

梯度兼容性验证

两种方法均支持反向传播,可通过以下代码验证:

# 计算损失并反向传播
loss = result.sum()
loss.backward()

# 查看输入张量的梯度(每个元素的梯度等于其在重叠块中出现的次数)
print(x.grad.squeeze())
# 输出: tensor([1., 2., 2., 2., 2., 2., 2., 2., 2., 2., 2., 1.]),符合预期

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 11:54:06