michaelcreatesstuff/llm-grounded-diffusion
2
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 