File size: 11,916 Bytes
17a8581
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5e83b2a
 
17a8581
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
import math
import os
from functools import partial
from typing import Optional, Tuple, Union

import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from torch.utils.checkpoint import checkpoint
from torch.nn.functional import scaled_dot_product_attention as slow_attn    # q, k, v: BHLc

from grn.models.rope import apply_rotary_emb
from grn.utils_t2iv.sequence_parallel import sp_all_to_all, SequenceParallelManager as sp_manager

try:
    from flash_attn.ops.rms_norm import rms_norm as rms_norm_impl
except ImportError:
    def rms_norm_impl(x, weight, epsilon):
        return (x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True).add_(epsilon))) * weight

def merge_states(states, splits, cfg):
    """
    pick key and value states for flash_attn_varlen_func
    Args:
        states: list of states
        splits: list of split sizes
        cfg: bool, use cfg or not"""
    if cfg:
        cond_len, uncond_len = 0, 0
        cond_states, uncond_states = [], []
        for stat_, split_ in zip(states, splits):
            cond, uncond = torch.split(stat_, split_, dim=2)
            cond_states.append(cond)
            uncond_states.append(uncond)
            cond_len += cond.shape[2]
            uncond_len += uncond.shape[2]
        return cond_states + uncond_states, [cond_len, uncond_len]
    else:
        cond_len = 0
        for stat_ in states:
            cond_len += stat_.shape[2]
        return states, [cond_len]

class FastRMSNorm(nn.Module):
    def __init__(self, C, eps=1e-6, elementwise_affine=True):
        super().__init__()
        self.C = C
        self.eps = eps
        self.elementwise_affine = elementwise_affine
        if self.elementwise_affine:
            self.weight = nn.Parameter(torch.ones(C))
        else:
            self.register_buffer('weight', torch.ones(C))
    
    def forward(self, x):
        src_type = x.dtype
        return rms_norm_impl(x.float(), self.weight, epsilon=self.eps).to(src_type)
    
    def extra_repr(self) -> str:
        return f'C={self.C}, eps={self.eps:g}, elementwise_affine={self.elementwise_affine}'


class WanLayerNorm(nn.LayerNorm):

    def __init__(self, dim, eps=1e-6, elementwise_affine=False):
        super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)

    def forward(self, x):
        r"""
        Args:
            x(Tensor): Shape [B, L, C]
        """
        return super().forward(x.float()).type_as(x)


class Qwen3MLP(nn.Module):
    def __init__(self, hidden_size, intermediate_size):
        super().__init__()
        self.hidden_size = hidden_size
        self.intermediate_size = intermediate_size
        self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
        self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
        self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
        self.act_fn = nn.SiLU()

    def forward(self, x):
        down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
        return down_proj

def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
    """
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    """
    batch, num_key_value_heads, slen, head_dim = hidden_states.shape
    if n_rep == 1:
        return hidden_states
    hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
    return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)

