CoolFace
Apppublic

svjack/ControlNet-Pose-Chinese

sourceHugging Faceupdated 3y agoView on Hugging Face
5likes
models.py836 linesDownload Raw Back to root
1import torch2import torch.nn as nn3import torch.nn.functional as F4 5from typing import List, Tuple, Union6from dataclasses import dataclass7from diffusers.utils.outputs import BaseOutput8from diffusers.configuration_utils import ConfigMixin, register_to_config9from diffusers.models.modeling_utils import ModelMixin10from diffusers.models.unet_2d_blocks import get_down_block as get_down_block_default11from diffusers.models.resnet import Mish, Upsample2D, Downsample2D, upsample_2d, downsample_2d, partial12from diffusers.models.cross_attention import CrossAttention, LoRALinearLayer # , LoRACrossAttnProcessor13 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    resnet_kernel_size=3,34):35    down_block_type = down_block_type[7:] if down_block_type.startswith("UNetRes") else down_block_type36    if down_block_type == "SimpleDownEncoderBlock2D":37        return SimpleDownEncoderBlock2D(38            num_layers=num_layers,39            in_channels=in_channels,40            out_channels=out_channels,41            add_downsample=add_downsample,42            convnet_eps=resnet_eps,43            convnet_act_fn=resnet_act_fn,44            convnet_groups=resnet_groups,45            downsample_padding=downsample_padding,46            convnet_time_scale_shift=resnet_time_scale_shift,47            convnet_kernel_size=resnet_kernel_size48        )49    else:50        return get_down_block_default(51            down_block_type,52            num_layers,53            in_channels,54            out_channels,55            temb_channels,56            add_downsample,57            resnet_eps,58            resnet_act_fn,59            attn_num_head_channels,60            resnet_groups=resnet_groups,61            cross_attention_dim=cross_attention_dim,62            downsample_padding=downsample_padding,63            dual_cross_attention=dual_cross_attention,64            use_linear_projection=use_linear_projection,65            only_cross_attention=only_cross_attention,66            upcast_attention=upcast_attention,67            resnet_time_scale_shift=resnet_time_scale_shift,68            # resnet_kernel_size=resnet_kernel_size69        )70 71 72class LoRACrossAttnProcessor(nn.Module):73    def __init__(74            self, 75            hidden_size, 76            cross_attention_dim=None, 77            rank=4, 78            post_add=False,79            key_states_skipped=False,80            value_states_skipped=False,81            output_states_skipped=False):82        super().__init__()83 84        self.hidden_size = hidden_size85        self.cross_attention_dim = cross_attention_dim86        self.rank = rank87        self.post_add = post_add88 89        self.to_q_lora = LoRALinearLayer(hidden_size, hidden_size, rank)90        if not key_states_skipped:91            self.to_k_lora = LoRALinearLayer(92                hidden_size if post_add else (cross_attention_dim or hidden_size), hidden_size, rank)93        if not value_states_skipped:94            self.to_v_lora = LoRALinearLayer(95                hidden_size if post_add else (cross_attention_dim or hidden_size), hidden_size, rank)96        if not output_states_skipped:97            self.to_out_lora = LoRALinearLayer(hidden_size, hidden_size, rank)98 99        self.key_states_skipped: bool = key_states_skipped100        self.value_states_skipped: bool = value_states_skipped101        self.output_states_skipped: bool = output_states_skipped102 103    def skip_key_states(self, is_skipped: bool = True):104        if is_skipped == False:105            assert hasattr(self, 'to_k_lora')106        self.key_states_skipped = is_skipped107 108    def skip_value_states(self, is_skipped: bool = True):109        if is_skipped == False:110            assert hasattr(self, 'to_q_lora')111        self.value_states_skipped = is_skipped112 113    def skip_output_states(self, is_skipped: bool = True):114        if is_skipped == False:115            assert hasattr(self, 'to_out_lora')116        self.output_states_skipped = is_skipped117 118    def __call__(119        self, attn: CrossAttention, hidden_states, encoder_hidden_states=None, attention_mask=None, scale=1.0120    ):121        batch_size, sequence_length, _ = hidden_states.shape122        attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length, batch_size)123 124        query = attn.to_q(hidden_states) 125        query = query + scale * self.to_q_lora(query if self.post_add else hidden_states)126        query = attn.head_to_batch_dim(query)127 128        encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states129 130        key = attn.to_k(encoder_hidden_states) 131        if not self.key_states_skipped:132            key = key + scale * self.to_k_lora(key if self.post_add else encoder_hidden_states)133        value = attn.to_v(encoder_hidden_states)134        if not self.value_states_skipped:135            value = value + scale * self.to_v_lora(value if self.post_add else encoder_hidden_states)136 137        key = attn.head_to_batch_dim(key)138        value = attn.head_to_batch_dim(value)139 140        attention_probs = attn.get_attention_scores(query, key, attention_mask)141        hidden_states = torch.bmm(attention_probs, value)142        hidden_states = attn.batch_to_head_dim(hidden_states)143 144        # linear proj145        out = attn.to_out[0](hidden_states)146        if not self.output_states_skipped:147            out = out + scale * self.to_out_lora(out if self.post_add else hidden_states)148        hidden_states = out149        # dropout150        hidden_states = attn.to_out[1](hidden_states)151 152        return hidden_states153 154 155class ControlLoRACrossAttnProcessor(LoRACrossAttnProcessor):156    def __init__(157            self, 158            hidden_size, 159            cross_attention_dim=None, 160            rank=4, 161            control_rank=None, 162            post_add=False, 163            concat_hidden=False,164            control_channels=None,165            control_self_add=True,166            key_states_skipped=False,167            value_states_skipped=False,168            output_states_skipped=False,169            **kwargs):170        super().__init__(171            hidden_size, 172            cross_attention_dim, 173            rank, 174            post_add=post_add,175            key_states_skipped=key_states_skipped,176            value_states_skipped=value_states_skipped,177            output_states_skipped=output_states_skipped)178 179        control_rank = rank if control_rank is None else control_rank180        control_channels = hidden_size if control_channels is None else control_channels181        self.concat_hidden = concat_hidden182        self.control_self_add = control_self_add if control_channels is None else False183        self.control_states: torch.Tensor = None184 185        self.to_control = LoRALinearLayer(186            control_channels + (hidden_size if concat_hidden else 0), 187            hidden_size, 188            control_rank)189        self.pre_loras: List[LoRACrossAttnProcessor] = []190        self.post_loras: List[LoRACrossAttnProcessor] = []191 192    def inject_pre_lora(self, lora_layer):193        self.pre_loras.append(lora_layer)194    195    def inject_post_lora(self, lora_layer):196        self.post_loras.append(lora_layer)197 198    def inject_control_states(self, control_states):199        self.control_states = control_states200 201    def process_control_states(self, hidden_states, scale=1.0):202        control_states = self.control_states.to(hidden_states.dtype)203        if hidden_states.ndim == 3 and control_states.ndim == 4:204            batch, _, height, width = control_states.shape205            control_states = control_states.permute(0, 2, 3, 1).reshape(batch, height * width, -1)206            self.control_states = control_states207        _control_states = control_states208        if self.concat_hidden:209            b1, b2 = control_states.shape[0], hidden_states.shape[0]210            if b1 != b2:211                control_states = control_states[:,None].repeat(1, b2//b1, *([1]*(len(control_states.shape)-1)))212                control_states = control_states.view(-1, *control_states.shape[2:])213            _control_states = torch.cat([hidden_states, control_states], -1)214        _control_states = scale * self.to_control(_control_states)215        if self.control_self_add:216            control_states = control_states + _control_states217        else:218            control_states = _control_states219 220        return control_states221 222    def __call__(223        self, attn: CrossAttention, hidden_states, encoder_hidden_states=None, attention_mask=None, scale=1.0224    ):225        pre_lora: LoRACrossAttnProcessor226        post_lora: LoRACrossAttnProcessor227        assert self.control_states is not None228 229        batch_size, sequence_length, _ = hidden_states.shape230        attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length)231        query = attn.to_q(hidden_states)232        for pre_lora in self.pre_loras:233            lora_in = query if pre_lora.post_add else hidden_states234            if isinstance(pre_lora, ControlLoRACrossAttnProcessor):235                lora_in = lora_in + pre_lora.process_control_states(hidden_states, scale)236            query = query + scale * pre_lora.to_q_lora(lora_in)237        query = query + scale * self.to_q_lora((238            query if self.post_add else hidden_states) + self.process_control_states(hidden_states, scale))239        for post_lora in self.post_loras:240            lora_in = query if post_lora.post_add else hidden_states241            if isinstance(post_lora, ControlLoRACrossAttnProcessor):242                lora_in = lora_in + post_lora.process_control_states(hidden_states, scale)243            query = query + scale * post_lora.to_q_lora(lora_in)244        query = attn.head_to_batch_dim(query)245 246        encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states247 248        key = attn.to_k(encoder_hidden_states)249        for pre_lora in self.pre_loras:250            if not pre_lora.key_states_skipped:251                key = key + scale * pre_lora.to_k_lora(key if pre_lora.post_add else encoder_hidden_states)252        if not self.key_states_skipped:253            key = key + scale * self.to_k_lora(key if self.post_add else encoder_hidden_states)254        for post_lora in self.post_loras:255            if not post_lora.key_states_skipped:256                key = key + scale * post_lora.to_k_lora(key if post_lora.post_add else encoder_hidden_states)257        value = attn.to_v(encoder_hidden_states)258        for pre_lora in self.pre_loras:259            if not pre_lora.value_states_skipped:260                value = value + pre_lora.to_v_lora(value if pre_lora.post_add else encoder_hidden_states)261        if not self.value_states_skipped:262            value = value + scale * self.to_v_lora(value if self.post_add else encoder_hidden_states)263        for post_lora in self.post_loras:264            if not post_lora.value_states_skipped:265                value = value + post_lora.to_v_lora(value if post_lora.post_add else encoder_hidden_states)266 267        key = attn.head_to_batch_dim(key)268        value = attn.head_to_batch_dim(value)269 270        attention_probs = attn.get_attention_scores(query, key, attention_mask)271        hidden_states = torch.bmm(attention_probs, value)272        hidden_states = attn.batch_to_head_dim(hidden_states)273 274        # linear proj275        out = attn.to_out[0](hidden_states)276        for pre_lora in self.pre_loras:277            if not pre_lora.output_states_skipped:278                out = out + scale * pre_lora.to_out_lora(out if pre_lora.post_add else hidden_states)279        out = out + scale * self.to_out_lora(out if self.post_add else hidden_states)280        for post_lora in self.post_loras:281            if not post_lora.output_states_skipped:282                out = out + scale * post_lora.to_out_lora(out if post_lora.post_add else hidden_states)283        hidden_states = out284        # dropout285        hidden_states = attn.to_out[1](hidden_states)286 287        return hidden_states288 289 290 291 292class ControlLoRACrossAttnProcessorV2(LoRACrossAttnProcessor):293    def __init__(294            self, 295            hidden_size, 296            cross_attention_dim=None, 297            rank=4, 298            control_rank=None, 299            control_channels=None,300            **kwargs):301        super().__init__(302            hidden_size, 303            cross_attention_dim, 304            rank, 305            post_add=False,306            key_states_skipped=True,307            value_states_skipped=True,308            output_states_skipped=False)309 310        control_rank = rank if control_rank is None else control_rank311        control_channels = hidden_size if control_channels is None else control_channels312        self.concat_hidden = True313        self.control_self_add = False314        self.control_states: torch.Tensor = None315 316        self.to_control = LoRALinearLayer(317            hidden_size + control_channels, 318            hidden_size, 319            control_rank)320        self.to_control_out = LoRALinearLayer(321            hidden_size + control_channels, 322            hidden_size, 323            control_rank)324        self.pre_loras: List[LoRACrossAttnProcessor] = []325        self.post_loras: List[LoRACrossAttnProcessor] = []326 327    def inject_pre_lora(self, lora_layer):328        self.pre_loras.append(lora_layer)329    330    def inject_post_lora(self, lora_layer):331        self.post_loras.append(lora_layer)332 333    def inject_control_states(self, control_states):334        self.control_states = control_states335 336    def process_control_states(self, hidden_states, scale=1.0, is_out=False):337        control_states = self.control_states.to(hidden_states.dtype)338        if hidden_states.ndim == 3 and control_states.ndim == 4:339            batch, _, height, width = control_states.shape340            control_states = control_states.permute(0, 2, 3, 1).reshape(batch, height * width, -1)341            self.control_states = control_states342        _control_states = control_states343        if self.concat_hidden:344            b1, b2 = control_states.shape[0], hidden_states.shape[0]345            if b1 != b2:346                control_states = control_states[:,None].repeat(1, b2//b1, *([1]*(len(control_states.shape)-1)))347                control_states = control_states.view(-1, *control_states.shape[2:])348            _control_states = torch.cat([hidden_states, control_states], -1)349        _control_states = scale * (self.to_control_out if is_out else self.to_control)(_control_states)350        if self.control_self_add:351            control_states = control_states + _control_states352        else:353            control_states = _control_states354 355        return control_states356 357    def __call__(358        self, attn: CrossAttention, hidden_states, encoder_hidden_states=None, attention_mask=None, scale=1.0359    ):360        pre_lora: LoRACrossAttnProcessor361        post_lora: LoRACrossAttnProcessor362        assert self.control_states is not None363 364        batch_size, sequence_length, _ = hidden_states.shape365        attention_mask = attn.prepare_attention_mask(attention_mask, sequence_length)366        for pre_lora in self.pre_loras:367            if isinstance(pre_lora, ControlLoRACrossAttnProcessorV2):368                hidden_states = hidden_states + pre_lora.process_control_states(hidden_states, scale)369        hidden_states = hidden_states + self.process_control_states(hidden_states, scale)370        for post_lora in self.post_loras:371            if isinstance(post_lora, ControlLoRACrossAttnProcessorV2):372                hidden_states = hidden_states + post_lora.process_control_states(hidden_states, scale)373        query = attn.to_q(hidden_states)374        for pre_lora in self.pre_loras:375            lora_in = query if pre_lora.post_add else hidden_states376            query = query + scale * pre_lora.to_q_lora(lora_in)377        query = query + scale * self.to_q_lora(query if self.post_add else hidden_states)378        for post_lora in self.post_loras:379            lora_in = query if post_lora.post_add else hidden_states380            query = query + scale * post_lora.to_q_lora(lora_in)381        query = attn.head_to_batch_dim(query)382 383        encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states384 385        key = attn.to_k(encoder_hidden_states)386        for pre_lora in self.pre_loras:387            if not pre_lora.key_states_skipped:388                key = key + scale * pre_lora.to_k_lora(key if pre_lora.post_add else encoder_hidden_states)389        if not self.key_states_skipped:390            key = key + scale * self.to_k_lora(key if self.post_add else encoder_hidden_states)391        for post_lora in self.post_loras:392            if not post_lora.key_states_skipped:393                key = key + scale * post_lora.to_k_lora(key if post_lora.post_add else encoder_hidden_states)394        value = attn.to_v(encoder_hidden_states)395        for pre_lora in self.pre_loras:396            if not pre_lora.value_states_skipped:397                value = value + pre_lora.to_v_lora(value if pre_lora.post_add else encoder_hidden_states)398        if not self.value_states_skipped:399            value = value + scale * self.to_v_lora(value if self.post_add else encoder_hidden_states)400        for post_lora in self.post_loras:401            if not post_lora.value_states_skipped:402                value = value + post_lora.to_v_lora(value if post_lora.post_add else encoder_hidden_states)403 404        key = attn.head_to_batch_dim(key)405        value = attn.head_to_batch_dim(value)406 407        attention_probs = attn.get_attention_scores(query, key, attention_mask)408        hidden_states = torch.bmm(attention_probs, value)409        hidden_states = attn.batch_to_head_dim(hidden_states)410 411        # linear proj412        for pre_lora in self.pre_loras:413            if isinstance(pre_lora, ControlLoRACrossAttnProcessorV2):414                hidden_states = hidden_states + pre_lora.process_control_states(hidden_states, scale, is_out=True)415        hidden_states = hidden_states + self.process_control_states(hidden_states, scale, is_out=True)416        for post_lora in self.post_loras:417            if isinstance(post_lora, ControlLoRACrossAttnProcessorV2):418                hidden_states = hidden_states + post_lora.process_control_states(hidden_states, scale, is_out=True)419        out = attn.to_out[0](hidden_states)420        for pre_lora in self.pre_loras:421            if not pre_lora.output_states_skipped:422                out = out + scale * pre_lora.to_out_lora(out if pre_lora.post_add else hidden_states)423        out = out + scale * self.to_out_lora(out if self.post_add else hidden_states)424        for post_lora in self.post_loras:425            if not post_lora.output_states_skipped:426                out = out + scale * post_lora.to_out_lora(out if post_lora.post_add else hidden_states)427        hidden_states = out428        # dropout429        hidden_states = attn.to_out[1](hidden_states)430 431        return hidden_states432 433 434class ConvBlock2D(nn.Module):435    def __init__(436        self,437        *,438        in_channels,439        out_channels=None,440        conv_kernel_size=3,441        dropout=0.0,442        temb_channels=512,443        groups=32,444        groups_out=None,445        pre_norm=True,446        eps=1e-6,447        non_linearity="swish",448        time_embedding_norm="default",449        kernel=None,450        output_scale_factor=1.0,451        up=False,452        down=False,453    ):454        super().__init__()455        self.pre_norm = pre_norm456        self.pre_norm = True457        self.in_channels = in_channels458        out_channels = in_channels if out_channels is None else out_channels459        self.out_channels = out_channels460        self.time_embedding_norm = time_embedding_norm461        self.up = up462        self.down = down463        self.output_scale_factor = output_scale_factor464 465        if groups_out is None:466            groups_out = groups467 468        self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)469 470        self.conv1 = torch.nn.Conv2d(in_channels, out_channels, kernel_size=conv_kernel_size, stride=1, padding=conv_kernel_size//2)471 472        if temb_channels is not None:473            if self.time_embedding_norm == "default":474                time_emb_proj_out_channels = out_channels475            elif self.time_embedding_norm == "scale_shift":476                time_emb_proj_out_channels = out_channels * 2477            else:478                raise ValueError(f"unknown time_embedding_norm : {self.time_embedding_norm} ")479 480            self.time_emb_proj = torch.nn.Linear(temb_channels, time_emb_proj_out_channels)481        else:482            self.time_emb_proj = None483 484        self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)485        self.dropout = torch.nn.Dropout(dropout)486 487        if non_linearity == "swish":488            self.nonlinearity = lambda x: F.silu(x)489        elif non_linearity == "mish":490            self.nonlinearity = Mish()491        elif non_linearity == "silu":492            self.nonlinearity = nn.SiLU()493 494        self.upsample = self.downsample = None495        if self.up:496            if kernel == "fir":497                fir_kernel = (1, 3, 3, 1)498                self.upsample = lambda x: upsample_2d(x, kernel=fir_kernel)499            elif kernel == "sde_vp":500                self.upsample = partial(F.interpolate, scale_factor=2.0, mode="nearest")501            else:502                self.upsample = Upsample2D(in_channels, use_conv=False)503        elif self.down:504            if kernel == "fir":505                fir_kernel = (1, 3, 3, 1)506                self.downsample = lambda x: downsample_2d(x, kernel=fir_kernel)507            elif kernel == "sde_vp":508                self.downsample = partial(F.avg_pool2d, kernel_size=2, stride=2)509            else:510                self.downsample = Downsample2D(in_channels, use_conv=False, padding=1, name="op")511 512    def forward(self, input_tensor, temb):513        hidden_states = input_tensor514 515        hidden_states = self.norm1(hidden_states)516        hidden_states = self.nonlinearity(hidden_states)517 518        if self.upsample is not None:519            # upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984520            if hidden_states.shape[0] >= 64:521                input_tensor = input_tensor.contiguous()522                hidden_states = hidden_states.contiguous()523            input_tensor = self.upsample(input_tensor)524            hidden_states = self.upsample(hidden_states)525        elif self.downsample is not None:526            input_tensor = self.downsample(input_tensor)527            hidden_states = self.downsample(hidden_states)528 529        hidden_states = self.conv1(hidden_states)530 531        if temb is not None:532            temb = self.time_emb_proj(self.nonlinearity(temb))[:, :, None, None]533 534        if temb is not None and self.time_embedding_norm == "default":535            hidden_states = hidden_states + temb536 537        hidden_states = self.norm2(hidden_states)538 539        if temb is not None and self.time_embedding_norm == "scale_shift":540            scale, shift = torch.chunk(temb, 2, dim=1)541            hidden_states = hidden_states * (1 + scale) + shift542 543        hidden_states = self.nonlinearity(hidden_states)544 545        output_tensor = self.dropout(hidden_states)546 547        return output_tensor548 549 550class SimpleDownEncoderBlock2D(nn.Module):551    def __init__(552        self,553        in_channels: int,554        out_channels: int,555        dropout: float = 0.0,556        num_layers: int = 1,557        convnet_eps: float = 1e-6,558        convnet_time_scale_shift: str = "default",559        convnet_act_fn: str = "swish",560        convnet_groups: int = 32,561        convnet_pre_norm: bool = True,562        convnet_kernel_size: int = 3,563        output_scale_factor=1.0,564        add_downsample=True,565        downsample_padding=1,566    ):567        super().__init__()568        convnets = []569 570        for i in range(num_layers):571            in_channels = in_channels if i == 0 else out_channels572            convnets.append(573                ConvBlock2D(574                    in_channels=in_channels,575                    out_channels=out_channels,576                    temb_channels=None,577                    eps=convnet_eps,578                    groups=convnet_groups,579                    dropout=dropout,580                    time_embedding_norm=convnet_time_scale_shift,581                    non_linearity=convnet_act_fn,582                    output_scale_factor=output_scale_factor,583                    pre_norm=convnet_pre_norm,584                    conv_kernel_size=convnet_kernel_size,585                )586            )587        in_channels = in_channels if num_layers == 0 else out_channels588 589        self.convnets = nn.ModuleList(convnets)590 591        if add_downsample:592            self.downsamplers = nn.ModuleList(593                [594                    Downsample2D(595                        in_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op"596                    )597                ]598            )599        else:600            self.downsamplers = None601 602    def forward(self, hidden_states):603        for convnet in self.convnets:604            hidden_states = convnet(hidden_states, temb=None)605 606        if self.downsamplers is not None:607            for downsampler in self.downsamplers:608                hidden_states = downsampler(hidden_states)609 610        return hidden_states611 612 613@dataclass614class ControlLoRAOutput(BaseOutput):615    control_states: Tuple[torch.FloatTensor]616 617 618class ControlLoRA(ModelMixin, ConfigMixin):619    @register_to_config620    def __init__(621        self,622        in_channels: int = 3,623        down_block_types: Tuple[str] = (624            "SimpleDownEncoderBlock2D",625            "SimpleDownEncoderBlock2D",626            "SimpleDownEncoderBlock2D",627            "SimpleDownEncoderBlock2D",628        ),629        block_out_channels: Tuple[int] = (32, 64, 128, 256),630        layers_per_block: int = 1,631        act_fn: str = "silu",632        norm_num_groups: int = 32,633        lora_pre_down_block_types: Tuple[str] = (634            None,635            "SimpleDownEncoderBlock2D",636            "SimpleDownEncoderBlock2D",637            "SimpleDownEncoderBlock2D",638        ),639        lora_pre_down_layers_per_block: int = 1,640        lora_pre_conv_skipped: bool = False,641        lora_pre_conv_types: Tuple[str] = (642            "SimpleDownEncoderBlock2D",643            "SimpleDownEncoderBlock2D",644            "SimpleDownEncoderBlock2D",645            "SimpleDownEncoderBlock2D",646        ),647        lora_pre_conv_layers_per_block: int = 1,648        lora_pre_conv_layers_kernel_size: int = 1,649        lora_block_in_channels: Tuple[int] = (256, 256, 256, 256),650        lora_block_out_channels: Tuple[int] = (320, 640, 1280, 1280),651        lora_cross_attention_dims: Tuple[List[int]] = (652            [None, 768, None, 768, None, 768, None, 768, None, 768], 653            [None, 768, None, 768, None, 768, None, 768, None, 768], 654            [None, 768, None, 768, None, 768, None, 768, None, 768], 655            [None, 768]656        ),657        lora_rank: int = 4,658        lora_control_rank: int = None,659        lora_post_add: bool = False,660        lora_concat_hidden: bool = False,661        lora_control_channels: Tuple[int] = (None, None, None, None),662        lora_control_self_add: bool = True,663        lora_key_states_skipped: bool = False,664        lora_value_states_skipped: bool = False,665        lora_output_states_skipped: bool = False,666        lora_control_version: int = 1667    ):668        super().__init__()669 670        lora_control_cls = ControlLoRACrossAttnProcessor671        if lora_control_version == 2:672            lora_control_cls = ControlLoRACrossAttnProcessorV2673 674        assert lora_block_in_channels[0] == block_out_channels[-1]675        676        if lora_pre_conv_skipped:677            lora_control_channels = lora_block_in_channels678            lora_control_self_add = False679 680        self.layers_per_block = layers_per_block681        self.lora_pre_down_layers_per_block = lora_pre_down_layers_per_block682        self.lora_pre_conv_layers_per_block = lora_pre_conv_layers_per_block683 684        self.conv_in = torch.nn.Conv2d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)685 686        self.down_blocks = nn.ModuleList([])687        self.pre_lora_layers = nn.ModuleList([])688        self.lora_layers = nn.ModuleList([])689 690        # pre_down691        pre_down_blocks = []692        output_channel = block_out_channels[0]693        for i, down_block_type in enumerate(down_block_types):694            input_channel = output_channel695            output_channel = block_out_channels[i]696            is_final_block = i == len(block_out_channels) - 1697 698            pre_down_block = get_down_block(699                down_block_type,700                num_layers=self.layers_per_block,701                in_channels=input_channel,702                out_channels=output_channel,703                add_downsample=not is_final_block,704                resnet_eps=1e-6,705                downsample_padding=0,706                resnet_act_fn=act_fn,707                resnet_groups=norm_num_groups,708                attn_num_head_channels=None,709                temb_channels=None,710            )711            pre_down_blocks.append(pre_down_block)712        self.down_blocks.append(nn.Sequential(*pre_down_blocks))713        self.pre_lora_layers.append(714            get_down_block(715                lora_pre_conv_types[0],716                num_layers=self.lora_pre_conv_layers_per_block,717                in_channels=lora_block_in_channels[0],718                out_channels=(719                    lora_block_out_channels[0] 720                    if lora_control_channels[0] is None 721                    else lora_control_channels[0]),722                add_downsample=False,723                resnet_eps=1e-6,724                downsample_padding=0,725                resnet_act_fn=act_fn,726                resnet_groups=norm_num_groups,727                attn_num_head_channels=None,728                temb_channels=None,729                resnet_kernel_size=lora_pre_conv_layers_kernel_size,730            ) if not lora_pre_conv_skipped else nn.Identity()731        )732        self.lora_layers.append(733            nn.ModuleList([734                lora_control_cls(735                    lora_block_out_channels[0], 736                    cross_attention_dim=cross_attention_dim, 737                    rank=lora_rank, 738                    control_rank=lora_control_rank,739                    post_add=lora_post_add,740                    concat_hidden=lora_concat_hidden,741                    control_channels=lora_control_channels[0],742                    control_self_add=lora_control_self_add,743                    key_states_skipped=lora_key_states_skipped,744                    value_states_skipped=lora_value_states_skipped,745                    output_states_skipped=lora_output_states_skipped)746                for cross_attention_dim in lora_cross_attention_dims[0]747            ])748        )749        750        # down751        output_channel = lora_block_in_channels[0]752        for i, down_block_type in enumerate(lora_pre_down_block_types):753            if i == 0:754                continue755            input_channel = output_channel756            output_channel = lora_block_in_channels[i]757 758            down_block = get_down_block(759                down_block_type,760                num_layers=self.lora_pre_down_layers_per_block,761                in_channels=input_channel,762                out_channels=output_channel,763                add_downsample=True,764                resnet_eps=1e-6,765                downsample_padding=0,766                resnet_act_fn=act_fn,767                resnet_groups=norm_num_groups,768                attn_num_head_channels=None,769                temb_channels=None,770            )771            self.down_blocks.append(down_block)772 773            self.pre_lora_layers.append(774                get_down_block(775                    lora_pre_conv_types[i],776                    num_layers=self.lora_pre_conv_layers_per_block,777                    in_channels=output_channel,778                    out_channels=(779                        lora_block_out_channels[i] 780                        if lora_control_channels[i] is None 781                        else lora_control_channels[i]),782                    add_downsample=False,783                    resnet_eps=1e-6,784                    downsample_padding=0,785                    resnet_act_fn=act_fn,786                    resnet_groups=norm_num_groups,787                    attn_num_head_channels=None,788                    temb_channels=None,789                    resnet_kernel_size=lora_pre_conv_layers_kernel_size,790                ) if not lora_pre_conv_skipped else nn.Identity()791            )792            self.lora_layers.append(793                nn.ModuleList([794                    lora_control_cls(795                        lora_block_out_channels[i], 796                        cross_attention_dim=cross_attention_dim, 797                        rank=lora_rank, 798                        control_rank=lora_control_rank,799                        post_add=lora_post_add,800                        concat_hidden=lora_concat_hidden,801                        control_channels=lora_control_channels[i],802                        control_self_add=lora_control_self_add,803                        key_states_skipped=lora_key_states_skipped,804                        value_states_skipped=lora_value_states_skipped,805                        output_states_skipped=lora_output_states_skipped)806                    for cross_attention_dim in lora_cross_attention_dims[i]807                ])808            )809 810    def forward(self, x: torch.FloatTensor, return_dict: bool = True) -> Union[ControlLoRAOutput, Tuple]:811        lora_layer: ControlLoRACrossAttnProcessor812        813        orig_dtype = x.dtype814        dtype = self.conv_in.weight.dtype815 816        h = x.to(dtype)817        h = self.conv_in(h)818        control_states_list = []819 820        # down821        for down_block, pre_lora_layer, lora_layer_list in zip(822            self.down_blocks, self.pre_lora_layers, self.lora_layers):823            h = down_block(h)824            control_states = pre_lora_layer(h)825            if isinstance(control_states, tuple):826                control_states = control_states[0]827            control_states = control_states.to(orig_dtype)828            for lora_layer in lora_layer_list:829                lora_layer.inject_control_states(control_states)830            control_states_list.append(control_states)831 832        if not return_dict:833            return tuple(control_states_list)834 835        return ControlLoRAOutput(control_states=tuple(control_states_list))836