hugging-apps/echo-memory
0
1from typing import Optional2 3import torch, math4import torch.nn5from einops import rearrange6from torch import nn7from functools import partial8from einops import rearrange9 10 11 12def attention(q, k, v, attn_mask, mode="torch"):13 q = q.transpose(1, 2)14 k = k.transpose(1, 2)15 v = v.transpose(1, 2)16 x = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)17 x = rearrange(x, "b n s d -> b s (n d)")18 return x19 20 21 22class MLP(nn.Module):23 """MLP as used in Vision Transformer, MLP-Mixer and related networks"""24 25 def __init__(26 self,27 in_channels,28 hidden_channels=None,29 out_features=None,30 act_layer=nn.GELU,31 norm_layer=None,32 bias=True,33 drop=0.0,34 use_conv=False,35 device=None,36 dtype=None,37 ):38 super().__init__()39 out_features = out_features or in_channels40 hidden_channels = hidden_channels or in_channels41 bias = (bias, bias)42 drop_probs = (drop, drop)43 linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear44 45 self.fc1 = linear_layer(46 in_channels, hidden_channels, bias=bias[0], device=device, dtype=dtype47 )48 self.act = act_layer()49 self.drop1 = nn.Dropout(drop_probs[0])50 self.norm = (51 norm_layer(hidden_channels, device=device, dtype=dtype)52 if norm_layer is not None53 else nn.Identity()54 )55 self.fc2 = linear_layer(56 hidden_channels, out_features, bias=bias[1], device=device, dtype=dtype57 )58 self.drop2 = nn.Dropout(drop_probs[1])59 60 def forward(self, x):61 x = self.fc1(x)62 x = self.act(x)63 x = self.drop1(x)64 x = self.norm(x)65 x = self.fc2(x)66 x = self.drop2(x)67 return x68 69 70class TextProjection(nn.Module):71 """72 Projects text embeddings. Also handles dropout for classifier-free guidance.73 74 Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py75 """76 77 def __init__(self, in_channels, hidden_size, act_layer, dtype=None, device=None):78 factory_kwargs = {"dtype": dtype, "device": device}79 super().__init__()80 self.linear_1 = nn.Linear(81 in_features=in_channels,82 out_features=hidden_size,83 bias=True,84 **factory_kwargs,85 )86 self.act_1 = act_layer()87 self.linear_2 = nn.Linear(88 in_features=hidden_size,89 out_features=hidden_size,90 bias=True,91 **factory_kwargs,92 )93 94 def forward(self, caption):95 hidden_states = self.linear_1(caption)96 hidden_states = self.act_1(hidden_states)97 hidden_states = self.linear_2(hidden_states)98 return hidden_states99 100 101class TimestepEmbedder(nn.Module):102 """103 Embeds scalar timesteps into vector representations.104 """105 106 def __init__(107 self,108 hidden_size,109 act_layer,110 frequency_embedding_size=256,111 max_period=10000,112 out_size=None,113 dtype=None,114 device=None,115 ):116 factory_kwargs = {"dtype": dtype, "device": device}117 super().__init__()118 self.frequency_embedding_size = frequency_embedding_size119 self.max_period = max_period120 if out_size is None:121 out_size = hidden_size122 123 self.mlp = nn.Sequential(124 nn.Linear(125 frequency_embedding_size, hidden_size, bias=True, **factory_kwargs126 ),127 act_layer(),128 nn.Linear(hidden_size, out_size, bias=True, **factory_kwargs),129 )130 nn.init.normal_(self.mlp[0].weight, std=0.02) # type: ignore131 nn.init.normal_(self.mlp[2].weight, std=0.02) # type: ignore132 133 @staticmethod134 def timestep_embedding(t, dim, max_period=10000):135 """136 Create sinusoidal timestep embeddings.137 138 Args:139 t (torch.Tensor): a 1-D Tensor of N indices, one per batch element. These may be fractional.140 dim (int): the dimension of the output.141 max_period (int): controls the minimum frequency of the embeddings.142 143 Returns:144 embedding (torch.Tensor): An (N, D) Tensor of positional embeddings.145 146 .. ref_link: https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py147 """148 half = dim // 2149 freqs = torch.exp(150 -math.log(max_period)151 * torch.arange(start=0, end=half, dtype=torch.float32)152 / half153 ).to(device=t.device)154 args = t[:, None].float() * freqs[None]155 embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)156 if dim % 2:157 embedding = torch.cat(158 [embedding, torch.zeros_like(embedding[:, :1])], dim=-1159 )160 return embedding161 162 def forward(self, t):163 t_freq = self.timestep_embedding(164 t, self.frequency_embedding_size, self.max_period165 ).type(t.dtype) # type: ignore166 t_emb = self.mlp(t_freq)167 return t_emb168 169 170def apply_gate(x, gate=None, tanh=False):171 """AI is creating summary for apply_gate172 173 Args:174 x (torch.Tensor): input tensor.175 gate (torch.Tensor, optional): gate tensor. Defaults to None.176 tanh (bool, optional): whether to use tanh function. Defaults to False.177 178 Returns:179 torch.Tensor: the output tensor after apply gate.180 """181 if gate is None:182 return x183 if tanh:184 return x * gate.unsqueeze(1).tanh()185 else:186 return x * gate.unsqueeze(1)187 188 189class RMSNorm(nn.Module):190 def __init__(191 self,192 dim: int,193 elementwise_affine=True,194 eps: float = 1e-6,195 device=None,196 dtype=None,197 ):198 """199 Initialize the RMSNorm normalization layer.200 201 Args:202 dim (int): The dimension of the input tensor.203 eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.204 205 Attributes:206 eps (float): A small value added to the denominator for numerical stability.207 weight (nn.Parameter): Learnable scaling parameter.208 209 """210 factory_kwargs = {"device": device, "dtype": dtype}211 super().__init__()212 self.eps = eps213 if elementwise_affine:214 self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))215 216 def _norm(self, x):217 """218 Apply the RMSNorm normalization to the input tensor.219 220 Args:221 x (torch.Tensor): The input tensor.222 223 Returns:224 torch.Tensor: The normalized tensor.225 226 """227 return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)228 229 def forward(self, x):230 """231 Forward pass through the RMSNorm layer.232 233 Args:234 x (torch.Tensor): The input tensor.235 236 Returns:237 torch.Tensor: The output tensor after applying RMSNorm.238 239 """240 output = self._norm(x.float()).type_as(x)241 if hasattr(self, "weight"):242 output = output * self.weight243 return output244 245 246def get_norm_layer(norm_layer):247 """248 Get the normalization layer.249 250 Args:251 norm_layer (str): The type of normalization layer.252 253 Returns:254 norm_layer (nn.Module): The normalization layer.255 """256 if norm_layer == "layer":257 return nn.LayerNorm258 elif norm_layer == "rms":259 return RMSNorm260 else:261 raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")262 263 264def get_activation_layer(act_type):265 """get activation layer266 267 Args:268 act_type (str): the activation type269 270 Returns:271 torch.nn.functional: the activation layer272 """273 if act_type == "gelu":274 return lambda: nn.GELU()275 elif act_type == "gelu_tanh":276 return lambda: nn.GELU(approximate="tanh")277 elif act_type == "relu":278 return nn.ReLU279 elif act_type == "silu":280 return nn.SiLU281 else:282 raise ValueError(f"Unknown activation type: {act_type}")283 284class IndividualTokenRefinerBlock(torch.nn.Module):285 def __init__(286 self,287 hidden_size,288 heads_num,289 mlp_width_ratio: str = 4.0,290 mlp_drop_rate: float = 0.0,291 act_type: str = "silu",292 qk_norm: bool = False,293 qk_norm_type: str = "layer",294 qkv_bias: bool = True,295 need_CA: bool = False,296 dtype: Optional[torch.dtype] = None,297 device: Optional[torch.device] = None,298 ):299 factory_kwargs = {"device": device, "dtype": dtype}300 super().__init__()301 self.need_CA = need_CA302 self.heads_num = heads_num303 head_dim = hidden_size // heads_num304 mlp_hidden_dim = int(hidden_size * mlp_width_ratio)305 306 self.norm1 = nn.LayerNorm(307 hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs308 )309 self.self_attn_qkv = nn.Linear(310 hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs311 )312 qk_norm_layer = get_norm_layer(qk_norm_type)313 self.self_attn_q_norm = (314 qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)315 if qk_norm316 else nn.Identity()317 )318 self.self_attn_k_norm = (319 qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)320 if qk_norm321 else nn.Identity()322 )323 self.self_attn_proj = nn.Linear(324 hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs325 )326 327 self.norm2 = nn.LayerNorm(328 hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs329 )330 act_layer = get_activation_layer(act_type)331 self.mlp = MLP(332 in_channels=hidden_size,333 hidden_channels=mlp_hidden_dim,334 act_layer=act_layer,335 drop=mlp_drop_rate,336 **factory_kwargs,337 )338 339 self.adaLN_modulation = nn.Sequential(340 act_layer(),341 nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),342 )343 344 if self.need_CA:345 self.cross_attnblock=CrossAttnBlock(hidden_size=hidden_size,346 heads_num=heads_num,347 mlp_width_ratio=mlp_width_ratio,348 mlp_drop_rate=mlp_drop_rate,349 act_type=act_type,350 qk_norm=qk_norm,351 qk_norm_type=qk_norm_type,352 qkv_bias=qkv_bias,353 **factory_kwargs,)354 # Zero-initialize the modulation355 nn.init.zeros_(self.adaLN_modulation[1].weight)356 nn.init.zeros_(self.adaLN_modulation[1].bias)357 358 def forward(359 self,360 x: torch.Tensor,361 c: torch.Tensor, # timestep_aware_representations + context_aware_representations362 attn_mask: torch.Tensor = None,363 y: torch.Tensor = None,364 ):365 gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)366 367 norm_x = self.norm1(x)368 qkv = self.self_attn_qkv(norm_x)369 q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)370 # Apply QK-Norm if needed371 q = self.self_attn_q_norm(q).to(v)372 k = self.self_attn_k_norm(k).to(v)373 374 # Self-Attention375 attn = attention(q, k, v, mode="torch", attn_mask=attn_mask)376 377 x = x + apply_gate(self.self_attn_proj(attn), gate_msa)378 379 if self.need_CA:380 x = self.cross_attnblock(x, c, attn_mask, y)381 382 # FFN Layer383 x = x + apply_gate(self.mlp(self.norm2(x)), gate_mlp)384 385 return x386 387 388 389 390class CrossAttnBlock(torch.nn.Module):391 def __init__(392 self,393 hidden_size,394 heads_num,395 mlp_width_ratio: str = 4.0,396 mlp_drop_rate: float = 0.0,397 act_type: str = "silu",398 qk_norm: bool = False,399 qk_norm_type: str = "layer",400 qkv_bias: bool = True,401 dtype: Optional[torch.dtype] = None,402 device: Optional[torch.device] = None,403 ):404 factory_kwargs = {"device": device, "dtype": dtype}405 super().__init__()406 self.heads_num = heads_num407 head_dim = hidden_size // heads_num408 409 self.norm1 = nn.LayerNorm(410 hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs411 )412 self.norm1_2 = nn.LayerNorm(413 hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs414 )415 self.self_attn_q = nn.Linear(416 hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs417 )418 self.self_attn_kv = nn.Linear(419 hidden_size, hidden_size*2, bias=qkv_bias, **factory_kwargs420 )421 qk_norm_layer = get_norm_layer(qk_norm_type)422 self.self_attn_q_norm = (423 qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)424 if qk_norm425 else nn.Identity()426 )427 self.self_attn_k_norm = (428 qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)429 if qk_norm430 else nn.Identity()431 )432 self.self_attn_proj = nn.Linear(433 hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs434 )435 436 self.norm2 = nn.LayerNorm(437 hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs438 )439 act_layer = get_activation_layer(act_type)440 441 self.adaLN_modulation = nn.Sequential(442 act_layer(),443 nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),444 )445 # Zero-initialize the modulation446 nn.init.zeros_(self.adaLN_modulation[1].weight)447 nn.init.zeros_(self.adaLN_modulation[1].bias)448 449 def forward(450 self,451 x: torch.Tensor,452 c: torch.Tensor, # timestep_aware_representations + context_aware_representations453 attn_mask: torch.Tensor = None,454 y: torch.Tensor=None,455 456 ):457 gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)458 459 norm_x = self.norm1(x)460 norm_y = self.norm1_2(y)461 q = self.self_attn_q(norm_x)462 q = rearrange(q, "B L (H D) -> B L H D", H=self.heads_num)463 kv = self.self_attn_kv(norm_y)464 k, v = rearrange(kv, "B L (K H D) -> K B L H D", K=2, H=self.heads_num)465 # Apply QK-Norm if needed466 q = self.self_attn_q_norm(q).to(v)467 k = self.self_attn_k_norm(k).to(v)468 469 # Self-Attention470 attn = attention(q, k, v, mode="torch", attn_mask=attn_mask)471 472 x = x + apply_gate(self.self_attn_proj(attn), gate_msa)473 474 return x475 476 477 478class IndividualTokenRefiner(torch.nn.Module):479 def __init__(480 self,481 hidden_size,482 heads_num,483 depth,484 mlp_width_ratio: float = 4.0,485 mlp_drop_rate: float = 0.0,486 act_type: str = "silu",487 qk_norm: bool = False,488 qk_norm_type: str = "layer",489 qkv_bias: bool = True,490 need_CA:bool=False,491 dtype: Optional[torch.dtype] = None,492 device: Optional[torch.device] = None,493 ): 494 495 factory_kwargs = {"device": device, "dtype": dtype}496 super().__init__()497 self.need_CA = need_CA498 self.blocks = nn.ModuleList(499 [500 IndividualTokenRefinerBlock(501 hidden_size=hidden_size,502 heads_num=heads_num,503 mlp_width_ratio=mlp_width_ratio,504 mlp_drop_rate=mlp_drop_rate,505 act_type=act_type,506 qk_norm=qk_norm,507 qk_norm_type=qk_norm_type,508 qkv_bias=qkv_bias,509 need_CA=self.need_CA,510 **factory_kwargs,511 )512 for _ in range(depth)513 ]514 )515 516 517 def forward(518 self,519 x: torch.Tensor,520 c: torch.LongTensor,521 mask: Optional[torch.Tensor] = None,522 y:torch.Tensor=None,523 ):524 self_attn_mask = None525 if mask is not None:526 batch_size = mask.shape[0]527 seq_len = mask.shape[1]528 mask = mask.to(x.device)529 # batch_size x 1 x seq_len x seq_len530 self_attn_mask_1 = mask.view(batch_size, 1, 1, seq_len).repeat(531 1, 1, seq_len, 1532 )533 # batch_size x 1 x seq_len x seq_len534 self_attn_mask_2 = self_attn_mask_1.transpose(2, 3)535 # batch_size x 1 x seq_len x seq_len, 1 for broadcasting of heads_num536 self_attn_mask = (self_attn_mask_1 & self_attn_mask_2).bool()537 # avoids self-attention weight being NaN for padding tokens538 self_attn_mask[:, :, :, 0] = True539 540 541 for block in self.blocks:542 x = block(x, c, self_attn_mask,y)543 544 return x545 546 547class SingleTokenRefiner(torch.nn.Module):548 """549 A single token refiner block for llm text embedding refine.550 """551 def __init__(552 self,553 in_channels,554 hidden_size,555 heads_num,556 depth,557 mlp_width_ratio: float = 4.0,558 mlp_drop_rate: float = 0.0,559 act_type: str = "silu",560 qk_norm: bool = False,561 qk_norm_type: str = "layer",562 qkv_bias: bool = True,563 need_CA:bool=False,564 attn_mode: str = "torch",565 dtype: Optional[torch.dtype] = None,566 device: Optional[torch.device] = None,567 ):568 factory_kwargs = {"device": device, "dtype": dtype}569 super().__init__()570 self.attn_mode = attn_mode571 self.need_CA = need_CA572 assert self.attn_mode == "torch", "Only support 'torch' mode for token refiner."573 574 self.input_embedder = nn.Linear(575 in_channels, hidden_size, bias=True, **factory_kwargs576 )577 if self.need_CA:578 self.input_embedder_CA = nn.Linear(579 in_channels, hidden_size, bias=True, **factory_kwargs580 )581 582 act_layer = get_activation_layer(act_type)583 # Build timestep embedding layer584 self.t_embedder = TimestepEmbedder(hidden_size, act_layer, **factory_kwargs)585 # Build context embedding layer586 self.c_embedder = TextProjection(587 in_channels, hidden_size, act_layer, **factory_kwargs588 )589 590 self.individual_token_refiner = IndividualTokenRefiner(591 hidden_size=hidden_size,592 heads_num=heads_num,593 depth=depth,594 mlp_width_ratio=mlp_width_ratio,595 mlp_drop_rate=mlp_drop_rate,596 act_type=act_type,597 qk_norm=qk_norm,598 qk_norm_type=qk_norm_type,599 qkv_bias=qkv_bias,600 need_CA=need_CA,601 **factory_kwargs,602 )603 604 def forward(605 self,606 x: torch.Tensor,607 t: torch.LongTensor,608 mask: Optional[torch.LongTensor] = None,609 y: torch.LongTensor=None,610 ):611 timestep_aware_representations = self.t_embedder(t)612 613 if mask is None:614 context_aware_representations = x.mean(dim=1)615 else:616 mask_float = mask.unsqueeze(-1) # [b, s1, 1]617 context_aware_representations = (x * mask_float).sum(618 dim=1619 ) / mask_float.sum(dim=1)620 context_aware_representations = self.c_embedder(context_aware_representations)621 c = timestep_aware_representations + context_aware_representations622 623 x = self.input_embedder(x)624 if self.need_CA:625 y = self.input_embedder_CA(y)626 x = self.individual_token_refiner(x, c, mask, y)627 else:628 x = self.individual_token_refiner(x, c, mask)629 630 return x631 632 633class Qwen2Connector(torch.nn.Module):634 def __init__(635 self,636 # biclip_dim=1024,637 in_channels=3584,638 hidden_size=4096,639 heads_num=32,640 depth=2,641 need_CA=False,642 device=None,643 dtype=torch.bfloat16,644 ):645 super().__init__()646 factory_kwargs = {"device": device, "dtype":dtype}647 648 self.S =SingleTokenRefiner(in_channels=in_channels,hidden_size=hidden_size,heads_num=heads_num,depth=depth,need_CA=need_CA,**factory_kwargs)649 self.global_proj_out=nn.Linear(in_channels,768)650 651 self.scale_factor = nn.Parameter(torch.zeros(1))652 with torch.no_grad():653 self.scale_factor.data += -(1 - 0.09)654 655 def forward(self, x,t,mask):656 mask_float = mask.unsqueeze(-1) # [b, s1, 1]657 x_mean = (x * mask_float).sum(658 dim=1659 ) / mask_float.sum(dim=1) * (1 + self.scale_factor.to(dtype=x.dtype, device=x.device))660 661 global_out=self.global_proj_out(x_mean)662 encoder_hidden_states = self.S(x,t,mask)663 return encoder_hidden_states,global_out664 665 @staticmethod666 def state_dict_converter():667 return Qwen2ConnectorStateDictConverter()668 669 670class Qwen2ConnectorStateDictConverter:671 def __init__(self):672 pass673 674 def from_diffusers(self, state_dict):675 return state_dict676 677 def from_civitai(self, state_dict):678 state_dict_ = {}679 for name, param in state_dict.items():680 if name.startswith("connector."):681 name_ = name[len("connector."):]682 state_dict_[name_] = param683 return state_dict_684 