class SelfAttention(nn.Module):
    def __init__(
        self, embed_dim=768, num_heads=12, num_key_value_heads=-1,
        use_flex_attn=False, qwen_qkvo_bias=False, **kwargs,
    ):
        """
        :param embed_dim: model's width
        :param num_heads: num heads of multi-head attention
        :param proj_drop: always 0 for testing
        :param tau: always 1
        :param cos_attn: always True: during attention, q and k will be L2-normalized and scaled by a head-wise learnable parameter self.scale_mul_1H11
        """
        super().__init__()
        assert embed_dim % num_heads == 0
        assert num_key_value_heads == -1 or num_heads % num_key_value_heads == 0
        
        self.num_heads, self.head_dim = num_heads, embed_dim // num_heads
        self.num_key_value_heads = num_key_value_heads if num_key_value_heads > 0 else num_heads
        self.q_proj = nn.Linear(embed_dim, self.num_heads*self.head_dim, bias=qwen_qkvo_bias)
        self.k_proj = nn.Linear(embed_dim, self.num_key_value_heads*self.head_dim, bias=qwen_qkvo_bias)
        self.v_proj = nn.Linear(embed_dim, self.num_key_value_heads*self.head_dim, bias=qwen_qkvo_bias)
        self.o_proj = nn.Linear(self.num_heads*self.head_dim, embed_dim, bias=qwen_qkvo_bias)
        self.q_norm = FastRMSNorm(self.head_dim)
        self.k_norm = FastRMSNorm(self.head_dim)
        self.num_key_value_groups = self.num_heads // self.num_key_value_heads
        self.scale = self.head_dim**-0.5
        
        self.caching = False    # kv caching: only used during inference
        self.cached_k = {}    # kv caching: only used during inference
        self.cached_v = {}    # kv caching: only used during inference
        self.cached_split_cond_uncond = {} # only used during inference

        self.use_flex_attn = use_flex_attn
    
    def kv_caching(self, enable: bool): # kv caching: only used during inference
        self.caching = enable
        self.cached_k = {}
        self.cached_v = {}
        self.cached_split_cond_uncond = {}

    # NOTE: attn_bias_or_two_vector is None during inference
    def forward(self, x, cu_seqlens, max_seqlen, attn_bias_or_two_vector: Union[torch.Tensor, Tuple[torch.IntTensor, torch.IntTensor]], attn_fn=None, rope2d_freqs_grid=[], scale_ind=0, context_info=None, last_diffusion_step=True, ref_text_scale_inds=[], use_cfg=False, split_cond_uncond=[], **kwargs):
        # x: fp32
        B, L, C = x.shape
        hidden_states = x
        input_shape = hidden_states.shape[:-1]
        hidden_shape = (*input_shape, -1, self.head_dim)
        query_states = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).contiguous()# batch, slen, heads, head_dim
        key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).contiguous() # batch, slen, num_key_value_heads, head_dim
        value_states = self.v_proj(hidden_states).view(hidden_shape).contiguous() # batch, slen, num_key_value_heads, head_dim

        if sp_manager.sp_on():
            # Headnum need to be sharded and L needs to be gathered
            # [B, H, raw_L/sp, C] --> [B, H/sp, raw_L, C]
            sdim = 1
            gdim = 2
            L = L * sp_manager.get_sp_size()
            C = C // sp_manager.get_sp_size()
            query_states = sp_all_to_all(query_states, sdim, gdim)
            key_states = sp_all_to_all(key_states, sdim, gdim)
            value_states = sp_all_to_all(value_states, sdim, gdim)

        query_states, key_states = apply_rotary_emb(query_states, key_states, rope2d_freqs_grid)
        key_states = repeat_kv(key_states, self.num_key_value_groups)
        value_states = repeat_kv(value_states, self.num_key_value_groups)

        if attn_bias_or_two_vector is None:
            # fa4, flash_attn_func input/output should be (batch_size, seqlen, nheads, headdim)
            from flash_attn.cute import flash_attn_varlen_func
            attn_output = flash_attn_varlen_func(
                q = query_states.squeeze(0),
                k = key_states.squeeze(0),
                v = value_states.squeeze(0),
                cu_seqlens_q=cu_seqlens,
                cu_seqlens_k=cu_seqlens,
                max_seqlen_q=max_seqlen,
                max_seqlen_k=max_seqlen,
                softmax_scale=self.scale,
            )
            attn_output = attn_output[0].reshape(B, L, C).contiguous()
        else:
            # slow attn
            attn_output = slow_attn(query=query_states.transpose(1, 2), key=key_states.transpose(1, 2), value=value_states.transpose(1, 2), scale=self.scale, attn_mask=attn_bias_or_two_vector, dropout_p=0).transpose(1, 2).reshape(B, L, C)

        if sp_manager.sp_on():
            # [B, raw_L, C/sp] --> [B, raw_L/sp, C]
            sdim = 1
            gdim = 2
            attn_output = sp_all_to_all(attn_output, sdim, gdim)

        attn_output = self.o_proj(attn_output)

        return attn_output
    
