Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
step1x_connector.py684 linesDownload Raw Back to models
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