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

如何通过切片操作实现相邻图像间的通道特征图传递?

解决方案

你要实现的特征跨帧迁移可以通过以下步骤完成:

  1. 先把输入的时间维度单独拆分出来,方便按时间步索引特征,不要直接把batch和时间步压成一维
  2. 对提取到的1通道特征做时间维度移位,对应把前一帧的特征移到后一帧的位置,第一帧没有前序特征可以用零填充
  3. 将历史特征和当前帧conv1输出的剩余可用通道拼接后,送入后续conv2层即可

注意你原有代码里conv2的输入通道设置为6是匹配需求的:conv1输出8通道,留1个通道传给下一帧,当前帧留5个自用,加上前一帧传过来的1个通道刚好凑齐6个输入通道。下面是修改后的可运行代码:

import cv2
import gym
import numpy as np
import matplotlib.pyplot as plt
import torch
import torch.nn as nn
import torch.optim as optim
import torch.autograd as autograd
import torch.nn.functional as F

N = 1 # Batch Size
T = 5 # Time Steps. This means that there are 5 frames in the video
C = 3 # RGB Channels
H = 144 # Height
W = 144 # Width

foo = torch.randn(N*T, C, H, W)
# 拆分出时间维度,形状变为 [N, T, C, H, W]
foo = foo.reshape(N, T, C, H, W)

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 8, 5)
        self.pool = nn.MaxPool2d(2, 2)
        # 输入通道为当前帧的5个通道 + 前一帧传过来的1个通道,共6个,和原有设置一致
        self.conv2 = nn.Conv2d(6, 16, 5)
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        N, T, C, H, W = x.shape
        # 先把所有帧过一遍conv1,形状变为 [N, T, 8, 140, 140]
        x = x.reshape(N*T, C, H, W)
        conv1_out = F.relu(self.conv1(x))
        conv1_out = conv1_out.reshape(N, T, 8, 140, 140)
        
        # 提取要迁移的1/8通道,形状[N, T, 1, 140, 140]
        transfer_feat = conv1_out[:, :, :1, :, :]
        # 构造前一帧的特征:整体后移一位,第一帧补0
        prev_transfer = torch.zeros_like(transfer_feat)
        # 从第2帧开始,赋值为前一帧的迁移特征
        prev_transfer[:, 1:, :, :, :] = transfer_feat[:, :-1, :, :, :]
        
        # 取当前帧conv1输出的后5个通道,加上前一帧的迁移特征,共6个通道,匹配conv2输入
        curr_feat = conv1_out[:, :, 1:6, :, :]
        conv2_input = torch.cat([curr_feat, prev_transfer], dim=2)
        # 压平时间维度过conv2
        conv2_input = conv2_input.reshape(N*T, 6, 140, 140)
        
        # 后续网络逻辑按原有需求补全即可
        x = self.pool(F.relu(self.conv2(conv2_input)))
        
        return x, transfer_feat

net = Net()
output, transfer_feat = net(foo)
print("输出特征形状:", output.shape)
print("迁移特征形状:", transfer_feat.shape)

如果你不想用索引赋值的方式,也可以用torch.roll实现特征移位,代码更简洁:

# 等价的移位实现,不需要手动逐帧赋值
prev_transfer = torch.roll(transfer_feat, shifts=1, dims=1)
prev_transfer[:, 0, :, :, :] = 0

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 10:00:01