class SelfAttnBlock(nn.Module):
    def __init__(
        self, embed_dim, num_heads, num_key_value_heads, mlp_ratio=4.,
        use_flex_attn=False,
        qwen_qkvo_bias=False, use_ada_layer_norm=False, **kwargs,
    ):
        super(SelfAttnBlock, self).__init__()
        self.C = embed_dim
        self.attn = SelfAttention(
            embed_dim=embed_dim, num_heads=num_heads, num_key_value_heads=num_key_value_heads,
            use_flex_attn=use_flex_attn, qwen_qkvo_bias=qwen_qkvo_bias, **kwargs,
        )
        self.mlp = Qwen3MLP(hidden_size=embed_dim, intermediate_size=round(embed_dim * mlp_ratio / 256) * 256)
        self.use_ada_layer_norm = use_ada_layer_norm
        if self.use_ada_layer_norm:
            self.modulation = nn.Parameter(torch.randn(1, 6, embed_dim) / embed_dim**0.5)
            self.input_layernorm = WanLayerNorm(embed_dim)
            self.post_attention_layernorm = WanLayerNorm(embed_dim)
        else:
            self.input_layernorm = FastRMSNorm(embed_dim)
            self.post_attention_layernorm = FastRMSNorm(embed_dim)
        
    # NOTE: attn_bias_or_two_vector is None during inference
    def forward(self, x, cu_seqlens, max_seqlen, e0, attn_bias_or_two_vector, attn_fn=None, rope2d_freqs_grid=[], scale_ind=0, context_info=None, last_diffusion_step=True, ref_text_scale_inds=[], use_cfg=False, split_cond_uncond=[], **kwargs):
        # x: [B,L,C]
        # e0: [B, L, 6, C]
        if self.use_ada_layer_norm:
            assert e0.dtype == torch.float32
            e = e0
            with torch.amp.autocast('cuda', dtype=torch.float32):
                e = (self.modulation.unsqueeze(0) + e).chunk(6, dim=2)
            residual = x
            hidden_states = x
            hidden_states = self.input_layernorm(hidden_states).float() * (1 + e[1].squeeze(2)) + e[0].squeeze(2)
            hidden_states = self.attn(hidden_states, cu_seqlens, max_seqlen, attn_bias_or_two_vector, attn_fn, rope2d_freqs_grid, scale_ind, context_info, last_diffusion_step, ref_text_scale_inds, use_cfg, split_cond_uncond, **kwargs)
            with torch.amp.autocast('cuda', dtype=torch.float32):
                hidden_states = residual + hidden_states * e[2].squeeze(2)
            # Fully Connected
            residual = hidden_states
            hidden_states = self.post_attention_layernorm(hidden_states).float() * (1 + e[4].squeeze(2)) + e[3].squeeze(2)
            hidden_states = self.mlp(hidden_states)
            with torch.amp.autocast('cuda', dtype=torch.float32):
                hidden_states = residual + hidden_states * e[5].squeeze(2)
        else:
            residual = x
            hidden_states = x
            hidden_states = self.input_layernorm(hidden_states)
            hidden_states = self.attn(hidden_states, cu_seqlens, max_seqlen, attn_bias_or_two_vector, attn_fn, rope2d_freqs_grid, scale_ind, context_info, last_diffusion_step, ref_text_scale_inds, use_cfg, split_cond_uncond, **kwargs)
            hidden_states = residual + hidden_states
            # Fully Connected
            residual = hidden_states
            hidden_states = self.post_attention_layernorm(hidden_states)
            hidden_states = self.mlp(hidden_states)
            hidden_states = residual + hidden_states
        return hidden_states