PyTorch Forward Hook异常:HF UNet下采样块触发顺序错误排查
问题描述
- 基于Hugging Face的UNet修改扩散模型,在Down Block和Up Block中添加条件分支
- 运行复现代码时出现两个异常:
- Forward Hook触发逻辑异常:最后一个Down Block的Hook先于第一个触发,且
ConditionResnet的call_count显示仅最后一个Down Block被调用(结果为([0, 0, 0, 1], [0, 0, 0, 0])) - 通道不匹配错误:RuntimeError提示输入为320通道,但卷积层权重期望1280通道
- Forward Hook触发逻辑异常:最后一个Down Block的Hook先于第一个触发,且
复现代码
import torch import torch.nn as nn import torch.nn.functional as F from diffusers import UNet2DConditionModel # config SD_MODEL = "runwayml/stable-diffusion-v1-5" DIM = 15 unet = UNet2DConditionModel.from_pretrained(SD_MODEL, subfolder="unet") bs = 2 timestep = torch.randint(0, 100, (bs,)) noise = torch.randn((bs, 4, 64, 64)) text_encoding = torch.randn((bs, 77, 768)) condition = torch.randn((bs, DIM)) DownOutput = tuple[torch.Tensor, tuple[torch.Tensor, torch.Tensor, torch.Tensor]] class ConditionResnet(nn.Module): def __init__(self, in_dim, out_dim): super().__init__() self.call_count = 0 self.projector = nn.Linear(in_dim, out_dim) self.conv1 = nn.Conv2d(out_dim, out_dim, kernel_size=3, stride=1, padding=1) self.non_linearity = F.silu def forward(self, out: torch.Tensor, condition: torch.Tensor) -> torch.Tensor: self.call_count += 1 input_vector = out out = self.conv1(out) + self.projector(condition)[:, :, None, None] return input_vector + self.non_linearity(out) # down blocks return tuples, so need slightly modified version class ConditionResnetDown(nn.Module): def __init__(self, in_dim, out_dim): super().__init__() self.condition_resnet = ConditionResnet(in_dim, out_dim) def forward(self, x: DownOutput, condition: torch.Tensor) -> DownOutput: return self.condition_resnet(x[0], condition), x[1] class UNetWithConditions(nn.Module): def __init__(self, unet: nn.Module, col_channels: int, down_block_sizes: list[int], up_block_sizes: list[int]): super().__init__() self.unet = unet self.down_block_condition_resnets = nn.ModuleList([ConditionResnetDown(col_channels, out_channel) for out_channel in down_block_sizes]) self.up_block_condition_resnets = nn.ModuleList([ConditionResnet(col_channels, out_channel) for out_channel in up_block_sizes]) self.condition = None # forward hooks for i in range(len(self.unet.down_blocks)): self.unet.down_blocks[i].register_forward_hook(lambda module, inputs, outputs: self.down_block_condition_resnets[i](outputs, self.condition)) for i in range(len(self.unet.up_blocks)): self.unet.up_blocks[i].register_forward_hook(lambda module, inputs, outputs: self.up_block_condition_resnets[i](outputs, self.condition)) def forward(self, noise, timestep, text_encoding, condition): self.condition = condition out = self.unet(noise, timestep, text_encoding).sample self.condition = None return out unet_with_conditions = UNetWithConditions(unet, DIM, [320, 640, 1280, 1280], [1280, 1280, 640, 320]) out2 = unet_with_conditions(noise, timestep, text_encoding, condition)
错误日志
--------------------------------------------------------------------------- RuntimeError Traceback (most recent call last) ~tmp/ipykernel_574/3305635741.py in <cell line: 2>() 1 unet_with_conditions = UNetWithConditions(unet, DIM, [320, 640, 1280, 1280], [1280, 1280, 640, 320]) ----> 2 out2 = unet_with_conditions(noise, timestep, text_encoding, condition) ~app/creator_content_publish_server/models/template_predict_title_trainer/torch_wrapper_layer.runfiles/pypi_torch/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs) 1192 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks 1193 or _global_forward_hooks or _global_forward_pre_hooks): -> 1194 return forward_call(*input, **kwargs) 1195 # Do not call functions when jit is used 1196 full_backward_hooks, non_full_backward_hooks = [], [] ~tmp/ipykernel_574/2058376754.py in forward(self, noise, timestep, text_encoding, condition) 59 def forward(self, noise, timestep, text_encoding, condition): 60 self.condition = condition ---> 61 out = self.unet(noise, timestep, text_encoding).sample 62 self.condition = None 63 return out ~app/creator_content_publish_server/models/template_predict_title_trainer/torch_wrapper_layer.runfiles/pypi_torch/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs) 1192 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks 1193 or _global_forward_hooks or _global_forward_pre_hooks): -> 1194 return forward_call(*input, **kwargs) 1195 # Do not call functions when jit is used 1196 full_backward_hooks, non_full_backward_hooks = [], [] ~nix/store/vzqny68wq33dcg4hkdala51n5vqhpnwc-python3-3.9.12/lib/python3.9/site-packages/diffusers/models/unet_2d_condition.py in forward(self, sample, timestep, encoder_hidden_states, class_labels, timestep_cond, attention_mask, cross_attention_kwargs, added_cond_kwargs, down_block_additional_residuals, mid_block_additional_residual, encoder_attention_mask, return_dict) 795 for downsample_block in self.down_blocks: 796 if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention: ---> 797 sample, res_samples = downsample_block( 798 hidden_states=sample, 799 temb=emb, ~app/creator_content_publish_server/models/template_predict_title_trainer/torch_wrapper_layer.runfiles/pypi_torch/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs) 1213 if _global_forward_hooks or self._forward_hooks: 1214 for hook in (*_global_forward_hooks.values(), *self._forward_hooks.values()): -> 1215 hook_result = hook(self, input, result) 1216 if hook_result is not None: 1217 result = hook_result ~tmp/ipykernel_574/2058376754.py in <lambda>(module, inputs, outputs) 53 # forward hooks 54 for i in range(len(self.unet.down_blocks)): ---> 55 self.unet.down_blocks[i].register_forward_hook(lambda module, inputs, outputs: self.down_block_condition_resnets[i](outputs, self.condition)) 56 for i in range(len(self.unet.up_blocks)): 57 self.unet.up_blocks[i].register_forward_hook(lambda module, inputs, outputs: self.up_block_condition_resnets[i](outputs, self.condition)) ~app/creator_content_publish_server/models/template_predict_title_trainer/torch_wrapper_layer.runfiles/pypi_torch/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs) 1192 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks 1193 or _global_forward_hooks or _global_forward_pre_hooks): -> 1194 return forward_call(*input, **kwargs) 1195 # Do not call functions when jit is used 1196 full_backward_hooks, non_full_backward_hooks = [], [] ~tmp/ipykernel_574/2058376754.py in forward(self, x, condition) 40 41 def forward(self, x: DownOutput, condition: torch.Tensor) -> DownOutput: ---> 42 return self.condition_resnet(x[0], condition), x[1] 43 44 class UNetWithConditions(nn.Module): ~app/creator_content_publish_server/models/template_predict_title_trainer/torch_wrapper_layer.runfiles/pypi_torch/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs) 1192 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks 1193 or _global_forward_hooks or _global_forward_pre_hooks): -> 1194 return forward_call(*input, **kwargs) 1195 # Do not call functions when jit is used 1196 full_backward_hooks, non_full_backward_hooks = [], [] ~tmp/ipykernel_574/2058376754.py in forward(self, out, condition) 30 self.call_count += 1 31 input_vector = out ---> 32 out = self.conv1(out) + self.projector(condition)[:, :, None, None] 33 return input_vector + self.non_linearity(out) 34 ~app/creator_content_publish_server/models/template_predict_title_trainer/torch_wrapper_layer.runfiles/pypi_torch/site-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs) 1192 if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks 1193 or _global_forward_hooks or _global_forward_pre_hooks): -> 1194 return forward_call(*input, **kwargs) 1195 # Do not call functions when jit is used 1196 full_backward_hooks, non_full_backward_hooks = [], [] ~app/creator_content_publish_server/models/template_predict_title_trainer/torch_wrapper_layer.runfiles/pypi_torch/site-packages/torch/nn/modules/conv.py in forward(self, input) 461 462 def forward(self, input: Tensor) -> Tensor: ---> 463 return self._conv_forward(input, self.weight, self.bias) 464 465 class Conv3d(_ConvNd): ~app/creator_content_publish_server/models/template_predict_title_trainer/torch_wrapper_layer.runfiles/pypi_torch/site-packages/torch/nn/modules/conv.py in _conv_forward(self, input, weight, bias) 457 weight, bias, self.stride, 458 _pair(0), self.dilation, self.groups) ---> 459 return F.conv2d(input, weight, bias, self.stride, 460 self.padding, self.dilation, self.groups) 461 RuntimeError: Given groups=1, weight of size [1280, 1280, 3, 3], expected input[2, 320, 32, 32] to have 1280 channels, but got 320 channels instead
解决方案
1. 修复Lambda闭包变量捕获问题
循环中注册Hook时,lambda表达式会延迟捕获循环变量i,导致所有Hook最终都使用循环结束时的i值(即最后一个索引),这是仅最后一个Down Block被调用的核心原因。
修复方法:用默认参数在定义lambda时捕获当前的i值:
# 替换原Hook注册代码 for i in range(len(self.unet.down_blocks)): self.unet.down_blocks[i].register_forward_hook(lambda module, inputs, outputs, idx=i: self.down_block_condition_resnets[idx](outputs, self.condition)) for i in range(len(self.unet.up_blocks)): self.unet.up_blocks[i].register_forward_hook(lambda module, inputs, outputs, idx=i: self.up_block_condition_resnets[idx](outputs, self.condition))
2. 修复通道不匹配问题
ConditionResnet的conv1输入通道数定义错误:当前代码混淆了条件向量维度和特征张量通道数,conv1的输入通道应与Down/Up Block输出的特征通道数一致,而非条件投影后的维度。
修正ConditionResnet及相关类的实现:
class ConditionResnet(nn.Module): def __init__(self, cond_dim, feat_dim): super().__init__() self.call_count = 0 # 条件向量投影到特征通道数 self.projector = nn.Linear(cond_dim, feat_dim) # 卷积层输入输出均为特征通道数 self.conv1 = nn.Conv2d(feat_dim, feat_dim, kernel_size=3, stride=1, padding=1) self.non_linearity = F.silu def forward(self, out: torch.Tensor, condition: torch.Tensor) -> torch.Tensor: self.call_count += 1 input_vector = out out = self.conv1(out) + self.projector(condition)[:, :, None, None] return input_vector + self.non_linearity(out) class ConditionResnetDown(nn.Module): def __init__(self, cond_dim, feat_dim): super().__init__() self.condition_resnet = ConditionResnet(cond_dim, feat_dim) def forward(self, x: DownOutput, condition: torch.Tensor) -> DownOutput: return self.condition_resnet(x[0], condition), x[1] class UNetWithConditions(nn.Module): def __init__(self, unet: nn.Module, cond_dim: int, down_feat_sizes: list[int], up_feat_sizes: list[int]): super().__init__() self.unet = unet self.down_block_condition_resnets = nn.ModuleList([ConditionResnetDown(cond_dim, feat_size) for feat_size in down_feat_sizes]) self.up_block_condition_resnets = nn.ModuleList([ConditionResnet(cond_dim, feat_size) for feat_size in up_feat_sizes]) self.condition = None # 修复后的Hook注册 for i in range(len(self.unet.down_blocks)): self.unet.down_blocks[i].register_forward_hook(lambda module, inputs, outputs, idx=i: self.down_block_condition_resnets[idx](outputs, self.condition)) for i in range(len(self.unet.up_blocks)): self.unet.up_blocks[i].register_forward_hook(lambda module, inputs, outputs, idx=i: self.up_block_condition_resnets[idx](outputs, self.condition)) def forward(self, noise, timestep, text_encoding, condition): self.condition = condition out = self.unet(noise, timestep, text_encoding).sample self.condition = None return out
3. 关于JIT编译和Hook传参的疑问解答
- 模型未被JIT编译:默认从HF加载的UNet是普通PyTorch模块,不会自动触发JIT编译
- Forward Hook本身不支持直接传入额外输入,但通过类实例变量(如
self.condition)传递参数的方式是可行的,问题根源不在Hook传参
相关产品推荐
相关产品推荐

