CoolFace
Apppublic

multimodalart/EchoMimic-zero

sourceHugging Faceupdated 2y agoView on Hugging Face
8likes
unet_3d_blocks.py874 linesDownload Raw Back to models
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