svjack/ControlNet-Pose-Chinese
5
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 