PyTorch实现DDPM-UNet代码报错修复求助:NotImplementedError问题
UNet-DDPM代码报错定位与修正
报错信息
NotImplementedError Traceback (most recent call last) <ipython-input-21-36a7bf177d9a> in <cell line: 60>() 62 t=torch.randint(0,1000,(3,)) 63 model=Unet(1000,128) ---> 64 y=model(x,t) 65 print(y.shape) 5 frames /usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py in _forward_unimplemented(self, *input) 350 registered hooks while the latter silently ignores them. 351 """ ---> 352 raise NotImplementedError(f'Module [{type(self).__name__}] is missing the required "forward" function') 353 354 NotImplementedError: Module [ModuleList] is missing the required "forward" function
原有UNet代码
import torch import torch.nn as nn # 假设依赖类已实现:ConvBnSiLu, EncoderBlock, DecoderBlock, ResidualBottleneck class ConvBnSiLu(nn.Module): def __init__(self, in_channels, out_channels, kernel_size, stride, padding): super().__init__() self.conv = nn.Conv2d(in_channels, out_channels, kernel_size, stride, padding) self.bn = nn.BatchNorm2d(out_channels) self.silu = nn.SiLU() def forward(self, x): return self.silu(self.bn(self.conv(x))) class EncoderBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.res1 = ResidualBottleneck(in_channels, out_channels, time_emb_dim) self.downsample = nn.Conv2d(out_channels, out_channels, 3, 2, 1) def forward(self, x, t_emb): x = self.res1(x, t_emb) skip = x x = self.downsample(x) return x, skip class DecoderBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.upsample = nn.ConvTranspose2d(in_channels, out_channels, 2, 2) self.res1 = ResidualBottleneck(out_channels*2, out_channels, time_emb_dim) def forward(self, x, skip, t_emb): x = self.upsample(x) x = torch.cat([x, skip], dim=1) x = self.res1(x, t_emb) return x class ResidualBottleneck(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.silu = nn.SiLU() self.bn1 = nn.BatchNorm2d(in_channels) self.conv1 = nn.Conv2d(in_channels, out_channels, 3, 1, 1) self.bn2 = nn.BatchNorm2d(out_channels) self.conv2 = nn.Conv2d(out_channels, out_channels, 3, 1, 1) self.time_mlp = nn.Sequential(nn.SiLU(), nn.Linear(time_emb_dim, out_channels)) self.residual_conv = nn.Conv2d(in_channels, out_channels, 1) if in_channels != out_channels else nn.Identity() def forward(self, x, t_emb): h = self.silu(self.bn1(x)) h = self.conv1(h) t_emb = self.time_mlp(t_emb)[:, :, None, None] h += t_emb h = self.silu(self.bn2(h)) h = self.conv2(h) return h + self.residual_conv(x) class Unet(nn.Module): def __init__(self,timesteps,time_embedding_dim,in_channels=3,out_channels=2,base_dim=32,dim_mults=[2,4,8,16]): super().__init__() assert isinstance(dim_mults,(list,tuple)) assert base_dim%2==0 channels=self._cal_channels(base_dim,dim_mults) self.init_conv=ConvBnSiLu(in_channels,base_dim,3,1,1) self.time_embedding=nn.Embedding(timesteps,time_embedding_dim) self.encoder_blocks=nn.ModuleList([EncoderBlock(c[0],c[1],time_embedding_dim) for c in channels]) self.decoder_blocks=nn.ModuleList([DecoderBlock(c[1],c[0],time_embedding_dim) for c in channels[::-1]]) self.mid_block=nn.Sequential(*[ResidualBottleneck(channels[-1][1],channels[-1][1]) for i in range(2)], ResidualBottleneck(channels[-1][1],channels[-1][1]//2)) self.final_conv=nn.Conv2d(in_channels=channels[0][0]//2,out_channels=out_channels,kernel_size=1) def forward(self,x,t=None): ''' Implement the data flow of the UNet architecture ''' # ---------- **** ---------- # # YOUR CODE HERE t = self.time_embedding #initial conv x1 = self.init_conv(x) #Down x2 = self.encoder_blocks(x1,t) x3 = self.encoder_blocks(x2[0],t) x4 = self.encoder_blocks(x3[0],t) x5 = self.encoder_blocks(x4[0],t) #Middle x6 = self.mid_block(x5[0]) #Up x = self.decoder_blocks(x6,x5[1],t) x = self.decoder_blocks(x,x4[1],t) x = self.decoder_blocks(x,x3[1],t) x = self.decoder_blocks(x,x2[1],t) x = self.decoder_blocks(x,x1[1],t) #final x = self.final_conv(x) # ---------- **** ---------- # return x def _cal_channels(self,base_dim,dim_mults): dims=[base_dim*x for x in dim_mults] dims.insert(0,base_dim) channels=[] for i in range(len(dims)-1): channels.append((dims[i],dims[i+1])) # in_channel, out_channel return channels if __name__=="__main__": x=torch.randn(3,3,224,224) t=torch.randint(0,1000,(3,)) model=Unet(1000,128) y=model(x,t) print(y.shape)
问题定位
- ModuleList调用错误:
encoder_blocks和decoder_blocks是nn.ModuleList容器,本身无forward方法,不能直接像函数调用,必须通过索引或遍历访问内部模块。 - 时间嵌入未编码:
t = self.time_embedding仅获取Embedding层对象,未传入时间步t生成嵌入向量,也未做维度映射和广播,无法和特征图融合。 - 编解码流程错误:重复调用整个
encoder_blocks列表,正确流程是依次调用每个EncoderBlock,保存每阶段的跳连接特征。 - 通道不匹配:MNIST是单通道灰度图,原有代码默认
in_channels=3,需改为1;DDPM输出为噪声预测,通道数应与输入一致(out_channels=1)。
修正后的核心代码
调整Unet类的初始化与forward方法
class Unet(nn.Module): def __init__(self,timesteps,time_embedding_dim,in_channels=1,out_channels=1,base_dim=32,dim_mults=[2,4,8,16]): super().__init__() assert isinstance(dim_mults,(list,tuple)) assert base_dim%2==0 channels=self._cal_channels(base_dim,dim_mults) self.init_conv=ConvBnSiLu(in_channels,base_dim,3,1,1) # 新增时间嵌入MLP映射 self.time_embedding=nn.Embedding(timesteps,time_embedding_dim) self.time_mlp = nn.Sequential( nn.Linear(time_embedding_dim, time_embedding_dim*4), nn.SiLU(), nn.Linear(time_embedding_dim*4, time_embedding_dim*4) ) self.encoder_blocks=nn.ModuleList([EncoderBlock(c[0],c[1],time_embedding_dim*4) for c in channels]) self.decoder_blocks=nn.ModuleList([DecoderBlock(c[1],c[0],time_embedding_dim*4) for c in channels[::-1]]) # 把mid_block改为ModuleList,支持传入t_emb self.mid_block=nn.ModuleList([ ResidualBottleneck(channels[-1][1],channels[-1][1],time_embedding_dim*4), ResidualBottleneck(channels[-1][1],channels[-1][1],time_embedding_dim*4), ResidualBottleneck(channels[-1][1],channels[-1][1]//2,time_embedding_dim*4) ]) self.final_conv=nn.Conv2d(in_channels=channels[0][0]//2,out_channels=out_channels,kernel_size=1) def forward(self,x,t=None): # 1. 处理时间嵌入 t_emb = self.time_embedding(t) t_emb = self.time_mlp(t_emb) # 2. 初始卷积 x = self.init_conv(x) skip_connections = [x] # 3. 编码器下采样,保存跳连接 for block in self.encoder_blocks: x, skip = block(x, t_emb) skip_connections.append(skip) # 4. 中间瓶颈层 for block in self.mid_block: x = block(x, t_emb) # 5. 解码器上采样,拼接跳连接 skip_connections = skip_connections[::-1][1:] for block, skip in zip(self.decoder_blocks, skip_connections): x = block(x, skip, t_emb) # 6. 最终输出 x = self.final_conv(x) return x if __name__=="__main__": x=torch.randn(3,1,28,28) t=torch.randint(0,1000,(3,)) model=Unet(1000,128) y=model(x,t) print(y.shape) # 预期输出: torch.Size([3, 1, 28, 28])
UNet-DDPM正确实现逻辑
- 时间嵌入融合:离散时间步通过Embedding转向量,再用MLP映射到高维度,广播到特征图空间维度,实现时间信息与空间特征的融合。
- 编码器下采样:每个EncoderBlock完成残差提取+下采样,保存每阶段特征作为跳连接,逐步降低特征图尺寸、提升通道数。
- 中间瓶颈层:针对最深层特征做密集处理,提取全局信息。
- 解码器上采样:每个DecoderBlock先上采样恢复尺寸,再拼接对应编码器的跳连接特征,通过残差块融合信息,逐步提升尺寸、降低通道数。
- 最终输出:用1x1卷积将特征映射到目标通道数(MNIST场景为1,对应预测噪声)。
内容的提问来源于stack exchange,提问作者Daniel
相关产品推荐
相关产品推荐

