Team Ai
Apppublic

michaelcreatesstuff/llm-grounded-diffusion

sourceHugging Faceupdated 3y agoView on Hugging Face
2likes
unet_2d_blocks.py794 linesDownload Raw Back to models
1# Copyright 2023 The HuggingFace Team. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7#     http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14from typing import Any, Dict, Optional, Tuple15 16import numpy as np17import torch18import torch.nn.functional as F19from torch import nn20 21from diffusers.utils import is_torch_version22from diffusers.models.dual_transformer_2d import DualTransformer2DModel23from diffusers.models.resnet import Downsample2D, ResnetBlock2D, Upsample2D24from .transformer_2d import Transformer2DModel25 26 27def get_down_block(28    down_block_type,29    num_layers,30    in_channels,31    out_channels,32    temb_channels,33    add_downsample,34    resnet_eps,35    resnet_act_fn,36    attn_num_head_channels,37    resnet_groups=None,38    cross_attention_dim=None,39    downsample_padding=None,40    dual_cross_attention=False,41    use_linear_projection=False,42    only_cross_attention=False,43    upcast_attention=False,44    resnet_time_scale_shift="default",45    resnet_skip_time_act=False,46    resnet_out_scale_factor=1.0,47    cross_attention_norm=None,48    use_gated_attention=False,49):50    down_block_type = down_block_type[7:] if down_block_type.startswith(51        "UNetRes") else down_block_type52    if down_block_type == "DownBlock2D":53        return DownBlock2D(54            num_layers=num_layers,55            in_channels=in_channels,56            out_channels=out_channels,57            temb_channels=temb_channels,58            add_downsample=add_downsample,59            resnet_eps=resnet_eps,60            resnet_act_fn=resnet_act_fn,61            resnet_groups=resnet_groups,62            downsample_padding=downsample_padding,63            resnet_time_scale_shift=resnet_time_scale_shift,64        )65    elif down_block_type == "CrossAttnDownBlock2D":66        if cross_attention_dim is None:67            raise ValueError(68                "cross_attention_dim must be specified for CrossAttnDownBlock2D")69        return CrossAttnDownBlock2D(70            num_layers=num_layers,71            in_channels=in_channels,72            out_channels=out_channels,73            temb_channels=temb_channels,74            add_downsample=add_downsample,75            resnet_eps=resnet_eps,76            resnet_act_fn=resnet_act_fn,77            resnet_groups=resnet_groups,78            downsample_padding=downsample_padding,79            cross_attention_dim=cross_attention_dim,80            attn_num_head_channels=attn_num_head_channels,81            dual_cross_attention=dual_cross_attention,82            use_linear_projection=use_linear_projection,83            only_cross_attention=only_cross_attention,84            upcast_attention=upcast_attention,85            resnet_time_scale_shift=resnet_time_scale_shift,86            use_gated_attention=use_gated_attention,87        )88 89    raise ValueError(f"{down_block_type} does not exist.")90 91 92def get_up_block(93    up_block_type,94    num_layers,95    in_channels,96    out_channels,97    prev_output_channel,98    temb_channels,99    add_upsample,100    resnet_eps,101    resnet_act_fn,102    attn_num_head_channels,103    resnet_groups=None,104    cross_attention_dim=None,105    dual_cross_attention=False,106    use_linear_projection=False,107    only_cross_attention=False,108    upcast_attention=False,109    resnet_time_scale_shift="default",110    resnet_skip_time_act=False,111    resnet_out_scale_factor=1.0,112    cross_attention_norm=None,113    use_gated_attention=False,114):115    up_block_type = up_block_type[7:] if up_block_type.startswith(116        "UNetRes") else up_block_type117    if up_block_type == "UpBlock2D":118        return UpBlock2D(119            num_layers=num_layers,120            in_channels=in_channels,121            out_channels=out_channels,122            prev_output_channel=prev_output_channel,123            temb_channels=temb_channels,124            add_upsample=add_upsample,125            resnet_eps=resnet_eps,126            resnet_act_fn=resnet_act_fn,127            resnet_groups=resnet_groups,128            resnet_time_scale_shift=resnet_time_scale_shift,129        )130    elif up_block_type == "CrossAttnUpBlock2D":131        if cross_attention_dim is None:132            raise ValueError(133                "cross_attention_dim must be specified for CrossAttnUpBlock2D")134        return CrossAttnUpBlock2D(135            num_layers=num_layers,136            in_channels=in_channels,137            out_channels=out_channels,138            prev_output_channel=prev_output_channel,139            temb_channels=temb_channels,140            add_upsample=add_upsample,141            resnet_eps=resnet_eps,142            resnet_act_fn=resnet_act_fn,143            resnet_groups=resnet_groups,144            cross_attention_dim=cross_attention_dim,145            attn_num_head_channels=attn_num_head_channels,146            dual_cross_attention=dual_cross_attention,147            use_linear_projection=use_linear_projection,148            only_cross_attention=only_cross_attention,149            upcast_attention=upcast_attention,150            resnet_time_scale_shift=resnet_time_scale_shift,151            use_gated_attention=use_gated_attention,152        )153 154    raise ValueError(f"{up_block_type} does not exist.")155 156 157class UNetMidBlock2DCrossAttn(nn.Module):158    def __init__(159        self,160        in_channels: int,161        temb_channels: int,162        dropout: float = 0.0,163        num_layers: int = 1,164        resnet_eps: float = 1e-6,165        resnet_time_scale_shift: str = "default",166        resnet_act_fn: str = "swish",167        resnet_groups: int = 32,168        resnet_pre_norm: bool = True,169        attn_num_head_channels=1,170        output_scale_factor=1.0,171        cross_attention_dim=1280,172        dual_cross_attention=False,173        use_linear_projection=False,174        upcast_attention=False,175        use_gated_attention=False,176    ):177        super().__init__()178 179        self.has_cross_attention = True180        self.attn_num_head_channels = attn_num_head_channels181        resnet_groups = resnet_groups if resnet_groups is not None else min(182            in_channels // 4, 32)183 184        # there is always at least one resnet185        resnets = [186            ResnetBlock2D(187                in_channels=in_channels,188                out_channels=in_channels,189                temb_channels=temb_channels,190                eps=resnet_eps,191                groups=resnet_groups,192                dropout=dropout,193                time_embedding_norm=resnet_time_scale_shift,194                non_linearity=resnet_act_fn,195                output_scale_factor=output_scale_factor,196                pre_norm=resnet_pre_norm,197            )198        ]199        attentions = []200 201        for _ in range(num_layers):202            if not dual_cross_attention:203                attentions.append(204                    Transformer2DModel(205                        attn_num_head_channels,206                        in_channels // attn_num_head_channels,207                        in_channels=in_channels,208                        num_layers=1,209                        cross_attention_dim=cross_attention_dim,210                        norm_num_groups=resnet_groups,211                        use_linear_projection=use_linear_projection,212                        upcast_attention=upcast_attention,213                        use_gated_attention=use_gated_attention,214                    )215                )216            else:217                attentions.append(218                    DualTransformer2DModel(219                        attn_num_head_channels,220                        in_channels // attn_num_head_channels,221                        in_channels=in_channels,222                        num_layers=1,223                        cross_attention_dim=cross_attention_dim,224                        norm_num_groups=resnet_groups,225                    )226                )227            resnets.append(228                ResnetBlock2D(229                    in_channels=in_channels,230                    out_channels=in_channels,231                    temb_channels=temb_channels,232                    eps=resnet_eps,233                    groups=resnet_groups,234                    dropout=dropout,235                    time_embedding_norm=resnet_time_scale_shift,236                    non_linearity=resnet_act_fn,237                    output_scale_factor=output_scale_factor,238                    pre_norm=resnet_pre_norm,239                )240            )241 242        self.attentions = nn.ModuleList(attentions)243        self.resnets = nn.ModuleList(resnets)244 245    def forward(246        self,247        hidden_states: torch.FloatTensor,248        temb: Optional[torch.FloatTensor] = None,249        encoder_hidden_states: Optional[torch.FloatTensor] = None,250        attention_mask: Optional[torch.FloatTensor] = None,251        cross_attention_kwargs: Optional[Dict[str, Any]] = None,252        encoder_attention_mask: Optional[torch.FloatTensor] = None,253        return_cross_attention_probs: bool = False,254    ) -> torch.FloatTensor:255        hidden_states = self.resnets[0](hidden_states, temb)256        cross_attention_probs_all = []257        base_attn_key = cross_attention_kwargs["attn_key"]258        for attn_key, (attn, resnet) in enumerate(zip(self.attentions, self.resnets[1:])):259            cross_attention_kwargs["attn_key"] = base_attn_key + [attn_key]260            hidden_states = attn(261                hidden_states,262                encoder_hidden_states=encoder_hidden_states,263                cross_attention_kwargs=cross_attention_kwargs,264                attention_mask=attention_mask,265                encoder_attention_mask=encoder_attention_mask,266                return_dict=False,267                return_cross_attention_probs=return_cross_attention_probs,268            )269            if return_cross_attention_probs:270                hidden_states, cross_attention_probs = hidden_states271                cross_attention_probs_all.append(cross_attention_probs)272            else:273                hidden_states = hidden_states[0]274            hidden_states = resnet(hidden_states, temb)275 276        if return_cross_attention_probs:277            return hidden_states, cross_attention_probs_all278        return hidden_states279 280 281class CrossAttnDownBlock2D(nn.Module):282    def __init__(283        self,284        in_channels: int,285        out_channels: int,286        temb_channels: int,287        dropout: float = 0.0,288        num_layers: int = 1,289        resnet_eps: float = 1e-6,290        resnet_time_scale_shift: str = "default",291        resnet_act_fn: str = "swish",292        resnet_groups: int = 32,293        resnet_pre_norm: bool = True,294        attn_num_head_channels=1,295        cross_attention_dim=1280,296        output_scale_factor=1.0,297        downsample_padding=1,298        add_downsample=True,299        dual_cross_attention=False,300        use_linear_projection=False,301        only_cross_attention=False,302        upcast_attention=False,303        use_gated_attention=False,304    ):305        super().__init__()306        resnets = []307        attentions = []308 309        self.has_cross_attention = True310        self.attn_num_head_channels = attn_num_head_channels311 312        for i in range(num_layers):313            in_channels = in_channels if i == 0 else out_channels314            resnets.append(315                ResnetBlock2D(316                    in_channels=in_channels,317                    out_channels=out_channels,318                    temb_channels=temb_channels,319                    eps=resnet_eps,320                    groups=resnet_groups,321                    dropout=dropout,322                    time_embedding_norm=resnet_time_scale_shift,323                    non_linearity=resnet_act_fn,324                    output_scale_factor=output_scale_factor,325                    pre_norm=resnet_pre_norm,326                )327            )328            if not dual_cross_attention:329                attentions.append(330                    Transformer2DModel(331                        attn_num_head_channels,332                        out_channels // attn_num_head_channels,333                        in_channels=out_channels,334                        num_layers=1,335                        cross_attention_dim=cross_attention_dim,336                        norm_num_groups=resnet_groups,337                        use_linear_projection=use_linear_projection,338                        only_cross_attention=only_cross_attention,339                        upcast_attention=upcast_attention,340                        use_gated_attention=use_gated_attention341                    )342                )343            else:344                attentions.append(345                    DualTransformer2DModel(346                        attn_num_head_channels,347                        out_channels // attn_num_head_channels,348                        in_channels=out_channels,349                        num_layers=1,350                        cross_attention_dim=cross_attention_dim,351                        norm_num_groups=resnet_groups,352                    )353                )354        self.attentions = nn.ModuleList(attentions)355        self.resnets = nn.ModuleList(resnets)356 357        if add_downsample:358            self.downsamplers = nn.ModuleList(359                [360                    Downsample2D(361                        out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op"362                    )363                ]364            )365        else:366            self.downsamplers = None367 368        self.gradient_checkpointing = False369 370    def forward(371        self,372        hidden_states: torch.FloatTensor,373        temb: Optional[torch.FloatTensor] = None,374        encoder_hidden_states: Optional[torch.FloatTensor] = None,375        attention_mask: Optional[torch.FloatTensor] = None,376        cross_attention_kwargs: Optional[Dict[str, Any]] = None,377        encoder_attention_mask: Optional[torch.FloatTensor] = None,378        return_cross_attention_probs: bool = False,379    ):380        output_states = ()381        cross_attention_probs_all = []382        base_attn_key = cross_attention_kwargs["attn_key"]383 384        for attn_key, (resnet, attn) in enumerate(zip(self.resnets, self.attentions)):385            386            cross_attention_kwargs["attn_key"] = base_attn_key + [attn_key]387            388            if self.training and self.gradient_checkpointing:389 390                def create_custom_forward(module, return_dict=None):391                    def custom_forward(*inputs):392                        if return_dict is not None:393                            return module(*inputs, return_dict=return_dict)394                        else:395                            return module(*inputs)396 397                    return custom_forward398 399                ckpt_kwargs: Dict[str, Any] = {400                    "use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}401                hidden_states = torch.utils.checkpoint.checkpoint(402                    create_custom_forward(resnet),403                    hidden_states,404                    temb,405                    **ckpt_kwargs,406                )407                hidden_states = torch.utils.checkpoint.checkpoint(408                    create_custom_forward(attn, return_dict=False),409                    hidden_states,410                    encoder_hidden_states,411                    None,  # timestep412                    None,  # class_labels413                    cross_attention_kwargs,414                    attention_mask,415                    encoder_attention_mask,416                    return_cross_attention_probs=return_cross_attention_probs,417                    **ckpt_kwargs,418                )419                if return_cross_attention_probs:420                    hidden_states, cross_attention_probs = hidden_states421                    cross_attention_probs_all.append(cross_attention_probs)422                else:423                    hidden_states = hidden_states[0]424            else:425                hidden_states = resnet(hidden_states, temb)426                hidden_states = attn(427                    hidden_states,428                    encoder_hidden_states=encoder_hidden_states,429                    cross_attention_kwargs=cross_attention_kwargs,430                    attention_mask=attention_mask,431                    encoder_attention_mask=encoder_attention_mask,432                    return_dict=False,433                    return_cross_attention_probs=return_cross_attention_probs,434                )435                if return_cross_attention_probs:436                    hidden_states, cross_attention_probs = hidden_states437                    cross_attention_probs_all.append(cross_attention_probs)438                else:439                    hidden_states = hidden_states[0]440 441            output_states = output_states + (hidden_states,)442 443        if self.downsamplers is not None:444            for downsampler in self.downsamplers:445                hidden_states = downsampler(hidden_states)446 447            output_states = output_states + (hidden_states,)448 449        if return_cross_attention_probs:450            return hidden_states, output_states, cross_attention_probs_all451        return hidden_states, output_states452 453 454class DownBlock2D(nn.Module):455    def __init__(456        self,457        in_channels: int,458        out_channels: int,459        temb_channels: int,460        dropout: float = 0.0,461        num_layers: int = 1,462        resnet_eps: float = 1e-6,463        resnet_time_scale_shift: str = "default",464        resnet_act_fn: str = "swish",465        resnet_groups: int = 32,466        resnet_pre_norm: bool = True,467        output_scale_factor=1.0,468        add_downsample=True,469        downsample_padding=1,470    ):471        super().__init__()472        resnets = []473 474        for i in range(num_layers):475            in_channels = in_channels if i == 0 else out_channels476            resnets.append(477                ResnetBlock2D(478                    in_channels=in_channels,479                    out_channels=out_channels,480                    temb_channels=temb_channels,481                    eps=resnet_eps,482                    groups=resnet_groups,483                    dropout=dropout,484                    time_embedding_norm=resnet_time_scale_shift,485                    non_linearity=resnet_act_fn,486                    output_scale_factor=output_scale_factor,487                    pre_norm=resnet_pre_norm,488                )489            )490 491        self.resnets = nn.ModuleList(resnets)492 493        if add_downsample:494            self.downsamplers = nn.ModuleList(495                [496                    Downsample2D(497                        out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op"498                    )499                ]500            )501        else:502            self.downsamplers = None503 504        self.gradient_checkpointing = False505 506    def forward(self, hidden_states, temb=None):507        output_states = ()508 509        for resnet in self.resnets:510            if self.training and self.gradient_checkpointing:511 512                def create_custom_forward(module):513                    def custom_forward(*inputs):514                        return module(*inputs)515 516                    return custom_forward517 518                if is_torch_version(">=", "1.11.0"):519                    hidden_states = torch.utils.checkpoint.checkpoint(520                        create_custom_forward(resnet), hidden_states, temb, use_reentrant=False521                    )522                else:523                    hidden_states = torch.utils.checkpoint.checkpoint(524                        create_custom_forward(resnet), hidden_states, temb525                    )526            else:527                hidden_states = resnet(hidden_states, temb)528 529            output_states = output_states + (hidden_states,)530 531        if self.downsamplers is not None:532            for downsampler in self.downsamplers:533                hidden_states = downsampler(hidden_states)534 535            output_states = output_states + (hidden_states,)536 537        return hidden_states, output_states538 539 540class CrossAttnUpBlock2D(nn.Module):541    def __init__(542        self,543        in_channels: int,544        out_channels: int,545        prev_output_channel: int,546        temb_channels: int,547        dropout: float = 0.0,548        num_layers: int = 1,549        resnet_eps: float = 1e-6,550        resnet_time_scale_shift: str = "default",551        resnet_act_fn: str = "swish",552        resnet_groups: int = 32,553        resnet_pre_norm: bool = True,554        attn_num_head_channels=1,555        cross_attention_dim=1280,556        output_scale_factor=1.0,557        add_upsample=True,558        dual_cross_attention=False,559        use_linear_projection=False,560        only_cross_attention=False,561        upcast_attention=False,562        use_gated_attention=False,563    ):564        super().__init__()565        resnets = []566        attentions = []567 568        self.has_cross_attention = True569        self.attn_num_head_channels = attn_num_head_channels570 571        for i in range(num_layers):572            res_skip_channels = in_channels if (573                i == num_layers - 1) else out_channels574            resnet_in_channels = prev_output_channel if i == 0 else out_channels575 576            resnets.append(577                ResnetBlock2D(578                    in_channels=resnet_in_channels + res_skip_channels,579                    out_channels=out_channels,580                    temb_channels=temb_channels,581                    eps=resnet_eps,582                    groups=resnet_groups,583                    dropout=dropout,584                    time_embedding_norm=resnet_time_scale_shift,585                    non_linearity=resnet_act_fn,586                    output_scale_factor=output_scale_factor,587                    pre_norm=resnet_pre_norm,588                )589            )590            if not dual_cross_attention:591                attentions.append(592                    Transformer2DModel(593                        attn_num_head_channels,594                        out_channels // attn_num_head_channels,595                        in_channels=out_channels,596                        num_layers=1,597                        cross_attention_dim=cross_attention_dim,598                        norm_num_groups=resnet_groups,599                        use_linear_projection=use_linear_projection,600                        only_cross_attention=only_cross_attention,601                        upcast_attention=upcast_attention,602                        use_gated_attention=use_gated_attention,603                    )604                )605            else:606                attentions.append(607                    DualTransformer2DModel(608                        attn_num_head_channels,609                        out_channels // attn_num_head_channels,610                        in_channels=out_channels,611                        num_layers=1,612                        cross_attention_dim=cross_attention_dim,613                        norm_num_groups=resnet_groups,614                    )615                )616        self.attentions = nn.ModuleList(attentions)617        self.resnets = nn.ModuleList(resnets)618 619        if add_upsample:620            self.upsamplers = nn.ModuleList(621                [Upsample2D(out_channels, use_conv=True, out_channels=out_channels)])622        else:623            self.upsamplers = None624 625        self.gradient_checkpointing = False626 627    def forward(628        self,629        hidden_states: torch.FloatTensor,630        res_hidden_states_tuple: Tuple[torch.FloatTensor, ...],631        temb: Optional[torch.FloatTensor] = None,632        encoder_hidden_states: Optional[torch.FloatTensor] = None,633        cross_attention_kwargs: Optional[Dict[str, Any]] = None,634        upsample_size: Optional[int] = None,635        attention_mask: Optional[torch.FloatTensor] = None,636        encoder_attention_mask: Optional[torch.FloatTensor] = None,637        return_cross_attention_probs: bool = False,638    ):639        cross_attention_probs_all = []640        base_attn_key = cross_attention_kwargs["attn_key"]641        642        for attn_key, (resnet, attn) in enumerate(zip(self.resnets, self.attentions)):643            cross_attention_kwargs["attn_key"] = base_attn_key + [attn_key]644            645            # pop res hidden states646            res_hidden_states = res_hidden_states_tuple[-1]647            res_hidden_states_tuple = res_hidden_states_tuple[:-1]648            hidden_states = torch.cat(649                [hidden_states, res_hidden_states], dim=1)650 651            if self.training and self.gradient_checkpointing:652 653                def create_custom_forward(module, return_dict=None):654                    def custom_forward(*inputs):655                        if return_dict is not None:656                            return module(*inputs, return_dict=return_dict)657                        else:658                            return module(*inputs)659 660                    return custom_forward661 662                ckpt_kwargs: Dict[str, Any] = {663                    "use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}664                hidden_states = torch.utils.checkpoint.checkpoint(665                    create_custom_forward(resnet),666                    hidden_states,667                    temb,668                    **ckpt_kwargs,669                )670                hidden_states = torch.utils.checkpoint.checkpoint(671                    create_custom_forward(attn, return_dict=False),672                    hidden_states,673                    encoder_hidden_states,674                    None,  # timestep675                    None,  # class_labels676                    cross_attention_kwargs,677                    attention_mask,678                    encoder_attention_mask,679                    **ckpt_kwargs,680                )681                if return_cross_attention_probs:682                    hidden_states, cross_attention_probs = hidden_states683                    cross_attention_probs_all.append(cross_attention_probs)684                else:685                    hidden_states = hidden_states[0]686            else:687                hidden_states = resnet(hidden_states, temb)688                hidden_states = attn(689                    hidden_states,690                    encoder_hidden_states=encoder_hidden_states,691                    cross_attention_kwargs=cross_attention_kwargs,692                    attention_mask=attention_mask,693                    encoder_attention_mask=encoder_attention_mask,694                    return_dict=False,695                    return_cross_attention_probs=return_cross_attention_probs,696                )697                if return_cross_attention_probs:698                    hidden_states, cross_attention_probs = hidden_states699                    cross_attention_probs_all.append(cross_attention_probs)700                else:701                    hidden_states = hidden_states[0]702 703        if self.upsamplers is not None:704            for upsampler in self.upsamplers:705                hidden_states = upsampler(hidden_states, upsample_size)706 707        if return_cross_attention_probs:708            return hidden_states, cross_attention_probs_all709        return hidden_states710 711 712class UpBlock2D(nn.Module):713    def __init__(714        self,715        in_channels: int,716        prev_output_channel: int,717        out_channels: int,718        temb_channels: int,719        dropout: float = 0.0,720        num_layers: int = 1,721        resnet_eps: float = 1e-6,722        resnet_time_scale_shift: str = "default",723        resnet_act_fn: str = "swish",724        resnet_groups: int = 32,725        resnet_pre_norm: bool = True,726        output_scale_factor=1.0,727        add_upsample=True,728    ):729        super().__init__()730        resnets = []731 732        for i in range(num_layers):733            res_skip_channels = in_channels if (734                i == num_layers - 1) else out_channels735            resnet_in_channels = prev_output_channel if i == 0 else out_channels736 737            resnets.append(738                ResnetBlock2D(739                    in_channels=resnet_in_channels + res_skip_channels,740                    out_channels=out_channels,741                    temb_channels=temb_channels,742                    eps=resnet_eps,743                    groups=resnet_groups,744                    dropout=dropout,745                    time_embedding_norm=resnet_time_scale_shift,746                    non_linearity=resnet_act_fn,747                    output_scale_factor=output_scale_factor,748                    pre_norm=resnet_pre_norm,749                )750            )751 752        self.resnets = nn.ModuleList(resnets)753 754        if add_upsample:755            self.upsamplers = nn.ModuleList(756                [Upsample2D(out_channels, use_conv=True, out_channels=out_channels)])757        else:758            self.upsamplers = None759 760        self.gradient_checkpointing = False761 762    def forward(self, hidden_states, res_hidden_states_tuple, temb=None, upsample_size=None):763        for resnet in self.resnets:764            # pop res hidden states765            res_hidden_states = res_hidden_states_tuple[-1]766            res_hidden_states_tuple = res_hidden_states_tuple[:-1]767            hidden_states = torch.cat(768                [hidden_states, res_hidden_states], dim=1)769 770            if self.training and self.gradient_checkpointing:771 772                def create_custom_forward(module):773                    def custom_forward(*inputs):774                        return module(*inputs)775 776                    return custom_forward777 778                if is_torch_version(">=", "1.11.0"):779                    hidden_states = torch.utils.checkpoint.checkpoint(780                        create_custom_forward(resnet), hidden_states, temb, use_reentrant=False781                    )782                else:783                    hidden_states = torch.utils.checkpoint.checkpoint(784                        create_custom_forward(resnet), hidden_states, temb785                    )786            else:787                hidden_states = resnet(hidden_states, temb)788 789        if self.upsamplers is not None:790            for upsampler in self.upsamplers:791                hidden_states = upsampler(hidden_states, upsample_size)792 793        return hidden_states794