Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
stepvideo_dit.py940 linesDownload Raw Back to models
1# Copyright 2025 StepFun Inc. All Rights Reserved.2# 3# Permission is hereby granted, free of charge, to any person obtaining a copy4# of this software and associated documentation files (the "Software"), to deal5# in the Software without restriction, including without limitation the rights6# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell7# copies of the Software, and to permit persons to whom the Software is8# furnished to do so, subject to the following conditions:9#10# The above copyright notice and this permission notice shall be included in all11# copies or substantial portions of the Software.12# ==============================================================================13from typing import Dict, Optional, Tuple, Union, List14import torch, math15from torch import nn16from einops import rearrange, repeat17from tqdm import tqdm18 19 20class RMSNorm(nn.Module):21    def __init__(22        self,23        dim: int,24        elementwise_affine=True,25        eps: float = 1e-6,26        device=None,27        dtype=None,28    ):29        """30        Initialize the RMSNorm normalization layer.31 32        Args:33            dim (int): The dimension of the input tensor.34            eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.35 36        Attributes:37            eps (float): A small value added to the denominator for numerical stability.38            weight (nn.Parameter): Learnable scaling parameter.39 40        """41        factory_kwargs = {"device": device, "dtype": dtype}42        super().__init__()43        self.eps = eps44        if elementwise_affine:45            self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))46 47    def _norm(self, x):48        """49        Apply the RMSNorm normalization to the input tensor.50 51        Args:52            x (torch.Tensor): The input tensor.53 54        Returns:55            torch.Tensor: The normalized tensor.56 57        """58        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)59 60    def forward(self, x):61        """62        Forward pass through the RMSNorm layer.63 64        Args:65            x (torch.Tensor): The input tensor.66 67        Returns:68            torch.Tensor: The output tensor after applying RMSNorm.69 70        """71        output = self._norm(x.float()).type_as(x)72        if hasattr(self, "weight"):73            output = output * self.weight74        return output75    76 77ACTIVATION_FUNCTIONS = {78    "swish": nn.SiLU(),79    "silu": nn.SiLU(),80    "mish": nn.Mish(),81    "gelu": nn.GELU(),82    "relu": nn.ReLU(),83}84 85 86def get_activation(act_fn: str) -> nn.Module:87    """Helper function to get activation function from string.88 89    Args:90        act_fn (str): Name of activation function.91 92    Returns:93        nn.Module: Activation function.94    """95 96    act_fn = act_fn.lower()97    if act_fn in ACTIVATION_FUNCTIONS:98        return ACTIVATION_FUNCTIONS[act_fn]99    else:100        raise ValueError(f"Unsupported activation function: {act_fn}")101 102 103def get_timestep_embedding(104    timesteps: torch.Tensor,105    embedding_dim: int,106    flip_sin_to_cos: bool = False,107    downscale_freq_shift: float = 1,108    scale: float = 1,109    max_period: int = 10000,110):111    """112    This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.113 114    :param timesteps: a 1-D Tensor of N indices, one per batch element.115                      These may be fractional.116    :param embedding_dim: the dimension of the output. :param max_period: controls the minimum frequency of the117    embeddings. :return: an [N x dim] Tensor of positional embeddings.118    """119    assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"120 121    half_dim = embedding_dim // 2122    exponent = -math.log(max_period) * torch.arange(123        start=0, end=half_dim, dtype=torch.float32, device=timesteps.device124    )125    exponent = exponent / (half_dim - downscale_freq_shift)126 127    emb = torch.exp(exponent)128    emb = timesteps[:, None].float() * emb[None, :]129 130    # scale embeddings131    emb = scale * emb132 133    # concat sine and cosine embeddings134    emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)135 136    # flip sine and cosine embeddings137    if flip_sin_to_cos:138        emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)139 140    # zero pad141    if embedding_dim % 2 == 1:142        emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))143    return emb144 145 146class Timesteps(nn.Module):147    def __init__(self, num_channels: int, flip_sin_to_cos: bool, downscale_freq_shift: float):148        super().__init__()149        self.num_channels = num_channels150        self.flip_sin_to_cos = flip_sin_to_cos151        self.downscale_freq_shift = downscale_freq_shift152 153    def forward(self, timesteps):154        t_emb = get_timestep_embedding(155            timesteps,156            self.num_channels,157            flip_sin_to_cos=self.flip_sin_to_cos,158            downscale_freq_shift=self.downscale_freq_shift,159        )160        return t_emb161 162 163class TimestepEmbedding(nn.Module):164    def __init__(165        self,166        in_channels: int,167        time_embed_dim: int,168        act_fn: str = "silu",169        out_dim: int = None,170        post_act_fn: Optional[str] = None,171        cond_proj_dim=None,172        sample_proj_bias=True173    ):174        super().__init__()175        linear_cls = nn.Linear176 177        self.linear_1 = linear_cls(178                in_channels, 179                time_embed_dim, 180                bias=sample_proj_bias,181            )182 183        if cond_proj_dim is not None:184            self.cond_proj = linear_cls(185                    cond_proj_dim, 186                    in_channels, 187                    bias=False,188                )189        else:190            self.cond_proj = None191 192        self.act = get_activation(act_fn)193 194        if out_dim is not None:195            time_embed_dim_out = out_dim196        else:197            time_embed_dim_out = time_embed_dim198            199        self.linear_2 = linear_cls(200                time_embed_dim, 201                time_embed_dim_out, 202                bias=sample_proj_bias, 203            )204 205        if post_act_fn is None:206            self.post_act = None207        else:208            self.post_act = get_activation(post_act_fn)209 210    def forward(self, sample, condition=None):211        if condition is not None:212            sample = sample + self.cond_proj(condition)213        sample = self.linear_1(sample)214 215        if self.act is not None:216            sample = self.act(sample)217 218        sample = self.linear_2(sample)219 220        if self.post_act is not None:221            sample = self.post_act(sample)222        return sample223 224 225class PixArtAlphaCombinedTimestepSizeEmbeddings(nn.Module):226    def __init__(self, embedding_dim, size_emb_dim, use_additional_conditions: bool = False):227        super().__init__()228 229        self.outdim = size_emb_dim230        self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)231        self.timestep_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)232 233        self.use_additional_conditions = use_additional_conditions234        if self.use_additional_conditions:235            self.additional_condition_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)236            self.resolution_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=size_emb_dim)237            self.nframe_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)238            self.fps_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=embedding_dim)239 240    def forward(self, timestep, resolution=None, nframe=None, fps=None):241        hidden_dtype = timestep.dtype242 243        timesteps_proj = self.time_proj(timestep)244        timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=hidden_dtype))  # (N, D)245 246        if self.use_additional_conditions:247            batch_size = timestep.shape[0]248            resolution_emb = self.additional_condition_proj(resolution.flatten()).to(hidden_dtype)249            resolution_emb = self.resolution_embedder(resolution_emb).reshape(batch_size, -1)250            nframe_emb = self.additional_condition_proj(nframe.flatten()).to(hidden_dtype)251            nframe_emb = self.nframe_embedder(nframe_emb).reshape(batch_size, -1)252            conditioning = timesteps_emb + resolution_emb + nframe_emb253 254            if fps is not None:255                fps_emb = self.additional_condition_proj(fps.flatten()).to(hidden_dtype)256                fps_emb = self.fps_embedder(fps_emb).reshape(batch_size, -1)257                conditioning = conditioning + fps_emb258        else:259            conditioning = timesteps_emb260 261        return conditioning262 263 264class AdaLayerNormSingle(nn.Module):265    r"""266        Norm layer adaptive layer norm single (adaLN-single).267 268        As proposed in PixArt-Alpha (see: https://arxiv.org/abs/2310.00426; Section 2.3).269 270        Parameters:271            embedding_dim (`int`): The size of each embedding vector.272            use_additional_conditions (`bool`): To use additional conditions for normalization or not.273    """274    def __init__(self, embedding_dim: int, use_additional_conditions: bool = False, time_step_rescale=1000):275        super().__init__()276 277        self.emb = PixArtAlphaCombinedTimestepSizeEmbeddings(278            embedding_dim, size_emb_dim=embedding_dim // 2, use_additional_conditions=use_additional_conditions279        )280 281        self.silu = nn.SiLU()282        self.linear = nn.Linear(embedding_dim, 6 * embedding_dim, bias=True)283 284        self.time_step_rescale = time_step_rescale  ## timestep usually in [0, 1], we rescale it to [0,1000] for stability285 286    def forward(287        self,288        timestep: torch.Tensor,289        added_cond_kwargs: Dict[str, torch.Tensor] = None,290    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:291        embedded_timestep = self.emb(timestep*self.time_step_rescale, **added_cond_kwargs)292 293        out = self.linear(self.silu(embedded_timestep))294 295        return out, embedded_timestep296    297 298class PixArtAlphaTextProjection(nn.Module):299    """300    Projects caption embeddings. Also handles dropout for classifier-free guidance.301 302    Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py303    """304 305    def __init__(self, in_features, hidden_size):306        super().__init__()307        self.linear_1 = nn.Linear(308                in_features, 309                hidden_size, 310                bias=True, 311            )        312        self.act_1 = nn.GELU(approximate="tanh")313        self.linear_2 = nn.Linear(314                hidden_size, 315                hidden_size, 316                bias=True, 317            )318 319    def forward(self, caption):320        hidden_states = self.linear_1(caption)321        hidden_states = self.act_1(hidden_states)322        hidden_states = self.linear_2(hidden_states)323        return hidden_states324 325 326class Attention(nn.Module):327    def __init__(self):328        super().__init__()329    330    def attn_processor(self, attn_type):331        if attn_type == 'torch':332            return self.torch_attn_func333        elif attn_type == 'parallel':334            return self.parallel_attn_func335        else:336            raise Exception('Not supported attention type...')337 338    def torch_attn_func(339        self,340        q,341        k,342        v,343        attn_mask=None,344        causal=False,345        drop_rate=0.0,346        **kwargs347    ):348 349        if attn_mask is not None and attn_mask.dtype != torch.bool:350            attn_mask = attn_mask.to(q.dtype)351            352        if attn_mask is not None and attn_mask.ndim == 3:   ## no head353            n_heads = q.shape[2]354            attn_mask = attn_mask.unsqueeze(1).repeat(1, n_heads, 1, 1)355        356        q, k, v = map(lambda x: rearrange(x, 'b s h d -> b h s d'), (q, k, v))357        if attn_mask is not None:358            attn_mask = attn_mask.to(q.device)359        x = torch.nn.functional.scaled_dot_product_attention(360            q, k, v, attn_mask=attn_mask, dropout_p=drop_rate, is_causal=causal361        )362        x = rearrange(x, 'b h s d -> b s h d')363        return x        364 365 366class RoPE1D:367    def __init__(self, freq=1e4, F0=1.0, scaling_factor=1.0):368        self.base = freq369        self.F0 = F0370        self.scaling_factor = scaling_factor371        self.cache = {}372 373    def get_cos_sin(self, D, seq_len, device, dtype):374        if (D, seq_len, device, dtype) not in self.cache:375            inv_freq = 1.0 / (self.base ** (torch.arange(0, D, 2).float().to(device) / D))376            t = torch.arange(seq_len, device=device, dtype=inv_freq.dtype)377            freqs = torch.einsum("i,j->ij", t, inv_freq).to(dtype)378            freqs = torch.cat((freqs, freqs), dim=-1)379            cos = freqs.cos()  # (Seq, Dim)380            sin = freqs.sin()381            self.cache[D, seq_len, device, dtype] = (cos, sin)382        return self.cache[D, seq_len, device, dtype]383 384    @staticmethod385    def rotate_half(x):386        x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2:]387        return torch.cat((-x2, x1), dim=-1)388 389    def apply_rope1d(self, tokens, pos1d, cos, sin):390        assert pos1d.ndim == 2391        cos = torch.nn.functional.embedding(pos1d, cos)[:, :, None, :]392        sin = torch.nn.functional.embedding(pos1d, sin)[:, :, None, :]393        return (tokens * cos) + (self.rotate_half(tokens) * sin)394 395    def __call__(self, tokens, positions):396        """397        input:398            * tokens: batch_size x ntokens x nheads x dim399            * positions: batch_size x ntokens (t position of each token)400        output:401            * tokens after applying RoPE2D (batch_size x ntokens x nheads x dim)402        """403        D = tokens.size(3)404        assert positions.ndim == 2  # Batch, Seq405        cos, sin = self.get_cos_sin(D, int(positions.max()) + 1, tokens.device, tokens.dtype)406        tokens = self.apply_rope1d(tokens, positions, cos, sin)407        return tokens408 409 410class RoPE3D(RoPE1D):411    def __init__(self, freq=1e4, F0=1.0, scaling_factor=1.0):412        super(RoPE3D, self).__init__(freq, F0, scaling_factor)413        self.position_cache = {}414 415    def get_mesh_3d(self, rope_positions, bsz):416        f, h, w = rope_positions417 418        if f"{f}-{h}-{w}" not in self.position_cache:419            x = torch.arange(f, device='cpu')420            y = torch.arange(h, device='cpu')421            z = torch.arange(w, device='cpu')422            self.position_cache[f"{f}-{h}-{w}"] = torch.cartesian_prod(x, y, z).view(1, f*h*w, 3).expand(bsz, -1, 3)423        return self.position_cache[f"{f}-{h}-{w}"]424     425    def __call__(self, tokens, rope_positions, ch_split, parallel=False):426        """427        input:428            * tokens: batch_size x ntokens x nheads x dim429            * rope_positions: list of (f, h, w)430        output:431            * tokens after applying RoPE2D (batch_size x ntokens x nheads x dim)432        """433        assert sum(ch_split) == tokens.size(-1); 434 435        mesh_grid = self.get_mesh_3d(rope_positions, bsz=tokens.shape[0])436        out = []437        for i, (D, x) in enumerate(zip(ch_split, torch.split(tokens, ch_split, dim=-1))):438            cos, sin = self.get_cos_sin(D, int(mesh_grid.max()) + 1, tokens.device, tokens.dtype)439            440            if parallel:441                pass442            else:443                mesh = mesh_grid[:, :, i].clone()444            x = self.apply_rope1d(x, mesh.to(tokens.device), cos, sin)445            out.append(x)446            447        tokens = torch.cat(out, dim=-1)448        return tokens449 450 451class SelfAttention(Attention):452    def __init__(self, hidden_dim, head_dim, bias=False, with_rope=True, with_qk_norm=True, attn_type='torch'):453        super().__init__()454        self.head_dim = head_dim455        self.n_heads = hidden_dim // head_dim456        457        self.wqkv = nn.Linear(hidden_dim, hidden_dim*3, bias=bias)458        self.wo = nn.Linear(hidden_dim, hidden_dim, bias=bias)459        460        self.with_rope = with_rope461        self.with_qk_norm = with_qk_norm462        if self.with_qk_norm:463            self.q_norm = RMSNorm(head_dim, elementwise_affine=True)464            self.k_norm = RMSNorm(head_dim, elementwise_affine=True)465        466        if self.with_rope:467            self.rope_3d = RoPE3D(freq=1e4, F0=1.0, scaling_factor=1.0)468            self.rope_ch_split = [64, 32, 32]469        470        self.core_attention = self.attn_processor(attn_type=attn_type)471        self.parallel = attn_type=='parallel'472        473    def apply_rope3d(self, x, fhw_positions, rope_ch_split, parallel=True):474        x = self.rope_3d(x, fhw_positions, rope_ch_split, parallel)475        return x476        477    def forward(478        self, 479        x,480        cu_seqlens=None,481        max_seqlen=None,482        rope_positions=None,483        attn_mask=None484    ):485        xqkv = self.wqkv(x) 486        xqkv = xqkv.view(*x.shape[:-1], self.n_heads, 3*self.head_dim)487 488        xq, xk, xv = torch.split(xqkv, [self.head_dim]*3, dim=-1)  ## seq_len, n, dim489    490        if self.with_qk_norm:491            xq = self.q_norm(xq)492            xk = self.k_norm(xk)493    494        if self.with_rope:495            xq = self.apply_rope3d(xq, rope_positions, self.rope_ch_split, parallel=self.parallel)496            xk = self.apply_rope3d(xk, rope_positions, self.rope_ch_split, parallel=self.parallel)497            498        output = self.core_attention(499                    xq,500                    xk,501                    xv,502                    cu_seqlens=cu_seqlens,503                    max_seqlen=max_seqlen,504                    attn_mask=attn_mask505                )506        output = rearrange(output, 'b s h d -> b s (h d)')507        output = self.wo(output)508        509        return output510    511    512class CrossAttention(Attention):513    def __init__(self, hidden_dim, head_dim, bias=False, with_qk_norm=True, attn_type='torch'):514        super().__init__()515        self.head_dim = head_dim516        self.n_heads = hidden_dim // head_dim517        518        self.wq = nn.Linear(hidden_dim, hidden_dim, bias=bias)519        self.wkv = nn.Linear(hidden_dim, hidden_dim*2, bias=bias)520        self.wo = nn.Linear(hidden_dim, hidden_dim, bias=bias)521        522        self.with_qk_norm = with_qk_norm523        if self.with_qk_norm:524            self.q_norm = RMSNorm(head_dim, elementwise_affine=True)525            self.k_norm = RMSNorm(head_dim, elementwise_affine=True)526        527        self.core_attention = self.attn_processor(attn_type=attn_type)528 529    def forward(530            self, 531            x: torch.Tensor,532            encoder_hidden_states: torch.Tensor,533            attn_mask=None534        ):535        xq = self.wq(x) 536        xq = xq.view(*xq.shape[:-1], self.n_heads, self.head_dim)537        538        xkv = self.wkv(encoder_hidden_states)539        xkv = xkv.view(*xkv.shape[:-1], self.n_heads, 2*self.head_dim)540 541        xk, xv = torch.split(xkv, [self.head_dim]*2, dim=-1)  ## seq_len, n, dim542    543        if self.with_qk_norm:544            xq = self.q_norm(xq)545            xk = self.k_norm(xk)546 547        output = self.core_attention(548                    xq,549                    xk,550                    xv,551                    attn_mask=attn_mask552                )553        554        output = rearrange(output, 'b s h d -> b s (h d)')555        output = self.wo(output)556        557        return output558 559    560class GELU(nn.Module):561    r"""562    GELU activation function with tanh approximation support with `approximate="tanh"`.563 564    Parameters:565        dim_in (`int`): The number of channels in the input.566        dim_out (`int`): The number of channels in the output.567        approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.568        bias (`bool`, defaults to True): Whether to use a bias in the linear layer.569    """570 571    def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True):572        super().__init__()573        self.proj = nn.Linear(dim_in, dim_out, bias=bias)574        self.approximate = approximate575 576    def gelu(self, gate: torch.Tensor) -> torch.Tensor:577        return torch.nn.functional.gelu(gate, approximate=self.approximate)578 579    def forward(self, hidden_states):580        hidden_states = self.proj(hidden_states)581        hidden_states = self.gelu(hidden_states)582        return hidden_states583    584    585class FeedForward(nn.Module):586    def __init__(587        self, 588        dim: int,589        inner_dim: Optional[int] = None,590        dim_out: Optional[int] = None,591        mult: int = 4,592        bias: bool = False,593    ):594        super().__init__()595        inner_dim = dim*mult if inner_dim is None else inner_dim596        dim_out = dim if dim_out is None else dim_out597        self.net = nn.ModuleList([598            GELU(dim, inner_dim, approximate="tanh", bias=bias),599            nn.Identity(),600            nn.Linear(inner_dim, dim_out, bias=bias)601        ])602        603        604    def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:605        for module in self.net:606            hidden_states = module(hidden_states)607        return hidden_states608    609 610def modulate(x, scale, shift):611    x = x * (1 + scale) + shift612    return x613 614 615def gate(x, gate):616    x = gate * x617    return x618 619 620class StepVideoTransformerBlock(nn.Module):621    r"""622    A basic Transformer block.623 624    Parameters:625        dim (`int`): The number of channels in the input and output.626        num_attention_heads (`int`): The number of heads to use for multi-head attention.627        attention_head_dim (`int`): The number of channels in each head.628        dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.629        cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.630        activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.631        num_embeds_ada_norm (:632            obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`.633        attention_bias (:634            obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.635        only_cross_attention (`bool`, *optional*):636            Whether to use only cross-attention layers. In this case two cross attention layers are used.637        double_self_attention (`bool`, *optional*):638            Whether to use two self-attention layers. In this case no cross attention layers are used.639        upcast_attention (`bool`, *optional*):640            Whether to upcast the attention computation to float32. This is useful for mixed precision training.641        norm_elementwise_affine (`bool`, *optional*, defaults to `True`):642            Whether to use learnable elementwise affine parameters for normalization.643        norm_type (`str`, *optional*, defaults to `"layer_norm"`):644            The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.645        final_dropout (`bool` *optional*, defaults to False):646            Whether to apply a final dropout after the last feed-forward layer.647        attention_type (`str`, *optional*, defaults to `"default"`):648            The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`.649        positional_embeddings (`str`, *optional*, defaults to `None`):650            The type of positional embeddings to apply to.651        num_positional_embeddings (`int`, *optional*, defaults to `None`):652            The maximum number of positional embeddings to apply.653    """654 655    def __init__(656        self,657        dim: int,658        attention_head_dim: int,659        norm_eps: float = 1e-5,660        ff_inner_dim: Optional[int] = None,661        ff_bias: bool = False,662        attention_type: str = 'parallel'663    ):664        super().__init__()665        self.dim = dim666        self.norm1 = nn.LayerNorm(dim, eps=norm_eps)667        self.attn1 = SelfAttention(dim, attention_head_dim, bias=False, with_rope=True, with_qk_norm=True, attn_type=attention_type)668        669        self.norm2 = nn.LayerNorm(dim, eps=norm_eps)670        self.attn2 = CrossAttention(dim, attention_head_dim, bias=False, with_qk_norm=True, attn_type='torch')671 672        self.ff = FeedForward(dim=dim, inner_dim=ff_inner_dim, dim_out=dim, bias=ff_bias)673 674        self.scale_shift_table = nn.Parameter(torch.randn(6, dim) /dim**0.5)675 676    @torch.no_grad()677    def forward(678        self,679        q: torch.Tensor,680        kv: Optional[torch.Tensor] = None,681        timestep: Optional[torch.LongTensor] =  None,682        attn_mask = None,683        rope_positions: list = None, 684    ) -> torch.Tensor:685        shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (686            torch.clone(chunk) for chunk in (self.scale_shift_table[None].to(dtype=q.dtype, device=q.device) + timestep.reshape(-1, 6, self.dim)).chunk(6, dim=1)687        )688        689        scale_shift_q = modulate(self.norm1(q), scale_msa, shift_msa)690 691        attn_q = self.attn1(692            scale_shift_q,693            rope_positions=rope_positions694        )695 696        q = gate(attn_q, gate_msa) + q697        698        attn_q = self.attn2(699                q,700                kv,701                attn_mask702            )703 704        q = attn_q + q705 706        scale_shift_q = modulate(self.norm2(q), scale_mlp, shift_mlp)707 708        ff_output = self.ff(scale_shift_q)709        710        q = gate(ff_output, gate_mlp) + q711        712        return q713    714    715class PatchEmbed(nn.Module):716    """2D Image to Patch Embedding"""717 718    def __init__(719        self,720        patch_size=64,721        in_channels=3,722        embed_dim=768,723        layer_norm=False,724        flatten=True,725        bias=True,726    ):727        super().__init__()728 729        self.flatten = flatten730        self.layer_norm = layer_norm731 732        self.proj = nn.Conv2d(733            in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size, bias=bias734        )735 736    def forward(self, latent):737        latent = self.proj(latent).to(latent.dtype)   738        if self.flatten:739            latent = latent.flatten(2).transpose(1, 2)  # BCHW -> BNC740        if self.layer_norm:741            latent = self.norm(latent)742 743        return latent744 745 746class StepVideoModel(torch.nn.Module):747    def __init__(748        self,749        num_attention_heads: int = 48,750        attention_head_dim: int = 128,751        in_channels: int = 64,752        out_channels: Optional[int] = 64,753        num_layers: int = 48,754        dropout: float = 0.0,755        patch_size: int = 1,756        norm_type: str = "ada_norm_single",757        norm_elementwise_affine: bool = False,758        norm_eps: float = 1e-6,759        use_additional_conditions: Optional[bool] = False,760        caption_channels: Optional[Union[int, List, Tuple]] = [6144, 1024],761        attention_type: Optional[str] = "torch",762    ):763        super().__init__()764 765        # Set some common variables used across the board.766        self.inner_dim = num_attention_heads * attention_head_dim767        self.out_channels = in_channels if out_channels is None else out_channels768 769        self.use_additional_conditions = use_additional_conditions770 771        self.pos_embed = PatchEmbed(772            patch_size=patch_size,773            in_channels=in_channels,774            embed_dim=self.inner_dim,775        )776 777        self.transformer_blocks = nn.ModuleList(778            [779                StepVideoTransformerBlock(780                    dim=self.inner_dim,781                    attention_head_dim=attention_head_dim,782                    attention_type=attention_type783                )784                for _ in range(num_layers)785            ]786        )787 788        # 3. Output blocks.789        self.norm_out = nn.LayerNorm(self.inner_dim, eps=norm_eps, elementwise_affine=norm_elementwise_affine)790        self.scale_shift_table = nn.Parameter(torch.randn(2, self.inner_dim) / self.inner_dim**0.5)791        self.proj_out = nn.Linear(self.inner_dim, patch_size * patch_size * self.out_channels)792        self.patch_size = patch_size793 794        self.adaln_single = AdaLayerNormSingle(795            self.inner_dim, use_additional_conditions=self.use_additional_conditions796        )797 798        if isinstance(caption_channels, int):799            caption_channel = caption_channels800        else:801            caption_channel, clip_channel = caption_channels802            self.clip_projection = nn.Linear(clip_channel, self.inner_dim) 803 804        self.caption_norm = nn.LayerNorm(caption_channel,  eps=norm_eps, elementwise_affine=norm_elementwise_affine)805        806        self.caption_projection = PixArtAlphaTextProjection(807            in_features=caption_channel, hidden_size=self.inner_dim808        )809        810        self.parallel = attention_type=='parallel'811 812    def patchfy(self, hidden_states):813        hidden_states = rearrange(hidden_states, 'b f c h w -> (b f) c h w')814        hidden_states = self.pos_embed(hidden_states)815        return hidden_states816 817    def prepare_attn_mask(self, encoder_attention_mask, encoder_hidden_states, q_seqlen):818        kv_seqlens = encoder_attention_mask.sum(dim=1).int()819        mask = torch.zeros([len(kv_seqlens), q_seqlen, max(kv_seqlens)], dtype=torch.bool, device=encoder_attention_mask.device)820        encoder_hidden_states = encoder_hidden_states[:,: max(kv_seqlens)]821        for i, kv_len in enumerate(kv_seqlens):822            mask[i, :, :kv_len] = 1823        return encoder_hidden_states, mask824        825        826    def block_forward(827        self,828        hidden_states,829        encoder_hidden_states=None,830        timestep=None,831        rope_positions=None,832        attn_mask=None,833        parallel=True834    ):835        for block in tqdm(self.transformer_blocks, desc="Transformer blocks"):836            hidden_states = block(837                hidden_states,838                encoder_hidden_states,839                timestep=timestep,840                attn_mask=attn_mask,841                rope_positions=rope_positions842            )843 844        return hidden_states845        846 847    @torch.inference_mode()848    def forward(849        self,850        hidden_states: torch.Tensor,851        encoder_hidden_states: Optional[torch.Tensor] = None,852        encoder_hidden_states_2: Optional[torch.Tensor] = None,853        timestep: Optional[torch.LongTensor] = None,854        added_cond_kwargs: Dict[str, torch.Tensor] = None,855        encoder_attention_mask: Optional[torch.Tensor] = None,856        fps: torch.Tensor=None,857        return_dict: bool = False,858    ):859        assert hidden_states.ndim==5; "hidden_states's shape should be (bsz, f, ch, h ,w)"860 861        bsz, frame, _, height, width = hidden_states.shape862        height, width = height // self.patch_size, width // self.patch_size863                864        hidden_states = self.patchfy(hidden_states) 865        len_frame = hidden_states.shape[1]866                867        if self.use_additional_conditions:868            added_cond_kwargs = {869                "resolution": torch.tensor([(height, width)]*bsz, device=hidden_states.device, dtype=hidden_states.dtype),870                "nframe": torch.tensor([frame]*bsz, device=hidden_states.device, dtype=hidden_states.dtype),871                "fps": fps872            }    873        else:874            added_cond_kwargs = {}875        876        timestep, embedded_timestep = self.adaln_single(877            timestep, added_cond_kwargs=added_cond_kwargs878        )879 880        encoder_hidden_states = self.caption_projection(self.caption_norm(encoder_hidden_states))881        882        if encoder_hidden_states_2 is not None and hasattr(self, 'clip_projection'):883            clip_embedding = self.clip_projection(encoder_hidden_states_2)884            encoder_hidden_states = torch.cat([clip_embedding, encoder_hidden_states], dim=1)885 886        hidden_states = rearrange(hidden_states, '(b f) l d->  b (f l) d', b=bsz, f=frame, l=len_frame).contiguous()887        encoder_hidden_states, attn_mask = self.prepare_attn_mask(encoder_attention_mask, encoder_hidden_states, q_seqlen=frame*len_frame)888        889        hidden_states = self.block_forward(890            hidden_states,891            encoder_hidden_states,892            timestep=timestep,893            rope_positions=[frame, height, width],894            attn_mask=attn_mask,895            parallel=self.parallel896        )897        898        hidden_states = rearrange(hidden_states, 'b (f l) d -> (b f) l d', b=bsz, f=frame, l=len_frame)899        900        embedded_timestep = repeat(embedded_timestep, 'b d -> (b f) d', f=frame).contiguous()901        902        shift, scale = (self.scale_shift_table[None].to(dtype=embedded_timestep.dtype, device=embedded_timestep.device) + embedded_timestep[:, None]).chunk(2, dim=1)903        hidden_states = self.norm_out(hidden_states)904        # Modulation905        hidden_states = hidden_states * (1 + scale) + shift906        hidden_states = self.proj_out(hidden_states)907        908        # unpatchify909        hidden_states = hidden_states.reshape(910            shape=(-1, height, width, self.patch_size, self.patch_size, self.out_channels)911        )912        913        hidden_states = rearrange(hidden_states, 'n h w p q c -> n c h p w q')914        output = hidden_states.reshape(915            shape=(-1, self.out_channels, height * self.patch_size, width * self.patch_size)916        )917 918        output = rearrange(output, '(b f) c h w -> b f c h w', f=frame)919 920        if return_dict:921            return {'x': output}922        return output923    924    @staticmethod925    def state_dict_converter():926        return StepVideoDiTStateDictConverter()927 928 929class StepVideoDiTStateDictConverter:930    def __init__(self):931        super().__init__()932 933    def from_diffusers(self, state_dict):934        return state_dict935    936    def from_civitai(self, state_dict):937        return state_dict938 939    940    
hugging-apps/echo-memory ยท Team Ai