multimodalart/EchoMimic-zero
8
1# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/unet_2d_blocks.py2 3import pdb4 5import torch6from torch import nn7 8from .motion_module import get_motion_module9 10# from .motion_module import get_motion_module11from .resnet import Downsample3D, ResnetBlock3D, Upsample3D12from .transformer_3d import Transformer3DModel13 14 15def get_down_block(16 down_block_type,17 num_layers,18 in_channels,19 out_channels,20 temb_channels,21 add_downsample,22 resnet_eps,23 resnet_act_fn,24 attn_num_head_channels,25 resnet_groups=None,26 cross_attention_dim=None,27 downsample_padding=None,28 dual_cross_attention=False,29 use_linear_projection=False,30 only_cross_attention=False,31 upcast_attention=False,32 resnet_time_scale_shift="default",33 unet_use_cross_frame_attention=None,34 unet_use_temporal_attention=None,35 use_inflated_groupnorm=None,36 use_motion_module=None,37 motion_module_type=None,38 motion_module_kwargs=None,39):40 down_block_type = (41 down_block_type[7:]42 if down_block_type.startswith("UNetRes")43 else down_block_type44 )45 if down_block_type == "DownBlock3D":46 return DownBlock3D(47 num_layers=num_layers,48 in_channels=in_channels,49 out_channels=out_channels,50 temb_channels=temb_channels,51 add_downsample=add_downsample,52 resnet_eps=resnet_eps,53 resnet_act_fn=resnet_act_fn,54 resnet_groups=resnet_groups,55 downsample_padding=downsample_padding,56 resnet_time_scale_shift=resnet_time_scale_shift,57 use_inflated_groupnorm=use_inflated_groupnorm,58 use_motion_module=use_motion_module,59 motion_module_type=motion_module_type,60 motion_module_kwargs=motion_module_kwargs,61 )62 elif down_block_type == "CrossAttnDownBlock3D":63 if cross_attention_dim is None:64 raise ValueError(65 "cross_attention_dim must be specified for CrossAttnDownBlock3D"66 )67 return CrossAttnDownBlock3D(68 num_layers=num_layers,69 in_channels=in_channels,70 out_channels=out_channels,71 temb_channels=temb_channels,72 add_downsample=add_downsample,73 resnet_eps=resnet_eps,74 resnet_act_fn=resnet_act_fn,75 resnet_groups=resnet_groups,76 downsample_padding=downsample_padding,77 cross_attention_dim=cross_attention_dim,78 attn_num_head_channels=attn_num_head_channels,79 dual_cross_attention=dual_cross_attention,80 use_linear_projection=use_linear_projection,81 only_cross_attention=only_cross_attention,82 upcast_attention=upcast_attention,83 resnet_time_scale_shift=resnet_time_scale_shift,84 unet_use_cross_frame_attention=unet_use_cross_frame_attention,85 unet_use_temporal_attention=unet_use_temporal_attention,86 use_inflated_groupnorm=use_inflated_groupnorm,87 use_motion_module=use_motion_module,88 motion_module_type=motion_module_type,89 motion_module_kwargs=motion_module_kwargs,90 )91 raise ValueError(f"{down_block_type} does not exist.")92 93 94def get_up_block(95 up_block_type,96 num_layers,97 in_channels,98 out_channels,99 prev_output_channel,100 temb_channels,101 add_upsample,102 resnet_eps,103 resnet_act_fn,104 attn_num_head_channels,105 resnet_groups=None,106 cross_attention_dim=None,107 dual_cross_attention=False,108 use_linear_projection=False,109 only_cross_attention=False,110 upcast_attention=False,111 resnet_time_scale_shift="default",112 unet_use_cross_frame_attention=None,113 unet_use_temporal_attention=None,114 use_inflated_groupnorm=None,115 use_motion_module=None,116 motion_module_type=None,117 motion_module_kwargs=None,118):119 up_block_type = (120 up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type121 )122 if up_block_type == "UpBlock3D":123 return UpBlock3D(124 num_layers=num_layers,125 in_channels=in_channels,126 out_channels=out_channels,127 prev_output_channel=prev_output_channel,128 temb_channels=temb_channels,129 add_upsample=add_upsample,130 resnet_eps=resnet_eps,131 resnet_act_fn=resnet_act_fn,132 resnet_groups=resnet_groups,133 resnet_time_scale_shift=resnet_time_scale_shift,134 use_inflated_groupnorm=use_inflated_groupnorm,135 use_motion_module=use_motion_module,136 motion_module_type=motion_module_type,137 motion_module_kwargs=motion_module_kwargs,138 )139 elif up_block_type == "CrossAttnUpBlock3D":140 if cross_attention_dim is None:141 raise ValueError(142 "cross_attention_dim must be specified for CrossAttnUpBlock3D"143 )144 return CrossAttnUpBlock3D(145 num_layers=num_layers,146 in_channels=in_channels,147 out_channels=out_channels,148 prev_output_channel=prev_output_channel,149 temb_channels=temb_channels,150 add_upsample=add_upsample,151 resnet_eps=resnet_eps,152 resnet_act_fn=resnet_act_fn,153 resnet_groups=resnet_groups,154 cross_attention_dim=cross_attention_dim,155 attn_num_head_channels=attn_num_head_channels,156 dual_cross_attention=dual_cross_attention,157 use_linear_projection=use_linear_projection,158 only_cross_attention=only_cross_attention,159 upcast_attention=upcast_attention,160 resnet_time_scale_shift=resnet_time_scale_shift,161 unet_use_cross_frame_attention=unet_use_cross_frame_attention,162 unet_use_temporal_attention=unet_use_temporal_attention,163 use_inflated_groupnorm=use_inflated_groupnorm,164 use_motion_module=use_motion_module,165 motion_module_type=motion_module_type,166 motion_module_kwargs=motion_module_kwargs,167 )168 raise ValueError(f"{up_block_type} does not exist.")169 170 171class UNetMidBlock3DCrossAttn(nn.Module):172 def __init__(173 self,174 in_channels: int,175 temb_channels: int,176 dropout: float = 0.0,177 num_layers: int = 1,178 resnet_eps: float = 1e-6,179 resnet_time_scale_shift: str = "default",180 resnet_act_fn: str = "swish",181 resnet_groups: int = 32,182 resnet_pre_norm: bool = True,183 attn_num_head_channels=1,184 output_scale_factor=1.0,185 cross_attention_dim=1280,186 dual_cross_attention=False,187 use_linear_projection=False,188 upcast_attention=False,189 unet_use_cross_frame_attention=None,190 unet_use_temporal_attention=None,191 use_inflated_groupnorm=None,192 use_motion_module=None,193 motion_module_type=None,194 motion_module_kwargs=None,195 ):196 super().__init__()197 198 self.has_cross_attention = True199 self.attn_num_head_channels = attn_num_head_channels200 resnet_groups = (201 resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)202 )203 204 # there is always at least one resnet205 resnets = [206 ResnetBlock3D(207 in_channels=in_channels,208 out_channels=in_channels,209 temb_channels=temb_channels,210 eps=resnet_eps,211 groups=resnet_groups,212 dropout=dropout,213 time_embedding_norm=resnet_time_scale_shift,214 non_linearity=resnet_act_fn,215 output_scale_factor=output_scale_factor,216 pre_norm=resnet_pre_norm,217 use_inflated_groupnorm=use_inflated_groupnorm,218 )219 ]220 attentions = []221 motion_modules = []222 223 for _ in range(num_layers):224 if dual_cross_attention:225 raise NotImplementedError226 attentions.append(227 Transformer3DModel(228 attn_num_head_channels,229 in_channels // attn_num_head_channels,230 in_channels=in_channels,231 num_layers=1,232 cross_attention_dim=cross_attention_dim,233 norm_num_groups=resnet_groups,234 use_linear_projection=use_linear_projection,235 upcast_attention=upcast_attention,236 unet_use_cross_frame_attention=unet_use_cross_frame_attention,237 unet_use_temporal_attention=unet_use_temporal_attention,238 )239 )240 motion_modules.append(241 get_motion_module(242 in_channels=in_channels,243 motion_module_type=motion_module_type,244 motion_module_kwargs=motion_module_kwargs,245 )246 if use_motion_module247 else None248 )249 resnets.append(250 ResnetBlock3D(251 in_channels=in_channels,252 out_channels=in_channels,253 temb_channels=temb_channels,254 eps=resnet_eps,255 groups=resnet_groups,256 dropout=dropout,257 time_embedding_norm=resnet_time_scale_shift,258 non_linearity=resnet_act_fn,259 output_scale_factor=output_scale_factor,260 pre_norm=resnet_pre_norm,261 use_inflated_groupnorm=use_inflated_groupnorm,262 )263 )264 265 self.attentions = nn.ModuleList(attentions)266 self.resnets = nn.ModuleList(resnets)267 self.motion_modules = nn.ModuleList(motion_modules)268 269 def forward(270 self,271 hidden_states,272 temb=None,273 encoder_hidden_states=None,274 audio_cond_fea = None,275 attention_mask=None,276 ):277 hidden_states = self.resnets[0](hidden_states, temb)278 for attn, resnet, motion_module in zip(279 self.attentions, self.resnets[1:], self.motion_modules280 ):281 hidden_states = attn(282 hidden_states,283 encoder_hidden_states=encoder_hidden_states,284 audio_cond_fea = audio_cond_fea285 ).sample286 hidden_states = (287 motion_module(288 hidden_states, temb, encoder_hidden_states=encoder_hidden_states289 )290 if motion_module is not None291 else hidden_states292 )293 hidden_states = resnet(hidden_states, temb)294 295 return hidden_states296 297 298class CrossAttnDownBlock3D(nn.Module):299 def __init__(300 self,301 in_channels: int,302 out_channels: int,303 temb_channels: int,304 dropout: float = 0.0,305 num_layers: int = 1,306 resnet_eps: float = 1e-6,307 resnet_time_scale_shift: str = "default",308 resnet_act_fn: str = "swish",309 resnet_groups: int = 32,310 resnet_pre_norm: bool = True,311 attn_num_head_channels=1,312 cross_attention_dim=1280,313 output_scale_factor=1.0,314 downsample_padding=1,315 add_downsample=True,316 dual_cross_attention=False,317 use_linear_projection=False,318 only_cross_attention=False,319 upcast_attention=False,320 unet_use_cross_frame_attention=None,321 unet_use_temporal_attention=None,322 use_inflated_groupnorm=None,323 use_motion_module=None,324 motion_module_type=None,325 motion_module_kwargs=None,326 ):327 super().__init__()328 resnets = []329 attentions = []330 motion_modules = []331 332 self.has_cross_attention = True333 self.attn_num_head_channels = attn_num_head_channels334 335 for i in range(num_layers):336 in_channels = in_channels if i == 0 else out_channels337 resnets.append(338 ResnetBlock3D(339 in_channels=in_channels,340 out_channels=out_channels,341 temb_channels=temb_channels,342 eps=resnet_eps,343 groups=resnet_groups,344 dropout=dropout,345 time_embedding_norm=resnet_time_scale_shift,346 non_linearity=resnet_act_fn,347 output_scale_factor=output_scale_factor,348 pre_norm=resnet_pre_norm,349 use_inflated_groupnorm=use_inflated_groupnorm,350 )351 )352 if dual_cross_attention:353 raise NotImplementedError354 attentions.append(355 Transformer3DModel(356 attn_num_head_channels,357 out_channels // attn_num_head_channels,358 in_channels=out_channels,359 num_layers=1,360 cross_attention_dim=cross_attention_dim,361 norm_num_groups=resnet_groups,362 use_linear_projection=use_linear_projection,363 only_cross_attention=only_cross_attention,364 upcast_attention=upcast_attention,365 unet_use_cross_frame_attention=unet_use_cross_frame_attention,366 unet_use_temporal_attention=unet_use_temporal_attention,367 )368 )369 motion_modules.append(370 get_motion_module(371 in_channels=out_channels,372 motion_module_type=motion_module_type,373 motion_module_kwargs=motion_module_kwargs,374 )375 if use_motion_module376 else None377 )378 379 self.attentions = nn.ModuleList(attentions)380 self.resnets = nn.ModuleList(resnets)381 self.motion_modules = nn.ModuleList(motion_modules)382 383 if add_downsample:384 self.downsamplers = nn.ModuleList(385 [386 Downsample3D(387 out_channels,388 use_conv=True,389 out_channels=out_channels,390 padding=downsample_padding,391 name="op",392 )393 ]394 )395 else:396 self.downsamplers = None397 398 self.gradient_checkpointing = False399 400 def forward(401 self,402 hidden_states,403 temb=None,404 encoder_hidden_states=None,405 audio_cond_fea=None,406 attention_mask=None,407 ):408 output_states = ()409 410 for i, (resnet, attn, motion_module) in enumerate(411 zip(self.resnets, self.attentions, self.motion_modules)412 ):413 # self.gradient_checkpointing = False414 if self.training and self.gradient_checkpointing:415 416 def create_custom_forward(module, return_dict=None):417 def custom_forward(*inputs):418 if return_dict is not None:419 return module(*inputs, return_dict=return_dict)420 else:421 return module(*inputs)422 423 return custom_forward424 425 hidden_states = torch.utils.checkpoint.checkpoint(426 create_custom_forward(resnet), hidden_states, temb427 )428 hidden_states = torch.utils.checkpoint.checkpoint(429 create_custom_forward(attn, return_dict=False),430 hidden_states,431 encoder_hidden_states,432 audio_cond_fea,433 )[0]434 435 # add motion module436 hidden_states = (437 motion_module(438 hidden_states, temb, encoder_hidden_states=encoder_hidden_states439 )440 if motion_module is not None441 else hidden_states442 )443 444 else:445 hidden_states = resnet(hidden_states, temb)446 hidden_states = attn(447 hidden_states,448 encoder_hidden_states=encoder_hidden_states,449 audio_cond_fea = audio_cond_fea,450 ).sample451 452 # add motion module453 hidden_states = (454 motion_module(455 hidden_states, temb, encoder_hidden_states=encoder_hidden_states456 )457 if motion_module is not None458 else hidden_states459 )460 461 output_states += (hidden_states,)462 463 if self.downsamplers is not None:464 for downsampler in self.downsamplers:465 hidden_states = downsampler(hidden_states)466 467 output_states += (hidden_states,)468 469 return hidden_states, output_states470 471 472class DownBlock3D(nn.Module):473 def __init__(474 self,475 in_channels: int,476 out_channels: int,477 temb_channels: int,478 dropout: float = 0.0,479 num_layers: int = 1,480 resnet_eps: float = 1e-6,481 resnet_time_scale_shift: str = "default",482 resnet_act_fn: str = "swish",483 resnet_groups: int = 32,484 resnet_pre_norm: bool = True,485 output_scale_factor=1.0,486 add_downsample=True,487 downsample_padding=1,488 use_inflated_groupnorm=None,489 use_motion_module=None,490 motion_module_type=None,491 motion_module_kwargs=None,492 ):493 super().__init__()494 resnets = []495 motion_modules = []496 497 # use_motion_module = False498 for i in range(num_layers):499 in_channels = in_channels if i == 0 else out_channels500 resnets.append(501 ResnetBlock3D(502 in_channels=in_channels,503 out_channels=out_channels,504 temb_channels=temb_channels,505 eps=resnet_eps,506 groups=resnet_groups,507 dropout=dropout,508 time_embedding_norm=resnet_time_scale_shift,509 non_linearity=resnet_act_fn,510 output_scale_factor=output_scale_factor,511 pre_norm=resnet_pre_norm,512 use_inflated_groupnorm=use_inflated_groupnorm,513 )514 )515 motion_modules.append(516 get_motion_module(517 in_channels=out_channels,518 motion_module_type=motion_module_type,519 motion_module_kwargs=motion_module_kwargs,520 )521 if use_motion_module522 else None523 )524 525 self.resnets = nn.ModuleList(resnets)526 self.motion_modules = nn.ModuleList(motion_modules)527 528 if add_downsample:529 self.downsamplers = nn.ModuleList(530 [531 Downsample3D(532 out_channels,533 use_conv=True,534 out_channels=out_channels,535 padding=downsample_padding,536 name="op",537 )538 ]539 )540 else:541 self.downsamplers = None542 543 self.gradient_checkpointing = False544 545 def forward(self, hidden_states, temb=None, encoder_hidden_states=None):546 output_states = ()547 548 for resnet, motion_module in zip(self.resnets, self.motion_modules):549 # print(f"DownBlock3D {self.gradient_checkpointing = }")550 if self.training and self.gradient_checkpointing:551 552 def create_custom_forward(module):553 def custom_forward(*inputs):554 return module(*inputs)555 556 return custom_forward557 558 hidden_states = torch.utils.checkpoint.checkpoint(559 create_custom_forward(resnet), hidden_states, temb560 )561 if motion_module is not None:562 hidden_states = torch.utils.checkpoint.checkpoint(563 create_custom_forward(motion_module),564 hidden_states.requires_grad_(),565 temb,566 encoder_hidden_states,567 )568 else:569 hidden_states = resnet(hidden_states, temb)570 571 # add motion module572 hidden_states = (573 motion_module(574 hidden_states, temb, encoder_hidden_states=encoder_hidden_states575 )576 if motion_module is not None577 else hidden_states578 )579 580 output_states += (hidden_states,)581 582 if self.downsamplers is not None:583 for downsampler in self.downsamplers:584 hidden_states = downsampler(hidden_states)585 586 output_states += (hidden_states,)587 588 return hidden_states, output_states589 590 591class CrossAttnUpBlock3D(nn.Module):592 def __init__(593 self,594 in_channels: int,595 out_channels: int,596 prev_output_channel: int,597 temb_channels: int,598 dropout: float = 0.0,599 num_layers: int = 1,600 resnet_eps: float = 1e-6,601 resnet_time_scale_shift: str = "default",602 resnet_act_fn: str = "swish",603 resnet_groups: int = 32,604 resnet_pre_norm: bool = True,605 attn_num_head_channels=1,606 cross_attention_dim=1280,607 output_scale_factor=1.0,608 add_upsample=True,609 dual_cross_attention=False,610 use_linear_projection=False,611 only_cross_attention=False,612 upcast_attention=False,613 unet_use_cross_frame_attention=None,614 unet_use_temporal_attention=None,615 use_motion_module=None,616 use_inflated_groupnorm=None,617 motion_module_type=None,618 motion_module_kwargs=None,619 ):620 super().__init__()621 resnets = []622 attentions = []623 motion_modules = []624 625 self.has_cross_attention = True626 self.attn_num_head_channels = attn_num_head_channels627 628 for i in range(num_layers):629 res_skip_channels = in_channels if (i == num_layers - 1) else out_channels630 resnet_in_channels = prev_output_channel if i == 0 else out_channels631 632 resnets.append(633 ResnetBlock3D(634 in_channels=resnet_in_channels + res_skip_channels,635 out_channels=out_channels,636 temb_channels=temb_channels,637 eps=resnet_eps,638 groups=resnet_groups,639 dropout=dropout,640 time_embedding_norm=resnet_time_scale_shift,641 non_linearity=resnet_act_fn,642 output_scale_factor=output_scale_factor,643 pre_norm=resnet_pre_norm,644 use_inflated_groupnorm=use_inflated_groupnorm,645 )646 )647 if dual_cross_attention:648 raise NotImplementedError649 attentions.append(650 Transformer3DModel(651 attn_num_head_channels,652 out_channels // attn_num_head_channels,653 in_channels=out_channels,654 num_layers=1,655 cross_attention_dim=cross_attention_dim,656 norm_num_groups=resnet_groups,657 use_linear_projection=use_linear_projection,658 only_cross_attention=only_cross_attention,659 upcast_attention=upcast_attention,660 unet_use_cross_frame_attention=unet_use_cross_frame_attention,661 unet_use_temporal_attention=unet_use_temporal_attention,662 )663 )664 motion_modules.append(665 get_motion_module(666 in_channels=out_channels,667 motion_module_type=motion_module_type,668 motion_module_kwargs=motion_module_kwargs,669 )670 if use_motion_module671 else None672 )673 674 self.attentions = nn.ModuleList(attentions)675 self.resnets = nn.ModuleList(resnets)676 self.motion_modules = nn.ModuleList(motion_modules)677 678 if add_upsample:679 self.upsamplers = nn.ModuleList(680 [Upsample3D(out_channels, use_conv=True, out_channels=out_channels)]681 )682 else:683 self.upsamplers = None684 685 self.gradient_checkpointing = False686 687 def forward(688 self,689 hidden_states,690 res_hidden_states_tuple,691 temb=None,692 encoder_hidden_states=None,693 audio_cond_fea=None,694 upsample_size=None,695 attention_mask=None,696 ):697 for i, (resnet, attn, motion_module) in enumerate(698 zip(self.resnets, self.attentions, self.motion_modules)699 ):700 # pop res hidden states701 res_hidden_states = res_hidden_states_tuple[-1]702 res_hidden_states_tuple = res_hidden_states_tuple[:-1]703 hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)704 705 if self.training and self.gradient_checkpointing:706 707 def create_custom_forward(module, return_dict=None):708 def custom_forward(*inputs):709 if return_dict is not None:710 return module(*inputs, return_dict=return_dict)711 else:712 return module(*inputs)713 714 return custom_forward715 716 hidden_states = torch.utils.checkpoint.checkpoint(717 create_custom_forward(resnet), hidden_states, temb718 )719 hidden_states = attn(720 hidden_states,721 encoder_hidden_states=encoder_hidden_states,722 audio_cond_fea = audio_cond_fea723 ).sample724 if motion_module is not None:725 # print("hidden_states shape1:", hidden_states.shape)726 hidden_states = torch.utils.checkpoint.checkpoint(727 create_custom_forward(motion_module),728 hidden_states.requires_grad_(),729 temb,730 encoder_hidden_states,731 )732 733 else:734 hidden_states = resnet(hidden_states, temb)735 hidden_states = attn(736 hidden_states,737 encoder_hidden_states=encoder_hidden_states,738 audio_cond_fea = audio_cond_fea739 ).sample740 # print("hidden_states shape1:", hidden_states.shape)741 # add motion module742 hidden_states = (743 motion_module(744 hidden_states, temb, encoder_hidden_states=encoder_hidden_states745 )746 if motion_module is not None747 else hidden_states748 )749 750 if self.upsamplers is not None:751 for upsampler in self.upsamplers:752 hidden_states = upsampler(hidden_states, upsample_size)753 754 return hidden_states755 756 757class UpBlock3D(nn.Module):758 def __init__(759 self,760 in_channels: int,761 prev_output_channel: int,762 out_channels: int,763 temb_channels: int,764 dropout: float = 0.0,765 num_layers: int = 1,766 resnet_eps: float = 1e-6,767 resnet_time_scale_shift: str = "default",768 resnet_act_fn: str = "swish",769 resnet_groups: int = 32,770 resnet_pre_norm: bool = True,771 output_scale_factor=1.0,772 add_upsample=True,773 use_inflated_groupnorm=None,774 use_motion_module=None,775 motion_module_type=None,776 motion_module_kwargs=None,777 ):778 super().__init__()779 resnets = []780 motion_modules = []781 782 # use_motion_module = False783 for i in range(num_layers):784 res_skip_channels = in_channels if (i == num_layers - 1) else out_channels785 resnet_in_channels = prev_output_channel if i == 0 else out_channels786 787 resnets.append(788 ResnetBlock3D(789 in_channels=resnet_in_channels + res_skip_channels,790 out_channels=out_channels,791 temb_channels=temb_channels,792 eps=resnet_eps,793 groups=resnet_groups,794 dropout=dropout,795 time_embedding_norm=resnet_time_scale_shift,796 non_linearity=resnet_act_fn,797 output_scale_factor=output_scale_factor,798 pre_norm=resnet_pre_norm,799 use_inflated_groupnorm=use_inflated_groupnorm,800 )801 )802 motion_modules.append(803 get_motion_module(804 in_channels=out_channels,805 motion_module_type=motion_module_type,806 motion_module_kwargs=motion_module_kwargs,807 )808 if use_motion_module809 else None810 )811 812 self.resnets = nn.ModuleList(resnets)813 self.motion_modules = nn.ModuleList(motion_modules)814 815 if add_upsample:816 self.upsamplers = nn.ModuleList(817 [Upsample3D(out_channels, use_conv=True, out_channels=out_channels)]818 )819 else:820 self.upsamplers = None821 822 self.gradient_checkpointing = False823 824 def forward(825 self,826 hidden_states,827 res_hidden_states_tuple,828 temb=None,829 upsample_size=None,830 encoder_hidden_states=None,831 ):832 for resnet, motion_module in zip(self.resnets, self.motion_modules):833 # pop res hidden states834 res_hidden_states = res_hidden_states_tuple[-1]835 res_hidden_states_tuple = res_hidden_states_tuple[:-1]836 hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)837 838 # print(f"UpBlock3D {self.gradient_checkpointing = }")839 if self.training and self.gradient_checkpointing:840 841 def create_custom_forward(module):842 def custom_forward(*inputs):843 return module(*inputs)844 845 return custom_forward846 847 hidden_states = torch.utils.checkpoint.checkpoint(848 create_custom_forward(resnet), hidden_states, temb849 )850 if motion_module is not None:851 hidden_states = torch.utils.checkpoint.checkpoint(852 create_custom_forward(motion_module),853 hidden_states.requires_grad_(),854 temb,855 encoder_hidden_states,856 )857 else:858 # print("hidden_states shape1:", hidden_states.shape)859 hidden_states = resnet(hidden_states, temb)860 # print("hidden_states shape2:", hidden_states.shape)861 hidden_states = (862 motion_module(863 hidden_states, temb, encoder_hidden_states=encoder_hidden_states864 )865 if motion_module is not None866 else hidden_states867 )868 869 if self.upsamplers is not None:870 for upsampler in self.upsamplers:871 hidden_states = upsampler(hidden_states, upsample_size)872 873 return hidden_states874 