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

如何根据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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 22:48:52