如何根据forward函数中输入x的维度大小设置encoder的t参数?
根据输入张量维度动态初始化Encoder参数的实现方法
核心思路:由于
__init__方法执行时还未获取到输入张量x,无法直接确定t参数,因此需要延迟初始化encoder,在第一次调用forward时根据x的维度创建encoder实例。修改后的代码实现:
import torch.nn as nn from my_folder import encoder class my_class(nn.Module): def __init__(self, in_channel=256): super(my_class, self).__init__() # 暂不初始化encoder,延迟到forward中处理 self.encoder = None self.in_channel = in_channel # 预先定义encoder的固定参数 self.fixed_h = 4 self.fixed_w = 6 self.fixed_patch_t = 2 def forward(self, x): # 提取x的第2维大小作为encoder的t参数 input_t = x.shape[2] # 仅在第一次forward时初始化encoder if self.encoder is None: self.encoder = encoder(t=input_t, h=self.fixed_h, w=self.fixed_w, patch_t=self.fixed_patch_t) # 将encoder参数同步到输入x所在设备(GPU/CPU) self.encoder = self.encoder.to(x.device) # 执行编码逻辑 y = self.encoder(x) return y
- 关键细节说明:
- 避免重复初始化:通过
self.encoder is None的判断,确保encoder仅在第一次调用forward时创建,后续复用同一实例,保证训练时参数可以正常更新。 - 设备同步:使用
to(x.device)确保encoder参数和输入x处于同一计算设备,避免张量设备不匹配的报错。 - 限制条件:此方案适用于训练过程中
x的第2维大小固定的场景;若需要支持动态变化的t,则需修改encoder类本身,使其能在forward中动态处理不同的t值(而非初始化时固定)。
- 避免重复初始化:通过
内容的提问来源于stack exchange,提问作者dtr43
相关产品推荐
相关产品推荐

