amonshano commited on
Commit
b66f552
·
verified ·
1 Parent(s): eafbe80

Add Echo-Memory codebase used for this run (CC BY 4.0, JD Echo Team) (part 3)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. code/flash-linear-attention/fla/models/sse/modeling_sse.py +437 -0
  2. code/flash-linear-attention/fla/models/transformer/__init__.py +12 -0
  3. code/flash-linear-attention/fla/models/transformer/configuration_transformer.py +85 -0
  4. code/flash-linear-attention/fla/models/transformer/modeling_transformer.py +356 -0
  5. code/flash-linear-attention/fla/models/utils.py +471 -0
  6. code/flash-linear-attention/fla/modules/__init__.py +32 -0
  7. code/flash-linear-attention/fla/modules/activations.py +555 -0
  8. code/flash-linear-attention/fla/modules/convolution.py +1167 -0
  9. code/flash-linear-attention/fla/modules/feature_map.py +298 -0
  10. code/flash-linear-attention/fla/modules/fused_bitlinear.py +633 -0
  11. code/flash-linear-attention/fla/modules/fused_cross_entropy.py +418 -0
  12. code/flash-linear-attention/fla/modules/fused_kl_div.py +322 -0
  13. code/flash-linear-attention/fla/modules/fused_linear_cross_entropy.py +630 -0
  14. code/flash-linear-attention/fla/modules/fused_norm_gate.py +1245 -0
  15. code/flash-linear-attention/fla/modules/grpo.py +412 -0
  16. code/flash-linear-attention/fla/modules/l2norm.py +287 -0
  17. code/flash-linear-attention/fla/modules/l2warp.py +37 -0
  18. code/flash-linear-attention/fla/modules/layernorm.py +1444 -0
  19. code/flash-linear-attention/fla/modules/layernorm_gated.py +527 -0
  20. code/flash-linear-attention/fla/modules/mlp.py +144 -0
  21. code/flash-linear-attention/fla/modules/parallel.py +53 -0
  22. code/flash-linear-attention/fla/modules/rotary.py +499 -0
  23. code/flash-linear-attention/fla/modules/token_shift.py +545 -0
  24. code/flash-linear-attention/fla/ops/__init__.py +54 -0
  25. code/flash-linear-attention/fla/ops/abc/__init__.py +6 -0
  26. code/flash-linear-attention/fla/ops/abc/chunk.py +1115 -0
  27. code/flash-linear-attention/fla/ops/abc/naive.py +94 -0
  28. code/flash-linear-attention/fla/ops/attn/__init__.py +6 -0
  29. code/flash-linear-attention/fla/ops/attn/decoding.py +181 -0
  30. code/flash-linear-attention/fla/ops/attn/parallel.py +728 -0
  31. code/flash-linear-attention/fla/ops/based/__init__.py +8 -0
  32. code/flash-linear-attention/fla/ops/based/fused_chunk.py +371 -0
  33. code/flash-linear-attention/fla/ops/based/naive.py +70 -0
  34. code/flash-linear-attention/fla/ops/based/parallel.py +406 -0
  35. code/flash-linear-attention/fla/ops/comba/__init__.py +7 -0
  36. code/flash-linear-attention/fla/ops/comba/chunk.py +340 -0
  37. code/flash-linear-attention/fla/ops/comba/fused_recurrent.py +330 -0
  38. code/flash-linear-attention/fla/ops/comba/utils.py +174 -0
  39. code/flash-linear-attention/fla/ops/comba/wy_fast.py +424 -0
  40. code/flash-linear-attention/fla/ops/common/__init__.py +0 -0
  41. code/flash-linear-attention/fla/ops/common/chunk_delta_h.py +533 -0
  42. code/flash-linear-attention/fla/ops/common/chunk_h.py +394 -0
  43. code/flash-linear-attention/fla/ops/common/chunk_h_parallel.py +554 -0
  44. code/flash-linear-attention/fla/ops/common/chunk_h_split.py +599 -0
  45. code/flash-linear-attention/fla/ops/common/chunk_o.py +689 -0
  46. code/flash-linear-attention/fla/ops/common/chunk_scaled_dot_kkt.py +124 -0
  47. code/flash-linear-attention/fla/ops/common/fused_chunk.py +636 -0
  48. code/flash-linear-attention/fla/ops/common/fused_recurrent.py +567 -0
  49. code/flash-linear-attention/fla/ops/delta_rule/README.md +90 -0
  50. code/flash-linear-attention/fla/ops/delta_rule/__init__.py +10 -0
code/flash-linear-attention/fla/models/sse/modeling_sse.py ADDED
@@ -0,0 +1,437 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from __future__ import annotations
3
+
4
+ import math
5
+ import warnings
6
+ from dataclasses import dataclass
7
+ from typing import TYPE_CHECKING, Optional, Tuple
8
+
9
+ import torch
10
+ import torch.nn as nn
11
+ from transformers.modeling_outputs import BaseModelOutputWithPast, MoeCausalLMOutputWithPast
12
+ from transformers.modeling_utils import PreTrainedModel
13
+ from transformers.utils import logging
14
+ from transformers.utils.deprecation import deprecate_kwarg
15
+
16
+ from fla.layers.attn import Attention
17
+ from fla.layers.sse import SSEGLA, SSEGDN
18
+ from fla.models.sse.configuration_sse import SSEConfig
19
+ from fla.models.utils import Cache, FLAGenerationMixin
20
+ from fla.modules import FusedCrossEntropyLoss, FusedLinearCrossEntropyLoss, RMSNorm
21
+ from fla.modules import GatedMLP as SSEMLP
22
+ from fla.modules.l2warp import l2_warp
23
+
24
+ if TYPE_CHECKING:
25
+ from transformers.processing_utils import Unpack
26
+
27
+
28
+ try:
29
+ from transformers.modeling_layers import GradientCheckpointingLayer
30
+ except ImportError:
31
+ from fla.models.modeling_layers import GradientCheckpointingLayer
32
+
33
+ logger = logging.get_logger(__name__)
34
+
35
+
36
+ class SSEBlock(GradientCheckpointingLayer):
37
+
38
+ def __init__(self, config: SSEConfig, layer_idx: int):
39
+ super().__init__()
40
+
41
+ self.config = config
42
+ self.layer_idx = layer_idx
43
+
44
+ self.attn_norm = (RMSNorm if config.fuse_norm else nn.RMSNorm)(config.hidden_size, eps=config.norm_eps)
45
+ if config.attn is not None and layer_idx in config.attn['layers']:
46
+ self.attn = Attention(
47
+ hidden_size=config.hidden_size,
48
+ num_heads=config.attn['num_heads'],
49
+ num_kv_heads=config.attn['num_kv_heads'],
50
+ qkv_bias=config.attn['qkv_bias'],
51
+ window_size=config.attn['window_size'],
52
+ rope_theta=config.attn['rope_theta'],
53
+ max_position_embeddings=config.max_position_embeddings,
54
+ layer_idx=layer_idx,
55
+ )
56
+ elif config.linear_attn_type == "gla":
57
+ self.attn = SSEGLA(
58
+ mode=config.attn_mode,
59
+ hidden_size=config.hidden_size,
60
+ expand_v=config.expand_v,
61
+ head_dim=config.head_dim,
62
+ num_heads=config.num_heads,
63
+ num_v_heads=config.num_v_heads,
64
+ use_output_gate=config.use_output_gate,
65
+ use_short_conv=config.use_short_conv,
66
+ conv_size=config.conv_size,
67
+ num_sparse_partition=config.num_sparse_partition,
68
+ num_writer=config.num_writer,
69
+ num_reader=config.num_reader,
70
+ sse_implementation=config.sse_implementation,
71
+ norm_eps=config.norm_eps,
72
+ layer_idx=layer_idx,
73
+ )
74
+ elif config.linear_attn_type == "gdn":
75
+ self.attn = SSEGDN(
76
+ mode=config.attn_mode,
77
+ hidden_size=config.hidden_size,
78
+ expand_v=config.expand_v,
79
+ head_dim=config.head_dim,
80
+ num_heads=config.num_heads,
81
+ num_v_heads=config.num_v_heads,
82
+ use_output_gate=config.use_output_gate,
83
+ use_short_conv=config.use_short_conv,
84
+ allow_neg_eigval=config.allow_neg_eigval,
85
+ conv_size=config.conv_size,
86
+ num_sparse_partition=config.num_sparse_partition,
87
+ num_writer=config.num_writer,
88
+ num_reader=config.num_reader,
89
+ sse_implementation=config.sse_implementation,
90
+ norm_eps=config.norm_eps,
91
+ layer_idx=layer_idx,
92
+ )
93
+ else:
94
+ raise ValueError(f"Unknown linear attention type: {config.linear_attn_type}")
95
+ self.mlp_norm = (RMSNorm if config.fuse_norm else nn.RMSNorm)(config.hidden_size, eps=config.norm_eps)
96
+ self.mlp = SSEMLP(
97
+ hidden_size=config.hidden_size,
98
+ hidden_ratio=config.hidden_ratio,
99
+ intermediate_size=config.intermediate_size,
100
+ hidden_act=config.hidden_act,
101
+ fuse_swiglu=config.fuse_swiglu,
102
+ )
103
+
104
+ def forward(
105
+ self,
106
+ hidden_states: torch.Tensor,
107
+ attention_mask: torch.Tensor | None = None,
108
+ past_key_values: Cache | list[torch.FloatTensor] | None = None,
109
+ use_cache: bool | None = False,
110
+ output_attentions: bool | None = False,
111
+ **kwargs: Unpack[dict],
112
+ ) -> tuple[torch.FloatTensor, tuple[torch.FloatTensor, torch.FloatTensor] | None]:
113
+ residual = hidden_states
114
+ hidden_states = self.attn_norm(hidden_states)
115
+ hidden_states, attentions, past_key_values = self.attn(
116
+ hidden_states=hidden_states,
117
+ attention_mask=attention_mask,
118
+ past_key_values=past_key_values,
119
+ use_cache=use_cache,
120
+ output_attentions=output_attentions,
121
+ **kwargs,
122
+ )
123
+ if self.config.fuse_norm:
124
+ hidden_states, residual = self.mlp_norm(hidden_states, residual, True)
125
+ else:
126
+ hidden_states = residual + hidden_states
127
+ residual = hidden_states
128
+ hidden_states = self.mlp_norm(hidden_states)
129
+ hidden_states = self.mlp(hidden_states, **kwargs)
130
+ hidden_states = residual + hidden_states
131
+
132
+ aux_loss = torch.zeros(()).to(hidden_states)
133
+ # Compatible with Attention output
134
+ if isinstance(attentions, tuple):
135
+ attentions, aux_loss = attentions
136
+
137
+ outputs = (hidden_states, attentions, past_key_values, aux_loss)
138
+
139
+ return outputs
140
+
141
+
142
+ class SSEPreTrainedModel(PreTrainedModel):
143
+
144
+ config_class = SSEConfig
145
+ base_model_prefix = 'model'
146
+ supports_gradient_checkpointing = True
147
+ _no_split_modules = ['SSEBlock']
148
+ _supports_cache_class = True
149
+
150
+ def __init__(self, *inputs, **kwargs):
151
+ super().__init__(*inputs, **kwargs)
152
+
153
+ def _init_weights(
154
+ self,
155
+ module: nn.Module,
156
+ prenorm_residual_strategy: str | None = None,
157
+ num_residuals_per_layer: int = 2,
158
+ ):
159
+ if isinstance(module, SSEGDN) and next(module.parameters()).device.type != 'meta':
160
+ with torch.no_grad():
161
+ module.A_log.copy_(nn.init.uniform_(module.A_log, a=0, b=16).log())
162
+ module.A_log._no_weight_decay = True
163
+ dt = torch.exp(
164
+ nn.init.uniform_(module.dt_bias) * (math.log(0.1) - math.log(0.001)) + math.log(0.001),
165
+ ).clamp(min=1e-4)
166
+ # Inverse of softplus: https://github.com/pytorch/pytorch/issues/72759
167
+ inv_dt = dt + torch.log(-torch.expm1(-dt))
168
+ module.dt_bias.copy_(inv_dt)
169
+ module.dt_bias._no_weight_decay = True
170
+
171
+ elif isinstance(module, (nn.Linear, nn.Conv1d)):
172
+ # Slightly different from the TF version which uses truncated_normal for initialization
173
+ # cf https://github.com/pytorch/pytorch/pull/5617
174
+ nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
175
+ if module.bias is not None:
176
+ nn.init.zeros_(module.bias)
177
+ elif isinstance(module, nn.Embedding):
178
+ nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
179
+ elif hasattr(module, 'reset_parameters'):
180
+ module.reset_parameters()
181
+
182
+ if prenorm_residual_strategy is not None:
183
+ # Reinitialize selected weights subject to the OpenAI GPT-2 Paper Scheme:
184
+ # > A modified initialization which accounts for the accumulation on the residual path with model depth. Scale
185
+ # > the weights of residual layers at initialization by a factor of 1/√N where N is the # of residual layers.
186
+ # > -- GPT-2 :: https://openai.com/blog/better-language-models/
187
+ #
188
+ # Reference (Megatron-LM): https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/model/gpt_model.py
189
+ p = None
190
+ if hasattr(module, 'o_proj'):
191
+ p = module.o_proj.weight
192
+ elif hasattr(module, 'down_proj'):
193
+ p = module.down_proj.weight
194
+ if p is not None:
195
+ # Special Scaled Initialization --> There are 2 Layer Norms per Transformer Block
196
+ # Following Pytorch init, except scale by 1/sqrt(2 * n_layer)
197
+ # We need to reinit p since this code could be called multiple times
198
+ # Having just p *= scale would repeatedly scale it down
199
+ if prenorm_residual_strategy == 'rescale':
200
+ nn.init.kaiming_uniform_(p, a=math.sqrt(5))
201
+ with torch.no_grad():
202
+ p /= math.sqrt(num_residuals_per_layer * self.config.num_hidden_layers)
203
+ elif prenorm_residual_strategy == 'zero':
204
+ nn.init.zeros_(p)
205
+ else:
206
+ raise ValueError(f"Invalid prenorm_residual_strategy: {prenorm_residual_strategy}")
207
+
208
+
209
+ @dataclass
210
+ class MoeModelOutputWithPastAndAuxLosses(BaseModelOutputWithPast):
211
+ """
212
+ Base class for model's outputs, with potential hidden states and attentions.
213
+
214
+ Args:
215
+ aux_losses (`Optional[Tuple[torch.FloatTensor]]`, *optional*, returned when `labels` is provided):
216
+ aux_losses for the sparse modules.
217
+ """
218
+
219
+ aux_losses: Optional[Tuple[torch.FloatTensor]] = None
220
+
221
+
222
+ class SSEModel(SSEPreTrainedModel):
223
+
224
+ def __init__(self, config: SSEConfig):
225
+ super().__init__(config)
226
+ self.padding_idx = config.pad_token_id
227
+ self.vocab_size = config.vocab_size
228
+
229
+ self.embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
230
+ self.layers = nn.ModuleList([SSEBlock(config, layer_idx) for layer_idx in range(config.num_hidden_layers)])
231
+ self.norm = (RMSNorm if config.fuse_norm else nn.RMSNorm)(config.hidden_size, eps=config.norm_eps)
232
+
233
+ self.gradient_checkpointing = False
234
+
235
+ self.post_init()
236
+
237
+ def get_input_embeddings(self):
238
+ return self.embeddings
239
+
240
+ def set_input_embeddings(self, value):
241
+ self.embeddings = value
242
+
243
+ def forward(
244
+ self,
245
+ input_ids: torch.LongTensor | None = None,
246
+ attention_mask: Optional[torch.Tensor] = None, # noqa
247
+ inputs_embeds: torch.FloatTensor | None = None,
248
+ past_key_values: Cache | list[torch.FloatTensor] | None = None,
249
+ use_cache: bool | None = None,
250
+ output_attentions: bool | None = None,
251
+ output_hidden_states: bool | None = None,
252
+ return_dict: bool | None = None,
253
+ **kwargs: Unpack[dict],
254
+ ) -> tuple | BaseModelOutputWithPast:
255
+ if output_attentions:
256
+ warnings.warn("`SSEModel` does not `output_attentions` now, setting it to `False`.")
257
+ output_attentions = False
258
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
259
+ output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
260
+ output_aux_losses = True
261
+ use_cache = use_cache if use_cache is not None else (self.config.use_cache if not self.training else False)
262
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
263
+
264
+ # retrieve input_ids and inputs_embeds
265
+ if input_ids is not None and inputs_embeds is not None:
266
+ raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
267
+ if input_ids is None and inputs_embeds is None:
268
+ raise ValueError("You have to specify either input_ids or inputs_embeds")
269
+
270
+ if inputs_embeds is None:
271
+ inputs_embeds = self.embeddings(input_ids)
272
+ hidden_states = inputs_embeds
273
+
274
+ if use_cache and not isinstance(past_key_values, Cache):
275
+ past_key_values = Cache.from_legacy_cache(past_key_values)
276
+
277
+ all_hidden_states = () if output_hidden_states else None
278
+ all_attns = () if output_attentions else None
279
+ all_aux_losses = () if output_aux_losses else None
280
+ for layer in self.layers:
281
+ if output_hidden_states:
282
+ all_hidden_states += (hidden_states,)
283
+
284
+ hidden_states, attentions, past_key_values, aux_loss = layer(
285
+ hidden_states,
286
+ attention_mask=attention_mask,
287
+ past_key_values=past_key_values,
288
+ use_cache=use_cache,
289
+ output_attentions=output_attentions,
290
+ **kwargs,
291
+ )
292
+
293
+ if output_attentions:
294
+ all_attns += (attentions,)
295
+
296
+ if output_aux_losses:
297
+ all_aux_losses += (aux_loss,)
298
+
299
+ hidden_states = self.norm(hidden_states)
300
+
301
+ # add hidden states from the last decoder layer
302
+ if output_hidden_states:
303
+ all_hidden_states += (hidden_states,)
304
+
305
+ if not return_dict:
306
+ return tuple(i for i in [hidden_states, past_key_values, all_hidden_states, all_attns, all_aux_losses] if i is not None)
307
+ return MoeModelOutputWithPastAndAuxLosses(
308
+ last_hidden_state=hidden_states,
309
+ past_key_values=past_key_values,
310
+ hidden_states=all_hidden_states,
311
+ attentions=all_attns,
312
+ aux_losses=all_aux_losses,
313
+ )
314
+
315
+
316
+ class SSEForCausalLM(SSEPreTrainedModel, FLAGenerationMixin):
317
+
318
+ _tied_weights_keys = ["lm_head.weight"]
319
+
320
+ def __init__(self, config):
321
+ super().__init__(config)
322
+ self.model = SSEModel(config)
323
+ self.vocab_size = config.vocab_size
324
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
325
+ self.criterion = None
326
+ self.aux_loss_coef = config.aux_loss_coef
327
+
328
+ # Initialize weights and apply final processing
329
+ self.post_init()
330
+
331
+ def get_input_embeddings(self):
332
+ return self.model.embeddings
333
+
334
+ def set_input_embeddings(self, value):
335
+ self.model.embeddings = value
336
+
337
+ def get_output_embeddings(self):
338
+ return self.lm_head
339
+
340
+ def set_output_embeddings(self, new_embeddings):
341
+ self.lm_head = new_embeddings
342
+
343
+ def set_decoder(self, decoder):
344
+ self.model = decoder
345
+
346
+ def get_decoder(self):
347
+ return self.model
348
+
349
+ def generate(self, *args, **kwargs):
350
+ try:
351
+ return super().generate(*args, **kwargs)
352
+ except AttributeError as exception:
353
+ if 'past_key_values' in str(exception):
354
+ raise AttributeError(
355
+ f"You tried to call `generate` with a decoding strategy that manipulates `past_key_values`, "
356
+ f"which is not supported for {self.__class__.__name__}. "
357
+ f"Try another generation strategy instead. "
358
+ f"For the available generation strategies, check this doc: "
359
+ f"https://huggingface.co/docs/transformers/en/generation_strategies#decoding-strategies",
360
+ )
361
+ else:
362
+ raise exception
363
+
364
+ @deprecate_kwarg("num_logits_to_keep", version="4.50", new_name="logits_to_keep")
365
+ def forward(
366
+ self,
367
+ input_ids: torch.LongTensor = None,
368
+ attention_mask: torch.Tensor | None = None,
369
+ inputs_embeds: torch.Tensor | None = None,
370
+ past_key_values: Cache | list[torch.FloatTensor] | None = None,
371
+ labels: torch.LongTensor | None = None,
372
+ use_cache: bool | None = None,
373
+ output_attentions: bool | None = None,
374
+ output_hidden_states: bool | None = None,
375
+ return_dict: bool | None = None,
376
+ logits_to_keep: int | None = 0,
377
+ **kwargs: Unpack[dict],
378
+ ) -> tuple | MoeCausalLMOutputWithPast:
379
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
380
+ output_hidden_states = (
381
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
382
+ )
383
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
384
+
385
+ outputs = self.model(
386
+ input_ids=input_ids,
387
+ attention_mask=attention_mask,
388
+ inputs_embeds=inputs_embeds,
389
+ past_key_values=past_key_values,
390
+ use_cache=use_cache,
391
+ output_attentions=output_attentions,
392
+ output_hidden_states=output_hidden_states,
393
+ return_dict=return_dict,
394
+ **kwargs,
395
+ )
396
+
397
+ hidden_states = outputs[0]
398
+
399
+ loss, aux_loss, logits = None, None, None
400
+ if not self.config.fuse_linear_cross_entropy or labels is None:
401
+ logits = self.lm_head(hidden_states if logits_to_keep is None else hidden_states[:, -logits_to_keep:])
402
+ if labels is not None:
403
+ if getattr(self, 'criterion', None) is None:
404
+ if self.config.fuse_linear_cross_entropy:
405
+ criterion = FusedLinearCrossEntropyLoss(use_l2warp=self.config.use_l2warp)
406
+ elif self.config.fuse_cross_entropy:
407
+ criterion = FusedCrossEntropyLoss(inplace_backward=True)
408
+ else:
409
+ criterion = nn.CrossEntropyLoss()
410
+ else:
411
+ criterion = self.criterion
412
+ labels = labels.to(hidden_states.device)
413
+ labels = torch.cat((labels[..., 1:], torch.full_like(labels[:, :1], criterion.ignore_index)), 1)
414
+ if self.config.fuse_linear_cross_entropy:
415
+ loss = criterion(hidden_states, labels, self.lm_head.weight, self.lm_head.bias)
416
+ else:
417
+ loss = criterion(logits.view(labels.numel(), -1), labels.view(-1))
418
+ loss = l2_warp(loss, logits) if self.config.use_l2warp else loss
419
+
420
+ aux_losses = outputs.aux_losses
421
+ compute_device = aux_losses[0].device
422
+ aux_loss = sum(layer_aux_loss.to(compute_device) for layer_aux_loss in aux_losses)
423
+
424
+ loss += self.aux_loss_coef * aux_loss.to(loss.device)
425
+
426
+ if not return_dict:
427
+ output = (logits,) + outputs[1:]
428
+ return (loss,) + output if loss is not None else output
429
+
430
+ return MoeCausalLMOutputWithPast(
431
+ loss=loss,
432
+ aux_loss=aux_loss,
433
+ logits=logits,
434
+ past_key_values=outputs.past_key_values,
435
+ hidden_states=outputs.hidden_states,
436
+ attentions=outputs.attentions,
437
+ )
code/flash-linear-attention/fla/models/transformer/__init__.py ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from transformers import AutoConfig, AutoModel, AutoModelForCausalLM
3
+
4
+ from fla.models.transformer.configuration_transformer import TransformerConfig
5
+ from fla.models.transformer.modeling_transformer import TransformerForCausalLM, TransformerModel
6
+
7
+ AutoConfig.register(TransformerConfig.model_type, TransformerConfig, exist_ok=True)
8
+ AutoModel.register(TransformerConfig, TransformerModel, exist_ok=True)
9
+ AutoModelForCausalLM.register(TransformerConfig, TransformerForCausalLM, exist_ok=True)
10
+
11
+
12
+ __all__ = ['TransformerConfig', 'TransformerForCausalLM', 'TransformerModel']
code/flash-linear-attention/fla/models/transformer/configuration_transformer.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import warnings
3
+
4
+ from transformers.configuration_utils import PretrainedConfig
5
+
6
+
7
+ class TransformerConfig(PretrainedConfig):
8
+
9
+ model_type = 'transformer'
10
+ keys_to_ignore_at_inference = ['past_key_values']
11
+
12
+ def __init__(
13
+ self,
14
+ hidden_size: int = 2048,
15
+ num_hidden_layers: int = 24,
16
+ num_heads: int = 32,
17
+ num_kv_heads: int | None = None,
18
+ qkv_bias: bool = False,
19
+ qk_norm: bool = False,
20
+ window_size: int | None = None,
21
+ rope_theta: float | None = 10000.,
22
+ max_position_embeddings: int = 2048,
23
+ hidden_ratio: int | None = 4,
24
+ intermediate_size: int | None = None,
25
+ hidden_act: str = "swish",
26
+ initializer_range: float = 0.02,
27
+ elementwise_affine: bool | None = True,
28
+ norm_eps: float = 1e-6,
29
+ use_cache: bool = True,
30
+ pad_token_id: int | None = None,
31
+ bos_token_id: int = 1,
32
+ eos_token_id: int = 2,
33
+ tie_word_embeddings: bool = False,
34
+ fuse_norm: bool = True,
35
+ fuse_swiglu: bool = True,
36
+ fuse_cross_entropy: bool = True,
37
+ fuse_linear_cross_entropy: bool = False,
38
+ use_l2warp: bool = False,
39
+ vocab_size: int = 32000,
40
+ **kwargs,
41
+ ):
42
+ self.hidden_size = hidden_size
43
+ self.num_hidden_layers = num_hidden_layers
44
+ self.num_heads = num_heads
45
+ self.num_kv_heads = num_kv_heads
46
+ self.qkv_bias = qkv_bias
47
+ self.qk_norm = qk_norm
48
+ self.window_size = window_size
49
+ self.rope_theta = rope_theta
50
+ self.max_position_embeddings = max_position_embeddings
51
+
52
+ self.hidden_ratio = hidden_ratio
53
+ self.intermediate_size = intermediate_size
54
+ self.hidden_act = hidden_act
55
+
56
+ self.initializer_range = initializer_range
57
+ self.elementwise_affine = elementwise_affine
58
+ self.norm_eps = norm_eps
59
+ self.use_cache = use_cache
60
+
61
+ self.fuse_norm = fuse_norm
62
+ self.fuse_swiglu = fuse_swiglu
63
+ self.fuse_cross_entropy = fuse_cross_entropy
64
+ self.fuse_linear_cross_entropy = fuse_linear_cross_entropy
65
+ self.use_l2warp = use_l2warp
66
+ self.vocab_size = vocab_size
67
+
68
+ if fuse_cross_entropy and fuse_linear_cross_entropy:
69
+ raise ValueError(
70
+ "`fuse_cross_entropy` and `fuse_linear_cross_entropy` cannot be True at the same time.",
71
+ )
72
+ if fuse_linear_cross_entropy:
73
+ warnings.warn(
74
+ "`fuse_linear_cross_entropy` is enabled, which can improves memory efficiency "
75
+ "at the potential cost of reduced precision. "
76
+ "If you observe issues like loss divergence, consider disabling this setting.",
77
+ )
78
+
79
+ super().__init__(
80
+ pad_token_id=pad_token_id,
81
+ bos_token_id=bos_token_id,
82
+ eos_token_id=eos_token_id,
83
+ tie_word_embeddings=tie_word_embeddings,
84
+ **kwargs,
85
+ )
code/flash-linear-attention/fla/models/transformer/modeling_transformer.py ADDED
@@ -0,0 +1,356 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from __future__ import annotations
3
+
4
+ import math
5
+ import warnings
6
+ from typing import TYPE_CHECKING, Any
7
+
8
+ import torch
9
+ import torch.nn as nn
10
+ from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
11
+ from transformers.modeling_utils import PreTrainedModel
12
+ from transformers.utils import logging
13
+ from transformers.utils.deprecation import deprecate_kwarg
14
+
15
+ from fla.layers.attn import Attention
16
+ from fla.models.transformer.configuration_transformer import TransformerConfig
17
+ from fla.models.utils import Cache, FLAGenerationMixin
18
+ from fla.modules import FusedCrossEntropyLoss, FusedLinearCrossEntropyLoss, RMSNorm
19
+ from fla.modules import GatedMLP as TransformerMLP
20
+ from fla.modules.l2warp import l2_warp
21
+
22
+ if TYPE_CHECKING:
23
+ from transformers.processing_utils import Unpack
24
+
25
+
26
+ try:
27
+ from transformers.modeling_layers import GradientCheckpointingLayer
28
+ except ImportError:
29
+ from fla.models.modeling_layers import GradientCheckpointingLayer
30
+
31
+ logger = logging.get_logger(__name__)
32
+
33
+
34
+ class TransformerBlock(GradientCheckpointingLayer):
35
+
36
+ def __init__(self, config: TransformerConfig, layer_idx: int):
37
+ super().__init__()
38
+
39
+ self.config = config
40
+ self.layer_idx = layer_idx
41
+
42
+ self.attn_norm = (RMSNorm if config.fuse_norm else nn.RMSNorm)(config.hidden_size, eps=config.norm_eps)
43
+ self.attn = Attention(
44
+ hidden_size=config.hidden_size,
45
+ num_heads=config.num_heads,
46
+ num_kv_heads=config.num_kv_heads,
47
+ qkv_bias=config.qkv_bias,
48
+ qk_norm=config.qk_norm,
49
+ window_size=config.window_size,
50
+ rope_theta=config.rope_theta,
51
+ max_position_embeddings=config.max_position_embeddings,
52
+ layer_idx=layer_idx,
53
+ )
54
+
55
+ self.mlp_norm = (RMSNorm if config.fuse_norm else nn.RMSNorm)(config.hidden_size, eps=config.norm_eps)
56
+ self.mlp = TransformerMLP(
57
+ hidden_size=config.hidden_size,
58
+ hidden_ratio=config.hidden_ratio,
59
+ intermediate_size=config.intermediate_size,
60
+ hidden_act=config.hidden_act,
61
+ fuse_swiglu=config.fuse_swiglu,
62
+ )
63
+
64
+ def forward(
65
+ self,
66
+ hidden_states: torch.Tensor,
67
+ attention_mask: torch.Tensor | None = None,
68
+ past_key_values: tuple[torch.Tensor] | None = None,
69
+ output_attentions: bool | None = False,
70
+ use_cache: bool | None = False,
71
+ **kwargs: Unpack[Any],
72
+ ) -> tuple[torch.FloatTensor, tuple[torch.FloatTensor, torch.FloatTensor] | None]:
73
+
74
+ residual = hidden_states
75
+ hidden_states = self.attn_norm(hidden_states)
76
+ hidden_states, attentions, past_key_values = self.attn(
77
+ hidden_states=hidden_states,
78
+ attention_mask=attention_mask,
79
+ past_key_values=past_key_values,
80
+ use_cache=use_cache,
81
+ output_attentions=output_attentions,
82
+ **kwargs,
83
+ )
84
+ if self.config.fuse_norm:
85
+ hidden_states, residual = self.mlp_norm(hidden_states, residual, True)
86
+ else:
87
+ hidden_states = residual + hidden_states
88
+ residual = hidden_states
89
+ hidden_states = self.mlp_norm(hidden_states)
90
+ hidden_states = self.mlp(hidden_states, **kwargs)
91
+ hidden_states = residual + hidden_states
92
+
93
+ outputs = (hidden_states,)
94
+
95
+ if output_attentions:
96
+ outputs += (attentions,)
97
+
98
+ if use_cache:
99
+ outputs += (past_key_values,)
100
+
101
+ return outputs
102
+
103
+
104
+ class TransformerPreTrainedModel(PreTrainedModel):
105
+
106
+ config_class = TransformerConfig
107
+ base_model_prefix = 'model'
108
+ supports_gradient_checkpointing = True
109
+ _no_split_modules = ['TransformerBlock']
110
+ _supports_cache_class = True
111
+
112
+ def __init__(self, *inputs, **kwargs):
113
+ super().__init__(*inputs, **kwargs)
114
+
115
+ def _init_weights(
116
+ self,
117
+ module: nn.Module,
118
+ rescale_prenorm_residual: bool = False,
119
+ num_residuals_per_layer: int = 2,
120
+ ):
121
+ if isinstance(module, (nn.Linear, nn.Conv1d)):
122
+ # Slightly different from the TF version which uses truncated_normal for initialization
123
+ # cf https://github.com/pytorch/pytorch/pull/5617
124
+ nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
125
+ if module.bias is not None:
126
+ nn.init.zeros_(module.bias)
127
+ elif isinstance(module, nn.Embedding):
128
+ nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
129
+ elif hasattr(module, 'reset_parameters'):
130
+ module.reset_parameters()
131
+
132
+ if rescale_prenorm_residual:
133
+ # Reinitialize selected weights subject to the OpenAI GPT-2 Paper Scheme:
134
+ # > A modified initialization which accounts for the accumulation on the residual path with model depth. Scale
135
+ # > the weights of residual layers at initialization by a factor of 1/√N where N is the # of residual layers.
136
+ # > -- GPT-2 :: https://openai.com/blog/better-language-models/
137
+ #
138
+ # Reference (Megatron-LM): https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/model/gpt_model.py
139
+ p = None
140
+ if hasattr(module, 'o_proj'):
141
+ p = module.o_proj.weight
142
+ elif hasattr(module, 'down_proj'):
143
+ p = module.down_proj.weight
144
+ if p is not None:
145
+ # Special Scaled Initialization --> There are 2 Layer Norms per Transformer Block
146
+ # Following Pytorch init, except scale by 1/sqrt(2 * n_layer)
147
+ # We need to reinit p since this code could be called multiple times
148
+ # Having just p *= scale would repeatedly scale it down
149
+ nn.init.kaiming_uniform_(p, a=math.sqrt(5))
150
+ with torch.no_grad():
151
+ p /= math.sqrt(num_residuals_per_layer * self.config.num_hidden_layers)
152
+
153
+
154
+ class TransformerModel(TransformerPreTrainedModel):
155
+
156
+ def __init__(
157
+ self,
158
+ config: TransformerConfig,
159
+ ) -> TransformerModel:
160
+ super().__init__(config)
161
+ self.padding_idx = config.pad_token_id
162
+ self.vocab_size = config.vocab_size
163
+
164
+ self.embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
165
+ self.layers = nn.ModuleList([TransformerBlock(config, layer_idx) for layer_idx in range(config.num_hidden_layers)])
166
+ self.norm = (RMSNorm if config.fuse_norm else nn.RMSNorm)(config.hidden_size, eps=config.norm_eps)
167
+
168
+ self.gradient_checkpointing = False
169
+
170
+ self.post_init()
171
+
172
+ def get_input_embeddings(self):
173
+ return self.embeddings
174
+
175
+ def set_input_embeddings(self, value):
176
+ self.embeddings = value
177
+
178
+ def forward(
179
+ self,
180
+ input_ids: torch.LongTensor | None = None,
181
+ attention_mask: torch.Tensor | None = None,
182
+ past_key_values: list[torch.FloatTensor] | None = None,
183
+ inputs_embeds: torch.FloatTensor | None = None,
184
+ use_cache: bool | None = None,
185
+ output_attentions: bool | None = None,
186
+ output_hidden_states: bool | None = None,
187
+ return_dict: bool | None = None,
188
+ **kwargs: Unpack[Any],
189
+ ) -> tuple | CausalLMOutputWithPast:
190
+ if output_attentions:
191
+ warnings.warn(
192
+ "`TransformerModel` does not support output attention weights now, so `output_attentions` is set to `False`.",
193
+ )
194
+ output_attentions = False
195
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
196
+ output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
197
+ use_cache = use_cache if use_cache is not None else (self.config.use_cache if not self.training else False)
198
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
199
+
200
+ # retrieve input_ids and inputs_embeds
201
+ if input_ids is not None and inputs_embeds is not None:
202
+ raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time")
203
+ elif input_ids is None and inputs_embeds is None:
204
+ raise ValueError("You have to specify either input_ids or inputs_embeds")
205
+
206
+ if use_cache and not isinstance(past_key_values, Cache):
207
+ past_key_values = Cache.from_legacy_cache(past_key_values)
208
+
209
+ if inputs_embeds is None:
210
+ inputs_embeds = self.embeddings(input_ids)
211
+
212
+ # embed positions
213
+ hidden_states = inputs_embeds
214
+
215
+ all_hidden_states = () if output_hidden_states else None
216
+ all_attns = () if output_attentions else None
217
+ next_cache = None
218
+
219
+ for layer in self.layers:
220
+ if output_hidden_states:
221
+ all_hidden_states += (hidden_states,)
222
+
223
+ layer_outputs = layer(
224
+ hidden_states,
225
+ attention_mask=attention_mask,
226
+ past_key_values=past_key_values,
227
+ output_attentions=output_attentions,
228
+ use_cache=use_cache,
229
+ **kwargs,
230
+ )
231
+
232
+ hidden_states = layer_outputs[0]
233
+
234
+ if use_cache:
235
+ next_cache = layer_outputs[2 if output_attentions else 1]
236
+
237
+ if output_attentions:
238
+ all_attns += (layer_outputs[1],)
239
+
240
+ hidden_states = self.norm(hidden_states)
241
+
242
+ # add hidden states from the last decoder layer
243
+ if output_hidden_states:
244
+ all_hidden_states += (hidden_states,)
245
+
246
+ if not return_dict:
247
+ return tuple(v for v in [hidden_states, next_cache, all_hidden_states, all_attns] if v is not None)
248
+
249
+ return BaseModelOutputWithPast(
250
+ last_hidden_state=hidden_states,
251
+ past_key_values=next_cache,
252
+ hidden_states=all_hidden_states,
253
+ attentions=all_attns,
254
+ )
255
+
256
+
257
+ class TransformerForCausalLM(TransformerPreTrainedModel, FLAGenerationMixin):
258
+
259
+ _tied_weights_keys = ["lm_head.weight"]
260
+
261
+ def __init__(self, config):
262
+ super().__init__(config)
263
+ self.model = TransformerModel(config)
264
+ self.vocab_size = config.vocab_size
265
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
266
+ self.criterion = None
267
+
268
+ # Initialize weights and apply final processing
269
+ self.post_init()
270
+
271
+ def get_input_embeddings(self):
272
+ return self.model.embeddings
273
+
274
+ def set_input_embeddings(self, value):
275
+ self.model.embeddings = value
276
+
277
+ def get_output_embeddings(self):
278
+ return self.lm_head
279
+
280
+ def set_output_embeddings(self, new_embeddings):
281
+ self.lm_head = new_embeddings
282
+
283
+ def set_decoder(self, decoder):
284
+ self.model = decoder
285
+
286
+ def get_decoder(self):
287
+ return self.model
288
+
289
+ @deprecate_kwarg("num_logits_to_keep", version="4.50", new_name="logits_to_keep")
290
+ def forward(
291
+ self,
292
+ input_ids: torch.LongTensor = None,
293
+ attention_mask: torch.Tensor | None = None,
294
+ past_key_values: Cache | list[torch.FloatTensor] | None = None,
295
+ inputs_embeds: torch.FloatTensor | None = None,
296
+ labels: torch.LongTensor | None = None,
297
+ use_cache: bool | None = None,
298
+ output_attentions: bool | None = None,
299
+ output_hidden_states: bool | None = None,
300
+ return_dict: bool | None = None,
301
+ logits_to_keep: int | None = 0,
302
+ **kwargs: Unpack[Any],
303
+ ) -> tuple | CausalLMOutputWithPast:
304
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
305
+ output_hidden_states = (
306
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
307
+ )
308
+ return_dict = return_dict if return_dict is not None else self.config.use_return_dict
309
+
310
+ outputs = self.model(
311
+ input_ids=input_ids,
312
+ attention_mask=attention_mask,
313
+ past_key_values=past_key_values,
314
+ inputs_embeds=inputs_embeds,
315
+ use_cache=use_cache,
316
+ output_attentions=output_attentions,
317
+ output_hidden_states=output_hidden_states,
318
+ return_dict=return_dict,
319
+ **kwargs,
320
+ )
321
+
322
+ hidden_states = outputs[0]
323
+
324
+ logits = None if self.config.fuse_linear_cross_entropy else self.lm_head(hidden_states[:, -logits_to_keep:])
325
+
326
+ loss = None
327
+ if labels is not None:
328
+ if getattr(self, 'criterion', None) is None:
329
+ if self.config.fuse_linear_cross_entropy:
330
+ criterion = FusedLinearCrossEntropyLoss(use_l2warp=self.config.use_l2warp)
331
+ elif self.config.fuse_cross_entropy:
332
+ criterion = FusedCrossEntropyLoss(inplace_backward=True)
333
+ else:
334
+ criterion = nn.CrossEntropyLoss()
335
+ else:
336
+ criterion = self.criterion
337
+ # Enable model parallelism
338
+ labels = labels.to(hidden_states.device)
339
+ labels = torch.cat((labels[..., 1:], torch.full_like(labels[:, :1], criterion.ignore_index)), 1)
340
+ if self.config.fuse_linear_cross_entropy:
341
+ loss = criterion(hidden_states, labels, self.lm_head.weight, self.lm_head.bias)
342
+ else:
343
+ loss = criterion(logits.view(labels.numel(), -1), labels.view(-1))
344
+ loss = l2_warp(loss, logits) if self.config.use_l2warp else loss
345
+
346
+ if not return_dict:
347
+ output = (logits,) + outputs[1:]
348
+ return (loss,) + output if loss is not None else output
349
+
350
+ return CausalLMOutputWithPast(
351
+ loss=loss,
352
+ logits=logits,
353
+ past_key_values=outputs.past_key_values,
354
+ hidden_states=outputs.hidden_states,
355
+ attentions=outputs.attentions,
356
+ )
code/flash-linear-attention/fla/models/utils.py ADDED
@@ -0,0 +1,471 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from __future__ import annotations
3
+
4
+ import inspect
5
+ from typing import Any
6
+
7
+ import torch
8
+ import transformers
9
+ from packaging import version
10
+ from transformers.cache_utils import Cache as HFCacheBase
11
+ from transformers.generation import GenerationMixin
12
+ from transformers.utils.deprecation import deprecate_kwarg
13
+
14
+ _TF_VERSION = transformers.__version__
15
+ _NEED_NEW = "4.53.3"
16
+ _IS_TRANSFORMERS_4_56_PLUS = version.parse(_TF_VERSION) >= version.parse("4.56.0")
17
+
18
+ if version.parse(_TF_VERSION) > version.parse(_NEED_NEW):
19
+ from transformers.cache_utils import CacheLayerMixin
20
+ else:
21
+ CacheLayerMixin = object
22
+
23
+
24
+ class FLALayer(CacheLayerMixin):
25
+ is_compileable = True
26
+ is_sliding = False
27
+
28
+ def __init__(self):
29
+ super().__init__()
30
+ self.state = None
31
+
32
+ def lazy_initialization(self, key_states: torch.Tensor):
33
+ self.state = None
34
+
35
+ def update(
36
+ self,
37
+ *,
38
+ recurrent_state: torch.Tensor | tuple[torch.Tensor, ...] | None = None,
39
+ attn_state: tuple[torch.Tensor, ...] | None = None,
40
+ conv_state: Any | None = None,
41
+ ffn_state: Any | None = None,
42
+ cache_kwargs: dict[str, Any] | None = None,
43
+ **_: Any,
44
+ ) -> dict[str, Any]:
45
+ if cache_kwargs is None:
46
+ cache_kwargs = {}
47
+ window_size = cache_kwargs.get("window_size")
48
+
49
+ if attn_state is not None and not isinstance(attn_state, (tuple, list)):
50
+ raise ValueError("`attn_state` must be a tuple/list of tensors")
51
+
52
+ if self.state is None:
53
+ self.state = {
54
+ "recurrent_state": None,
55
+ "attn_state": None,
56
+ "conv_state": None,
57
+ "ffn_state": None,
58
+ }
59
+
60
+ if recurrent_state is not None:
61
+ self.state["recurrent_state"] = recurrent_state
62
+
63
+ if attn_state is not None:
64
+ input_size = attn_state[0].shape[1]
65
+ if self.state["attn_state"] is None:
66
+ if window_size is not None and input_size > window_size:
67
+ attn_state = tuple(x[:, -window_size:].contiguous() for x in attn_state)
68
+ self.state["attn_state"] = tuple(attn_state)
69
+ else:
70
+ old = self.state["attn_state"]
71
+ if window_size is not None and old[0].shape[1] >= window_size:
72
+ new_tuple = []
73
+ for old_x, new_x in zip(old, attn_state, strict=False):
74
+ rolled = old_x.roll(-input_size, dims=1)
75
+ tail = new_x[:, -window_size:]
76
+ rolled[:, -tail.shape[1]:] = tail
77
+ new_tuple.append(rolled)
78
+ self.state["attn_state"] = tuple(new_tuple)
79
+ else:
80
+ self.state["attn_state"] = tuple(
81
+ torch.cat([old_x, new_x], dim=1) for old_x, new_x in zip(old, attn_state, strict=False)
82
+ )
83
+
84
+ if conv_state is not None:
85
+ self.state["conv_state"] = conv_state
86
+ if ffn_state is not None:
87
+ self.state["ffn_state"] = ffn_state
88
+
89
+ if not hasattr(self, 'device'):
90
+ self.device = 'cpu'
91
+ for state in (recurrent_state, attn_state, conv_state, ffn_state):
92
+ if state is not None:
93
+ self.device = state.device if isinstance(state, torch.Tensor) else state[0].device
94
+ break
95
+
96
+ return self.state
97
+
98
+ def get_seq_length(self, cache_position=None) -> int:
99
+ # we do not store seen_tokens here
100
+ return 0
101
+
102
+ def get_max_cache_shape(self) -> int:
103
+ return -1
104
+
105
+ def get_mask_sizes(self, cache_position: torch.Tensor) -> tuple[int, int]:
106
+ return 0, 0
107
+
108
+ def offload(self):
109
+ if self.state is None:
110
+ return
111
+
112
+ def to_cpu(x):
113
+ return x.to("cpu", non_blocking=True) if isinstance(x, torch.Tensor) else x
114
+ for k in ("recurrent_state", "attn_state", "conv_state", "ffn_state"):
115
+ v = self.state.get(k, None)
116
+ if v is None:
117
+ continue
118
+ if isinstance(v, (tuple, list)):
119
+ self.state[k] = tuple(to_cpu(t) for t in v)
120
+ else:
121
+ self.state[k] = to_cpu(v)
122
+
123
+ def prefetch(self):
124
+ if self.state is None:
125
+ return
126
+
127
+ def to_dev(x):
128
+ return x.to(self.device, non_blocking=True) if isinstance(x, torch.Tensor) else x
129
+ for k in ("recurrent_state", "attn_state", "conv_state", "ffn_state"):
130
+ v = self.state.get(k, None)
131
+ if v is None:
132
+ continue
133
+ if isinstance(v, (tuple, list)):
134
+ self.state[k] = tuple(to_dev(t) for t in v)
135
+ else:
136
+ self.state[k] = to_dev(v)
137
+
138
+ def reset(self):
139
+ pass
140
+
141
+
142
+ class LegacyFLACache(HFCacheBase):
143
+ """
144
+ A cache used for storing hidden states produced by flash linear attention models.
145
+
146
+ It stores the states of each layer as the tensor of shape `[batch_size, key_dim, value_dim]`.
147
+ """
148
+
149
+ is_compileable = True
150
+
151
+ def __init__(
152
+ self,
153
+ seen_tokens: int = 0,
154
+ ) -> LegacyFLACache:
155
+ super().__init__()
156
+
157
+ self.states: list[dict[str, Any]] = []
158
+
159
+ self._seen_tokens = seen_tokens # Used in `generate` to keep tally of how many tokens the cache has seen
160
+
161
+ def __getitem__(self, layer_idx: int) -> dict[str, Any]:
162
+ if layer_idx < len(self):
163
+ return self.states[layer_idx]
164
+ else:
165
+ raise KeyError(f"Cache only has {len(self)} layers, attempted to access layer with index {layer_idx}")
166
+
167
+ def __iter__(self):
168
+ yield from self.states
169
+
170
+ def __len__(self):
171
+ return len(self.states)
172
+
173
+ def update(
174
+ self,
175
+ recurrent_state: tuple[torch.Tensor] | None = None,
176
+ attn_state: tuple[torch.Tensor] | None = None,
177
+ conv_state: tuple[torch.Tensor] | None = None,
178
+ ffn_state: tuple[torch.Tensor] | None = None,
179
+ layer_idx: int = 0,
180
+ offset: int | None = 1,
181
+ cache_kwargs: dict[str, Any] | None = None,
182
+ ) -> dict[str, Any]:
183
+ """
184
+ Args:
185
+ recurrent_state (`torch.Tensor`):
186
+ The new recurrent state to cache.
187
+ attn_state (`tuple[torch.Tensor]`):
188
+ The new attention key/value states to cache.
189
+ conv_state (`tuple[torch.Tensor]`):
190
+ The new convolution state to cache.
191
+ ffn_state (`tuple[torch.Tensor]`):
192
+ The new feed-forward state to cache.
193
+ layer_idx (`int`, defaults to 0):
194
+ The index of the layer to cache the states for.
195
+ offset (`int`, defaults to 1):
196
+ The number of new tokens being processed.
197
+ cache_kwargs (`Dict[str, Any]`):
198
+ Additional arguments for the cache subclass.
199
+
200
+ Return:
201
+ Dictionary of the updated state.
202
+ """
203
+
204
+ if cache_kwargs is None:
205
+ cache_kwargs = {}
206
+ if attn_state is not None:
207
+ input_size = attn_state[0].shape[1]
208
+ window_size = cache_kwargs.get('window_size')
209
+ if not isinstance(attn_state, (tuple, list)):
210
+ raise ValueError("`attn_state` must be a tuple of tensors for key/value states")
211
+ if len(self.states) <= layer_idx:
212
+ # update the number of seen tokens
213
+ if layer_idx == 0:
214
+ self._seen_tokens += offset
215
+ if attn_state is not None:
216
+ if window_size is not None and input_size > window_size:
217
+ attn_state = [state[:, -window_size:].contiguous() for state in attn_state]
218
+ state = dict(
219
+ recurrent_state=recurrent_state,
220
+ attn_state=attn_state,
221
+ conv_state=conv_state,
222
+ ffn_state=ffn_state,
223
+ )
224
+ self.states.append(state)
225
+ else:
226
+ # update the number of seen tokens
227
+ if layer_idx == len(self.states) - 1:
228
+ self._seen_tokens += offset
229
+ state = self.states[layer_idx]
230
+ if recurrent_state is not None:
231
+ state['recurrent_state'] = recurrent_state
232
+ if attn_state is not None:
233
+ if window_size is not None and state['attn_state'][0].shape[1] == window_size:
234
+ for i, (old_state, new_state) in enumerate(zip(state['attn_state'], attn_state, strict=False)):
235
+ # DO NOT allocate new memory if the cache is full
236
+ # roll the key/value states to the left by `input_size`
237
+ old_state = old_state.roll(-input_size, 1)
238
+ # replace the last `input_size` tokens with the new key/value states
239
+ old_state[:, -input_size:] = new_state
240
+ state['attn_state'][i] = old_state
241
+ else:
242
+ attn_state = [
243
+ torch.cat([old_state, new_state], 1)
244
+ for old_state, new_state in zip(state['attn_state'], attn_state, strict=False)
245
+ ]
246
+ state['attn_state'] = attn_state
247
+ if conv_state is not None:
248
+ state['conv_state'] = conv_state
249
+ if ffn_state is not None:
250
+ state['ffn_state'] = ffn_state
251
+
252
+ return state
253
+
254
+ def get_seq_length(self, layer_idx: int | None = 0) -> int:
255
+ """Returns the sequence length of the cached states. A layer index can be optionally passed."""
256
+ if len(self.states) <= layer_idx:
257
+ return 0
258
+ return self._seen_tokens
259
+
260
+ def get_max_cache_shape(self) -> int | None:
261
+ """Returns the maximum sequence length of the cached states. Cache does not have a maximum length."""
262
+ return None
263
+
264
+ def to_legacy_cache(self) -> tuple:
265
+ return tuple(self.states)
266
+
267
+ @classmethod
268
+ @torch.compiler.disable
269
+ def from_legacy_cache(
270
+ cls,
271
+ past_key_values: tuple | None = None,
272
+ seen_tokens: int = 0,
273
+ ) -> LegacyFLACache:
274
+ """Converts a cache in the legacy cache format into an equivalent `Cache`."""
275
+
276
+ cache = cls(seen_tokens)
277
+ if isinstance(past_key_values, list):
278
+ for layer_idx in range(len(past_key_values)):
279
+ cache.states.append(past_key_values[layer_idx])
280
+ return cache
281
+
282
+
283
+ class FLACache(HFCacheBase):
284
+ """
285
+ A cache used for storing hidden states produced by flash linear attention models.
286
+
287
+ It stores the states of each layer as the tensor of shape `[batch_size, key_dim, value_dim]`.
288
+ """
289
+
290
+ is_compileable = True
291
+
292
+ def __init__(self, seen_tokens: int = 0, **kwargs):
293
+ parent_init = super().__init__
294
+ sig = inspect.signature(parent_init)
295
+ param_names = list(sig.parameters.keys())
296
+
297
+ if 'layer_class_to_replicate' in param_names:
298
+ self.use_layer_class_to_replicate = True
299
+ super().__init__(layer_class_to_replicate=FLALayer, **kwargs)
300
+ elif 'layer_classes' in param_names:
301
+ self.use_layer_class_to_replicate = False
302
+ super().__init__(layer_classes=FLALayer, **kwargs)
303
+ else:
304
+ raise TypeError(
305
+ "FLA cache initialization failed: HFCacheBase.__init__ accepts neither "
306
+ "'layer_class_to_replicate' nor 'layer_classes'. This might be caused by an incompatible "
307
+ "transformers version. Please check your transformers>=4.36.0",
308
+ )
309
+ self._seen_tokens = int(seen_tokens)
310
+
311
+ def update(
312
+ self,
313
+ recurrent_state: tuple[torch.Tensor] | None = None,
314
+ attn_state: tuple[torch.Tensor] | None = None,
315
+ conv_state: tuple[torch.Tensor] | None = None,
316
+ ffn_state: tuple[torch.Tensor] | None = None,
317
+ layer_idx: int = 0,
318
+ offset: int | None = 1,
319
+ cache_kwargs: dict[str, Any] | None = None,
320
+ ) -> dict[str, Any]:
321
+ if not self.use_layer_class_to_replicate:
322
+ self.append_new_layers(layer_idx)
323
+ else:
324
+ while len(self.layers) <= layer_idx:
325
+ self.layers.append(self.layer_class_to_replicate())
326
+ if layer_idx == 0:
327
+ self._seen_tokens += int(offset)
328
+
329
+ return self.layers[layer_idx].update(
330
+ recurrent_state=recurrent_state,
331
+ attn_state=attn_state,
332
+ conv_state=conv_state,
333
+ ffn_state=ffn_state,
334
+ cache_kwargs=cache_kwargs,
335
+ )
336
+
337
+ def __getitem__(self, layer_idx: int) -> dict[str, Any]:
338
+ if layer_idx >= len(self.layers):
339
+ raise KeyError(f"Cache only have {len(self.layers)} layers, however accessed {layer_idx} out of bounds")
340
+ return self.layers[layer_idx].state
341
+
342
+ def __iter__(self):
343
+ for i in range(len(self.layers)):
344
+ yield self[i]
345
+
346
+ def __len__(self):
347
+ return super().__len__()
348
+
349
+ def get_seq_length(self, layer_idx: int | None = 0, cache_position=None) -> int:
350
+ if len(self.layers) <= (layer_idx or 0):
351
+ return 0
352
+ return self._seen_tokens
353
+
354
+ def get_max_cache_shape(self, layer_idx: int = 0) -> int:
355
+ return -1
356
+
357
+ def get_mask_sizes(self, cache_position: torch.Tensor, layer_idx: int) -> tuple[int, int]:
358
+ # Respect your global seen_tokens semantics
359
+ # kv_length = past_seen + current_query_length
360
+ query_len = int(cache_position.shape[0]) if cache_position is not None else 0
361
+ kv_length = int(self._seen_tokens) + query_len
362
+ return kv_length, 0
363
+
364
+ def to_legacy_cache(self) -> tuple[dict[str, Any], ...]:
365
+ return tuple(self[i] for i in range(len(self.layers)))
366
+
367
+ @classmethod
368
+ @torch.compiler.disable
369
+ def from_legacy_cache(
370
+ cls,
371
+ past_key_values: tuple[dict[str, Any], ...] | None = None,
372
+ seen_tokens: int = 0,
373
+ **kwargs,
374
+ ) -> FLACache:
375
+ cache = cls(seen_tokens=seen_tokens, **kwargs)
376
+ if isinstance(past_key_values, (list, tuple)):
377
+ for i, st in enumerate(past_key_values):
378
+ while len(cache.layers) <= i:
379
+ cache.layers.append(cache.layer_class_to_replicate())
380
+ cache.layers[i].state = dict(st)
381
+ return cache
382
+
383
+
384
+ class FLAGenerationMixin(GenerationMixin):
385
+ """
386
+ Flash Linear Attention Generation Mixin that provides version-compatible generation methods.
387
+ This mixin handles transformers library version differences, particularly for prepare_inputs_for_generation.
388
+ """
389
+
390
+ def __init__(self, *args, **kwargs):
391
+ super().__init__(*args, **kwargs)
392
+
393
+ @deprecate_kwarg("num_logits_to_keep", version="4.50", new_name="logits_to_keep")
394
+ def prepare_inputs_for_generation(
395
+ self,
396
+ input_ids: torch.LongTensor = None,
397
+ past_key_values: HFCacheBase | None = None,
398
+ attention_mask: torch.Tensor | None = None,
399
+ inputs_embeds: torch.Tensor | None = None,
400
+ use_cache: bool = True,
401
+ logits_to_keep: int | None = None,
402
+ cache_position: torch.LongTensor | None = None,
403
+ **kwargs,
404
+ ):
405
+ # Use pre-computed version comparison for performance
406
+ if _IS_TRANSFORMERS_4_56_PLUS:
407
+ # For transformers 4.56.0+, use cache_position-based logic
408
+ model_inputs = {}
409
+
410
+ # Handle cache-dependent input preparation
411
+ if past_key_values is not None:
412
+ model_inputs["past_key_values"] = past_key_values
413
+
414
+ # Use the new cache-dependent input preparation method if available
415
+ if hasattr(self, '_cache_dependant_input_preparation') and cache_position is not None:
416
+ inputs_embeds, input_ids = self._cache_dependant_input_preparation(
417
+ input_ids, inputs_embeds, cache_position,
418
+ )
419
+ elif cache_position is not None:
420
+ # Fallback: manually slice using cache_position
421
+ if input_ids is not None and input_ids.shape[1] != cache_position.shape[0]:
422
+ input_ids = input_ids[:, cache_position]
423
+ elif hasattr(past_key_values, '__len__') and len(past_key_values) > 0:
424
+ # Ultimate fallback to old behavior
425
+ input_ids = input_ids[:, -1:]
426
+
427
+ # Handle input format (similar to base class logic)
428
+ if inputs_embeds is not None and (cache_position is None or len(cache_position) == inputs_embeds.shape[1]):
429
+ model_inputs['inputs_embeds'] = inputs_embeds
430
+ model_inputs['input_ids'] = None
431
+ else:
432
+ model_inputs['input_ids'] = input_ids.contiguous() if input_ids is not None else None
433
+ model_inputs['inputs_embeds'] = None
434
+
435
+ model_inputs['cache_position'] = cache_position
436
+
437
+ else:
438
+ # For older transformers versions, use the original logic
439
+ model_inputs = {}
440
+ # only last token for `inputs_ids` if the `past_key_values` is not empty.
441
+ if past_key_values is not None and hasattr(past_key_values, '__len__') and len(past_key_values) > 0:
442
+ input_ids = input_ids[:, -1:]
443
+ # if `inputs_embeds` are passed, we only want to use them in the 1st generation step
444
+ if inputs_embeds is not None and hasattr(past_key_values, '__len__') and len(past_key_values) == 0:
445
+ model_inputs = {'inputs_embeds': inputs_embeds}
446
+ else:
447
+ # The `contiguous()` here is necessary to have a static stride during decoding. torchdynamo otherwise
448
+ # recompiles graphs as the stride of the inputs is a guard.
449
+ # Ref: https://github.com/huggingface/transformers/pull/29114
450
+ # TODO: use `next_tokens` directly instead.
451
+ model_inputs = {'input_ids': input_ids.contiguous()}
452
+
453
+ if logits_to_keep is not None:
454
+ model_inputs['logits_to_keep'] = logits_to_keep
455
+
456
+ model_inputs.update({
457
+ 'past_key_values': past_key_values,
458
+ 'use_cache': use_cache,
459
+ 'attention_mask': attention_mask,
460
+ })
461
+ return model_inputs
462
+
463
+
464
+ if version.parse(_TF_VERSION) > version.parse(_NEED_NEW):
465
+ class Cache(FLACache):
466
+ def __init__(self, seen_tokens: int = 0, **kwargs: Any) -> None:
467
+ super().__init__(seen_tokens=seen_tokens, **kwargs)
468
+ else:
469
+ class Cache(LegacyFLACache):
470
+ def __init__(self, seen_tokens: int = 0, **kwargs: Any) -> None:
471
+ super().__init__(seen_tokens=seen_tokens)
code/flash-linear-attention/fla/modules/__init__.py ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from fla.modules.convolution import ImplicitLongConvolution, LongConvolution, ShortConvolution
3
+ from fla.modules.fused_bitlinear import BitLinear, FusedBitLinear
4
+ from fla.modules.fused_cross_entropy import FusedCrossEntropyLoss
5
+ from fla.modules.fused_kl_div import FusedKLDivLoss
6
+ from fla.modules.fused_linear_cross_entropy import FusedLinearCrossEntropyLoss
7
+ from fla.modules.fused_norm_gate import (
8
+ FusedLayerNormGated,
9
+ FusedLayerNormSwishGate,
10
+ FusedLayerNormSwishGateLinear,
11
+ FusedRMSNormGated,
12
+ FusedRMSNormSwishGate,
13
+ FusedRMSNormSwishGateLinear,
14
+ )
15
+ from fla.modules.l2norm import L2Norm
16
+ from fla.modules.layernorm import GroupNorm, GroupNormLinear, LayerNorm, LayerNormLinear, RMSNorm, RMSNormLinear
17
+ from fla.modules.mlp import GatedMLP
18
+ from fla.modules.rotary import RotaryEmbedding
19
+ from fla.modules.token_shift import TokenShift
20
+
21
+ __all__ = [
22
+ 'ImplicitLongConvolution', 'LongConvolution', 'ShortConvolution',
23
+ 'BitLinear', 'FusedBitLinear',
24
+ 'FusedCrossEntropyLoss', 'FusedLinearCrossEntropyLoss', 'FusedKLDivLoss',
25
+ 'L2Norm',
26
+ 'GroupNorm', 'GroupNormLinear', 'LayerNorm', 'LayerNormLinear', 'RMSNorm', 'RMSNormLinear',
27
+ 'FusedLayerNormGated', 'FusedLayerNormSwishGate', 'FusedLayerNormSwishGateLinear',
28
+ 'FusedRMSNormGated', 'FusedRMSNormSwishGate', 'FusedRMSNormSwishGateLinear',
29
+ 'GatedMLP',
30
+ 'RotaryEmbedding',
31
+ 'TokenShift',
32
+ ]
code/flash-linear-attention/fla/modules/activations.py ADDED
@@ -0,0 +1,555 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Tri Dao, Yu Zhang, Songlin Yang.
2
+
3
+ import torch
4
+ import torch.nn.functional as F
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils.op import exp, log
9
+ from fla.utils import autocast_custom_bwd, autocast_custom_fwd, autotune_cache_kwargs, input_guard, is_amd
10
+
11
+ try:
12
+ from torch.distributed.tensor import DTensor
13
+ except (ImportError, AttributeError):
14
+ DTensor = None
15
+
16
+ NUM_WARPS_AUTOTUNE = [1, 2, 4, 8, 16] if is_amd else [1, 2, 4, 8, 16, 32]
17
+
18
+
19
+ @triton.autotune(
20
+ configs=[
21
+ triton.Config({'B': bs}, num_warps=num_warps)
22
+ for bs in [512, 1024, 2048, 4096, 8192]
23
+ for num_warps in NUM_WARPS_AUTOTUNE
24
+ ],
25
+ key=['D'],
26
+ **autotune_cache_kwargs,
27
+ )
28
+ @triton.jit(do_not_specialize=['T'])
29
+ def sigmoid_fwd_kernel(
30
+ x, y,
31
+ T,
32
+ B: tl.constexpr,
33
+ D: tl.constexpr,
34
+ ):
35
+ pid = tl.program_id(0)
36
+ offs = pid * B + tl.arange(0, B)
37
+ mask = offs < T
38
+ x_val = tl.load(x + offs, mask=mask, other=0.).to(tl.float32)
39
+ y_val = 1.0 / (1.0 + exp(-x_val))
40
+ tl.store(y + offs, y_val.to(y.dtype.element_ty), mask=mask)
41
+
42
+
43
+ @triton.autotune(
44
+ configs=[
45
+ triton.Config({'B': bs}, num_warps=num_warps)
46
+ for bs in [512, 1024, 2048, 4096, 8192]
47
+ for num_warps in NUM_WARPS_AUTOTUNE
48
+ ],
49
+ key=['D'],
50
+ **autotune_cache_kwargs,
51
+ )
52
+ @triton.jit(do_not_specialize=['T'])
53
+ def sigmoid_bwd_kernel(
54
+ x, dy, dx,
55
+ T,
56
+ B: tl.constexpr,
57
+ D: tl.constexpr,
58
+ ):
59
+ pid = tl.program_id(0)
60
+ offs = pid * B + tl.arange(0, B)
61
+ mask = offs < T
62
+ x_val = tl.load(x + offs, mask=mask, other=0.).to(tl.float32)
63
+ g_val = tl.load(dy + offs, mask=mask, other=0.).to(tl.float32)
64
+ s = 1.0 / (1.0 + exp(-x_val))
65
+ dx_val = g_val * s * (1.0 - s)
66
+ tl.store(dx + offs, dx_val.to(dx.dtype.element_ty), mask=mask)
67
+
68
+
69
+ def sigmoid_fwd(x: torch.Tensor) -> torch.Tensor:
70
+ T, D = x.numel(), x.shape[-1]
71
+ y = torch.empty_like(x)
72
+ sigmoid_fwd_kernel[lambda meta: (triton.cdiv(T, meta['B']),)](x, y, T=T, D=D)
73
+ return y
74
+
75
+
76
+ def sigmoid_bwd(x: torch.Tensor, dy: torch.Tensor) -> torch.Tensor:
77
+ T, D = x.numel(), x.shape[-1]
78
+ dx = torch.empty_like(x)
79
+ sigmoid_bwd_kernel[lambda meta: (triton.cdiv(T, meta['B']),)](x, dy, dx, T=T, D=D)
80
+ return dx
81
+
82
+
83
+ class SigmoidFunction(torch.autograd.Function):
84
+
85
+ @staticmethod
86
+ def forward(ctx, x):
87
+ ctx.save_for_backward(x)
88
+ return sigmoid_fwd(x)
89
+
90
+ @staticmethod
91
+ def backward(ctx, dout):
92
+ x, = ctx.saved_tensors
93
+ return sigmoid_bwd(x, dout)
94
+
95
+
96
+ sigmoid = SigmoidFunction.apply
97
+
98
+
99
+ @triton.autotune(
100
+ configs=[
101
+ triton.Config({'B': bs}, num_warps=num_warps)
102
+ for bs in [512, 1024, 2048, 4096, 8192]
103
+ for num_warps in NUM_WARPS_AUTOTUNE
104
+ ],
105
+ key=['D'],
106
+ **autotune_cache_kwargs,
107
+ )
108
+ @triton.jit(do_not_specialize=['T'])
109
+ def logsigmoid_fwd_kernel(
110
+ x,
111
+ y,
112
+ temperature,
113
+ T,
114
+ B: tl.constexpr,
115
+ D: tl.constexpr,
116
+ ):
117
+ i = tl.program_id(0)
118
+ o_i = i * B + tl.arange(0, B)
119
+ m_i = o_i < T
120
+
121
+ b_x = tl.load(x + o_i, mask=m_i, other=0.).to(tl.float32)
122
+ b_m = tl.minimum(0., b_x)
123
+ b_z = 1. + exp(-tl.abs(b_x))
124
+ b_y = (b_m - log(b_z)) / temperature
125
+ tl.store(y + o_i, b_y.to(y.dtype.element_ty), mask=m_i)
126
+
127
+
128
+ @triton.autotune(
129
+ configs=[
130
+ triton.Config({'B': bs}, num_warps=num_warps)
131
+ for bs in [512, 1024, 2048, 4096, 8192]
132
+ for num_warps in NUM_WARPS_AUTOTUNE
133
+ ],
134
+ key=['D'],
135
+ **autotune_cache_kwargs,
136
+ )
137
+ @triton.jit(do_not_specialize=['T'])
138
+ def logsigmoid_bwd_kernel(
139
+ x,
140
+ dx,
141
+ dy,
142
+ temperature,
143
+ T,
144
+ B: tl.constexpr,
145
+ D: tl.constexpr,
146
+ ):
147
+ i = tl.program_id(0)
148
+ o_i = i * B + tl.arange(0, B)
149
+ m_i = o_i < T
150
+
151
+ b_x = tl.load(x + o_i, mask=m_i, other=0.).to(tl.float32)
152
+ b_dy = tl.load(dy + o_i, mask=m_i, other=0.).to(tl.float32)
153
+ b_dx = b_dy * ((1. - tl.sigmoid(b_x)) / temperature)
154
+ tl.store(dx + o_i, b_dx.to(dx.dtype.element_ty), mask=m_i)
155
+
156
+
157
+ def logsigmoid_fwd(x: torch.Tensor, temperature: float = 1.) -> torch.Tensor:
158
+ T, D = x.numel(), x.shape[-1]
159
+ y = torch.empty_like(x)
160
+ logsigmoid_fwd_kernel[lambda meta: (triton.cdiv(T, meta['B']),)](
161
+ x=x,
162
+ y=y,
163
+ temperature=temperature,
164
+ T=T,
165
+ D=D,
166
+ )
167
+ return y
168
+
169
+
170
+ def logsigmoid_bwd(x: torch.Tensor, dy: torch.Tensor, temperature: float = 1.) -> torch.Tensor:
171
+ T, D = x.numel(), x.shape[-1]
172
+ dx = torch.empty_like(x)
173
+ logsigmoid_bwd_kernel[lambda meta: (triton.cdiv(T, meta['B']),)](
174
+ x=x,
175
+ dx=dx,
176
+ dy=dy,
177
+ temperature=temperature,
178
+ T=T,
179
+ D=D,
180
+ )
181
+ return dx
182
+
183
+
184
+ class LogSigmoidFunction(torch.autograd.Function):
185
+
186
+ @staticmethod
187
+ @input_guard
188
+ def forward(ctx, x, temperature):
189
+ ctx.save_for_backward(x)
190
+ ctx.temperature = temperature
191
+ return logsigmoid_fwd(x, temperature)
192
+
193
+ @staticmethod
194
+ @input_guard
195
+ def backward(ctx, dy):
196
+ x, = ctx.saved_tensors
197
+ return logsigmoid_bwd(x, dy, ctx.temperature), None
198
+
199
+
200
+ def logsigmoid(x: torch.Tensor, temperature: float = 1.) -> torch.Tensor:
201
+ return LogSigmoidFunction.apply(x, temperature)
202
+
203
+
204
+ @triton.autotune(
205
+ configs=[
206
+ triton.Config({'B': bs}, num_warps=num_warps)
207
+ for bs in [512, 1024, 2048, 4096, 8192]
208
+ for num_warps in NUM_WARPS_AUTOTUNE
209
+ ],
210
+ key=['D'],
211
+ **autotune_cache_kwargs,
212
+ )
213
+ @triton.jit(do_not_specialize=['T'])
214
+ def swish_fwd_kernel(
215
+ x, y,
216
+ T,
217
+ B: tl.constexpr,
218
+ D: tl.constexpr,
219
+ ):
220
+ pid = tl.program_id(0)
221
+ offs = pid * B + tl.arange(0, B)
222
+ mask = offs < T
223
+ x_val = tl.load(x + offs, mask=mask, other=0.).to(tl.float32)
224
+ s = 1.0 / (1.0 + exp(-x_val))
225
+ y_val = x_val * s
226
+ tl.store(y + offs, y_val.to(y.dtype.element_ty), mask=mask)
227
+
228
+
229
+ @triton.autotune(
230
+ configs=[
231
+ triton.Config({'B': bs}, num_warps=num_warps)
232
+ for bs in [512, 1024, 2048, 4096, 8192]
233
+ for num_warps in NUM_WARPS_AUTOTUNE
234
+ ],
235
+ key=['D'],
236
+ **autotune_cache_kwargs,
237
+ )
238
+ @triton.jit(do_not_specialize=['T'])
239
+ def swish_bwd_kernel(
240
+ x, dy, dx,
241
+ T,
242
+ B: tl.constexpr,
243
+ D: tl.constexpr,
244
+ ):
245
+ pid = tl.program_id(0)
246
+ offs = pid * B + tl.arange(0, B)
247
+ mask = offs < T
248
+ x_val = tl.load(x + offs, mask=mask, other=0.).to(tl.float32)
249
+ g_val = tl.load(dy + offs, mask=mask, other=0.).to(tl.float32)
250
+ s = 1.0 / (1.0 + exp(-x_val))
251
+ dx_val = g_val * s * (1.0 + x_val * (1.0 - s))
252
+ tl.store(dx + offs, dx_val.to(dx.dtype.element_ty), mask=mask)
253
+
254
+
255
+ def swish_fwd(x: torch.Tensor) -> torch.Tensor:
256
+ T, D = x.numel(), x.shape[-1]
257
+ y = torch.empty_like(x)
258
+ swish_fwd_kernel[lambda meta: (triton.cdiv(T, meta['B']),)](x, y, T=T, D=D)
259
+ return y
260
+
261
+
262
+ def swish_bwd(x: torch.Tensor, dy: torch.Tensor) -> torch.Tensor:
263
+ T, D = x.numel(), x.shape[-1]
264
+ dx = torch.empty_like(x)
265
+ swish_bwd_kernel[lambda meta: (triton.cdiv(T, meta['B']),)](x, dy, dx, T=T, D=D)
266
+ return dx
267
+
268
+
269
+ class SwishFunction(torch.autograd.Function):
270
+
271
+ @staticmethod
272
+ def forward(ctx, x):
273
+ ctx.save_for_backward(x)
274
+ return swish_fwd(x)
275
+
276
+ @staticmethod
277
+ def backward(ctx, dout):
278
+ x, = ctx.saved_tensors
279
+ return swish_bwd(x, dout)
280
+
281
+
282
+ swish = SwishFunction.apply
283
+
284
+ # 1/sqrt(2*pi)-> 0.3989423
285
+ # 1/sqrt(2) -> 0.70710678
286
+ # sqrt(2/pi) -> 0.79788456
287
+
288
+
289
+ # this function is tanh approximation of gelu
290
+ # actual gelu is:
291
+ # x * 0.5 * (1.0 + torch.erf(x * 0.70710678))
292
+ @torch.compile
293
+ def bias_gelu(y, bias):
294
+ x = bias + y
295
+ return (x * 0.5 * (1.0 + torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x)))).to(dtype=y.dtype)
296
+
297
+
298
+ # gradient of tanh approximation of gelu
299
+ # gradient of actual gelu is:
300
+ # 0.5 * (1. + torch.erf(x * 0.70710678)) + 0.3989423 * x * torch.exp(-0.5 * x * x)
301
+ @torch.compile
302
+ def bias_gelu_bwd(g, y, bias):
303
+ """Assume that y has shape (B, D=D) and bias has shape (D)"""
304
+ x = bias + y
305
+ tanh_out = torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x))
306
+ # sqrt(2/pi) * 3 * 0.044715 -> 0.1070322243
307
+ ff = 0.5 * x * ((1 - tanh_out * tanh_out) * (0.79788456 + 0.1070322243 * x * x)) + 0.5 * (
308
+ 1 + tanh_out
309
+ )
310
+ grad_y = ff * g
311
+ return grad_y.to(dtype=y.dtype), grad_y.sum(dim=(0), dtype=bias.dtype)
312
+
313
+
314
+ class GeLUFunction(torch.autograd.Function):
315
+
316
+ @staticmethod
317
+ # bias is an optional argument
318
+ def forward(ctx, input, bias):
319
+ ctx.save_for_backward(input, bias)
320
+ return bias_gelu(input, bias)
321
+
322
+ @staticmethod
323
+ def backward(ctx, grad_output):
324
+ input, bias = ctx.saved_tensors
325
+ tmp = bias_gelu_bwd(grad_output, input, bias)
326
+ return tmp, tmp
327
+
328
+
329
+ bias_gelu_impl = GeLUFunction.apply
330
+
331
+
332
+ # this function is tanh approximation of gelu
333
+ # actual gelu is:
334
+ # x * 0.5 * (1.0 + torch.erf(x * 0.70710678))
335
+ @torch.compile
336
+ def gelu_fwd(x):
337
+ return (x * 0.5 * (1.0 + torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x)))).to(dtype=x.dtype)
338
+
339
+
340
+ # gradient of tanh approximation of gelu
341
+ # gradient of actual gelu is:
342
+ # 0.5 * (1. + torch.erf(x * 0.70710678)) + 0.3989423 * x * torch.exp(-0.5 * x * x)
343
+ @torch.compile
344
+ def gelu_bwd(g, x):
345
+ tanh_out = torch.tanh(0.79788456 * x * (1 + 0.044715 * x * x))
346
+ # sqrt(2/pi) * 3 * 0.044715 -> 0.1070322243
347
+ ff = 0.5 * x * ((1 - tanh_out * tanh_out) * (0.79788456 + 0.1070322243 * x * x)) + 0.5 * (
348
+ 1 + tanh_out
349
+ )
350
+ return (ff * g).to(dtype=x.dtype)
351
+
352
+
353
+ class FastGeLUFunction(torch.autograd.Function):
354
+ @staticmethod
355
+ # bias is an optional argument
356
+ def forward(ctx, input):
357
+ ctx.save_for_backward(input)
358
+ return gelu_fwd(input)
359
+
360
+ @staticmethod
361
+ def backward(ctx, grad_output):
362
+ (input,) = ctx.saved_tensors
363
+ tmp = gelu_bwd(grad_output, input)
364
+ return tmp
365
+
366
+
367
+ fast_gelu_impl = FastGeLUFunction.apply
368
+
369
+
370
+ @torch.compile
371
+ def relu_bwd(g, x):
372
+ return torch.where(x >= 0, g, 0.0).to(dtype=x.dtype)
373
+
374
+
375
+ @torch.compile
376
+ def sqrelu_fwd(x):
377
+ r = F.relu(x.float())
378
+ return (r * r).to(dtype=x.dtype)
379
+
380
+
381
+ @torch.compile
382
+ def sqrelu_bwd(g, x):
383
+ return (2.0 * g * F.relu(x.float())).to(dtype=x.dtype)
384
+
385
+
386
+ class SquaredReLUFunction(torch.autograd.Function):
387
+
388
+ @staticmethod
389
+ def forward(ctx, input):
390
+ ctx.save_for_backward(input)
391
+ return sqrelu_fwd(input)
392
+
393
+ @staticmethod
394
+ def backward(ctx, grad_output):
395
+ input, = ctx.saved_tensors
396
+ return sqrelu_bwd(grad_output, input)
397
+
398
+
399
+ sqrelu = SquaredReLUFunction.apply
400
+
401
+
402
+ @triton.autotune(
403
+ configs=[
404
+ triton.Config({'B': bs}, num_warps=num_warps)
405
+ for bs in [512, 1024, 2048, 4096, 8192]
406
+ for num_warps in NUM_WARPS_AUTOTUNE
407
+ ],
408
+ key=['D'],
409
+ **autotune_cache_kwargs,
410
+ )
411
+ @triton.jit(do_not_specialize=['T'])
412
+ def swiglu_fwd_kernel(
413
+ x, y, z,
414
+ T,
415
+ B: tl.constexpr,
416
+ D: tl.constexpr,
417
+ ):
418
+ pid = tl.program_id(0)
419
+ offs = pid * B + tl.arange(0, B)
420
+ mask = offs < T
421
+ x_val = tl.load(x + offs, mask=mask, other=0.).to(tl.float32)
422
+ y_val = tl.load(y + offs, mask=mask, other=0.).to(tl.float32)
423
+ s = 1.0 / (1.0 + exp(-x_val))
424
+ z_val = x_val * s * y_val
425
+ tl.store(z + offs, z_val.to(z.dtype.element_ty), mask=mask)
426
+
427
+
428
+ @triton.heuristics({
429
+ 'HAS_WEIGHT': lambda args: args['z'] is not None,
430
+ })
431
+ @triton.autotune(
432
+ configs=[
433
+ triton.Config({'B': bs}, num_warps=num_warps)
434
+ for bs in [512, 1024, 2048, 4096, 8192]
435
+ for num_warps in NUM_WARPS_AUTOTUNE
436
+ ],
437
+ key=['D'],
438
+ **autotune_cache_kwargs,
439
+ )
440
+ @triton.jit(do_not_specialize=['T'])
441
+ def swiglu_fwdbwd_kernel(
442
+ x, y, g, dx, dy, z,
443
+ T,
444
+ B: tl.constexpr,
445
+ D: tl.constexpr,
446
+ HAS_WEIGHT: tl.constexpr,
447
+ ):
448
+ pid = tl.program_id(0)
449
+ offs = pid * B + tl.arange(0, B)
450
+ mask = offs < T
451
+ x_val = tl.load(x + offs, mask=mask, other=0.).to(tl.float32)
452
+ y_val = tl.load(y + offs, mask=mask, other=0.).to(tl.float32)
453
+ g_val = tl.load(g + offs, mask=mask, other=0.).to(tl.float32)
454
+
455
+ s = 1.0 / (1.0 + exp(-x_val))
456
+ x_s = x_val * s
457
+ dx_val = g_val * s * (1.0 + x_val * (1.0 - s)) * y_val
458
+ dy_val = g_val * x_s
459
+
460
+ tl.store(dx + offs, dx_val.to(dx.dtype.element_ty), mask=mask)
461
+ tl.store(dy + offs, dy_val.to(dy.dtype.element_ty), mask=mask)
462
+ if HAS_WEIGHT:
463
+ z_val = x_s * y_val
464
+ tl.store(z + offs, z_val.to(z.dtype.element_ty), mask=mask)
465
+
466
+
467
+ def swiglu_fwd(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor:
468
+ T, D = x.numel(), x.shape[-1]
469
+ z = torch.empty_like(x)
470
+ swiglu_fwd_kernel[lambda meta: (triton.cdiv(T, meta['B']),)](x, y, z, T=T, D=D)
471
+ return z
472
+
473
+
474
+ def swiglu_fwdbwd(x: torch.Tensor, y: torch.Tensor, g: torch.Tensor, use_weight: bool = False):
475
+ T, D = x.numel(), x.shape[-1]
476
+ dx = torch.empty_like(x)
477
+ dy = torch.empty_like(x)
478
+ if use_weight:
479
+ # recomputed for weight grad
480
+ z = torch.empty_like(x)
481
+ else:
482
+ z = None
483
+ swiglu_fwdbwd_kernel[lambda meta: (triton.cdiv(T, meta['B']),)](x, y, g, dx, dy, z, T=T, D=D)
484
+ if use_weight:
485
+ return dx, dy, z
486
+ return dx, dy
487
+
488
+
489
+ class SwiGLUFunction(torch.autograd.Function):
490
+ r"""
491
+ Swish-Gated Linear Unit (SwiGLU) function.
492
+
493
+ .. math::
494
+ \text{SwiGLU}(x, y) = swish(x) * y = \frac{x}{1 + \exp(-x)} * y
495
+ """
496
+
497
+ @staticmethod
498
+ def forward(ctx, x, y):
499
+ ctx.save_for_backward(x, y)
500
+ return swiglu_fwd(x, y)
501
+
502
+ @staticmethod
503
+ def backward(ctx, dout):
504
+ x, y = ctx.saved_tensors
505
+ return swiglu_fwdbwd(x, y, dout)
506
+
507
+
508
+ class SwiGLULinearFunction(torch.autograd.Function):
509
+ r"""
510
+ Swish-Gated Linear Unit (SwiGLU) function followed by a linear transformation.
511
+
512
+ .. math::
513
+ \text{SwiGLULinear}(x, y, W, b) = (swish(x) * y) W + b
514
+
515
+ This simple wrap discards the intermediate results of SwiGLU(x, y) to save memory.
516
+ """
517
+
518
+ @staticmethod
519
+ @autocast_custom_fwd
520
+ def forward(ctx, x, y, weight, bias):
521
+ z = swiglu_fwd(x, y)
522
+ out = F.linear(z, weight, bias)
523
+ # We don't store z, will be recomputed in the backward pass to save memory
524
+ ctx.save_for_backward(x, y, weight)
525
+ ctx.linear_bias_is_none = bias is None
526
+ return out
527
+
528
+ @staticmethod
529
+ @autocast_custom_bwd
530
+ def backward(ctx, dout, *args):
531
+ x, y, weight = ctx.saved_tensors
532
+ dout = dout.reshape(-1, dout.shape[-1])
533
+ dz = F.linear(dout, weight.t()).view_as(x)
534
+ dx, dy, z = swiglu_fwdbwd(x, y, dz, use_weight=True)
535
+ dlinear_weight = torch.einsum("bo,bi->oi", dout, z.reshape(-1, z.shape[-1]))
536
+ dlinear_bias = None if ctx.linear_bias_is_none else dout.sum(0)
537
+ return dx, dy, dlinear_weight, dlinear_bias
538
+
539
+
540
+ swiglu = SwiGLUFunction.apply
541
+
542
+
543
+ swiglu_linear = SwiGLULinearFunction.apply
544
+
545
+
546
+ ACT2FN = {
547
+ 'relu': F.relu,
548
+ 'sigmoid': sigmoid,
549
+ 'logsigmoid': logsigmoid,
550
+ 'silu': swish,
551
+ 'swish': swish,
552
+ 'sqrelu': sqrelu,
553
+ 'gelu': fast_gelu_impl,
554
+ 'bias_gelu': bias_gelu_impl,
555
+ }
code/flash-linear-attention/fla/modules/convolution.py ADDED
@@ -0,0 +1,1167 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+ import math
4
+ import warnings
5
+
6
+ import torch
7
+ import torch.nn as nn
8
+ import torch.nn.functional as F
9
+ import triton
10
+ import triton.language as tl
11
+ from einops import rearrange
12
+
13
+ from fla.ops.utils import prepare_chunk_indices, prepare_sequence_ids
14
+ from fla.utils import autotune_cache_kwargs, get_multiprocessor_count, input_guard, is_amd
15
+
16
+ NUM_WARPS_AUTOTUNE = [2, 4, 8, 16] if is_amd else [4, 8, 16, 32]
17
+ STATIC_WARPS = 32 if not is_amd else 16
18
+
19
+
20
+ try:
21
+ from causal_conv1d import causal_conv1d_fn
22
+ from causal_conv1d import causal_conv1d_update as causal_conv1d_update_cuda
23
+ except ImportError:
24
+ causal_conv1d_fn = None
25
+ causal_conv1d_update_cuda = None
26
+
27
+
28
+ @triton.heuristics({
29
+ 'HAS_WEIGHT': lambda args: args['weight'] is not None,
30
+ 'HAS_BIAS': lambda args: args['bias'] is not None,
31
+ 'HAS_RESIDUAL': lambda args: args['residual'] is not None,
32
+ 'USE_INITIAL_STATE': lambda args: args['initial_state'] is not None,
33
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
34
+ })
35
+ @triton.autotune(
36
+ configs=[
37
+ triton.Config({'BD': BD}, num_warps=num_warps)
38
+ for BD in [16, 32, 64, 128]
39
+ for num_warps in NUM_WARPS_AUTOTUNE
40
+ ],
41
+ key=['D', 'W', 'NB'],
42
+ **autotune_cache_kwargs,
43
+ )
44
+ @triton.jit
45
+ def causal_conv1d_fwd_kernel(
46
+ x,
47
+ y,
48
+ weight,
49
+ bias,
50
+ residual,
51
+ cu_seqlens,
52
+ initial_state,
53
+ chunk_indices,
54
+ B,
55
+ T,
56
+ D: tl.constexpr,
57
+ W: tl.constexpr,
58
+ BT: tl.constexpr,
59
+ BW: tl.constexpr,
60
+ BD: tl.constexpr,
61
+ NB: tl.constexpr,
62
+ ACTIVATION: tl.constexpr,
63
+ HAS_WEIGHT: tl.constexpr,
64
+ HAS_BIAS: tl.constexpr,
65
+ HAS_RESIDUAL: tl.constexpr,
66
+ USE_INITIAL_STATE: tl.constexpr,
67
+ IS_VARLEN: tl.constexpr,
68
+ ):
69
+ i_d, i_t, i_b = tl.program_id(0), tl.program_id(1), tl.program_id(2)
70
+
71
+ if IS_VARLEN:
72
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
73
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
74
+ T = eos - bos
75
+ else:
76
+ i_n = i_b
77
+ bos, eos = (i_b * T).to(tl.int64), (i_b * T + T).to(tl.int64)
78
+
79
+ o_d = i_d * BD + tl.arange(0, BD)
80
+ o_w = tl.arange(0, BW) + W - BW
81
+ m_d = o_d < D
82
+ m_w = o_w >= 0
83
+
84
+ if HAS_WEIGHT:
85
+ # [BD, BW]
86
+ b_w = tl.load(weight + o_d[:, None] * W + o_w, mask=m_d[:, None] & m_w, other=0).to(tl.float32)
87
+
88
+ b_y = tl.zeros((BT, BD), dtype=tl.float32)
89
+ if not USE_INITIAL_STATE:
90
+ for i_w in tl.static_range(-W + 1, 1):
91
+ p_yi = tl.make_block_ptr(x + bos * D, (T, D), (D, 1), (i_t * BT + i_w, i_d * BD), (BT, BD), (1, 0))
92
+ # [BT, BD]
93
+ b_yi = tl.load(p_yi, boundary_check=(0, 1)).to(tl.float32)
94
+ if HAS_WEIGHT:
95
+ b_yi *= tl.sum(b_w * (o_w == (i_w + W - 1)), 1)
96
+ b_y += b_yi
97
+ elif i_t * BT >= W:
98
+ # to make Triton compiler happy, we need to copy codes
99
+ for i_w in tl.static_range(-W + 1, 1):
100
+ p_yi = tl.make_block_ptr(x + bos * D, (T, D), (D, 1), (i_t * BT + i_w, i_d * BD), (BT, BD), (1, 0))
101
+ # [BT, BD]
102
+ b_yi = tl.load(p_yi, boundary_check=(0, 1)).to(tl.float32)
103
+ if HAS_WEIGHT:
104
+ b_yi *= tl.sum(b_w * (o_w == (i_w + W - 1)), 1)
105
+ b_y += b_yi
106
+ else:
107
+ o_t = i_t * BT + tl.arange(0, BT)
108
+ for i_w in tl.static_range(-W + 1, 1):
109
+ o_x = o_t + i_w
110
+ m_x = ((o_x >= 0) & (o_x < T))[:, None] & m_d
111
+ m_c = ((o_x + W >= 0) & (o_x < 0))[:, None] & m_d
112
+
113
+ b_yi = tl.load(x + bos * D + o_x[:, None] * D + o_d, mask=m_x, other=0).to(tl.float32)
114
+
115
+ b_yi += tl.load(initial_state + i_n * D*W + o_d * W + (o_x + W)[:, None], mask=m_c, other=0).to(tl.float32)
116
+
117
+ if HAS_WEIGHT:
118
+ b_yi *= tl.sum(b_w * (o_w == (i_w + W - 1)), 1)
119
+ b_y += b_yi
120
+
121
+ if HAS_BIAS:
122
+ b_y += tl.load(bias + o_d, mask=m_d).to(tl.float32)
123
+
124
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
125
+ b_y = b_y * tl.sigmoid(b_y)
126
+
127
+ if HAS_RESIDUAL:
128
+ p_residual = tl.make_block_ptr(residual + bos * D, (T, D), (D, 1), (i_t * BT, i_d * BD), (BT, BD), (1, 0))
129
+ b_residual = tl.load(p_residual, boundary_check=(0, 1))
130
+ b_y += b_residual
131
+
132
+ p_y = tl.make_block_ptr(y + bos * D, (T, D), (D, 1), (i_t * BT, i_d * BD), (BT, BD), (1, 0))
133
+ tl.store(p_y, tl.cast(b_y, dtype=p_y.dtype.element_ty, fp_downcast_rounding='rtne'), boundary_check=(0, 1))
134
+
135
+
136
+ @triton.heuristics({
137
+ 'HAS_WEIGHT': lambda args: args['dw'] is not None,
138
+ 'HAS_BIAS': lambda args: args['db'] is not None,
139
+ 'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
140
+ 'USE_FINAL_STATE': lambda args: args['dht'] is not None,
141
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
142
+ })
143
+ @triton.autotune(
144
+ configs=[
145
+ triton.Config({'BD': BD}, num_warps=num_warps)
146
+ for BD in [16, 32, 64, 128]
147
+ for num_warps in [4, 8, 16, 32]
148
+ ],
149
+ key=['D', 'W', 'NB'],
150
+ **autotune_cache_kwargs,
151
+ )
152
+ @triton.jit
153
+ def causal_conv1d_bwd_kernel(
154
+ x,
155
+ y,
156
+ weight,
157
+ initial_state,
158
+ dh0,
159
+ dht,
160
+ dy,
161
+ dx,
162
+ dw,
163
+ db,
164
+ cu_seqlens,
165
+ chunk_indices,
166
+ B,
167
+ T,
168
+ D: tl.constexpr,
169
+ W: tl.constexpr,
170
+ BT: tl.constexpr,
171
+ BW: tl.constexpr,
172
+ BD: tl.constexpr,
173
+ NB: tl.constexpr,
174
+ ACTIVATION: tl.constexpr,
175
+ HAS_WEIGHT: tl.constexpr,
176
+ HAS_BIAS: tl.constexpr,
177
+ USE_INITIAL_STATE: tl.constexpr,
178
+ USE_FINAL_STATE: tl.constexpr,
179
+ IS_VARLEN: tl.constexpr,
180
+ ):
181
+ i_d, i_t, i_b = tl.program_id(0), tl.program_id(1), tl.program_id(2)
182
+ if IS_VARLEN:
183
+ i_tg = i_t
184
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
185
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
186
+ T = eos - bos
187
+ else:
188
+ i_tg = i_b * tl.num_programs(1) + i_t
189
+ i_n = i_b
190
+ bos, eos = (i_b * T).to(tl.int64), (i_b * T + T).to(tl.int64)
191
+
192
+ o_d = i_d * BD + tl.arange(0, BD)
193
+ o_w = tl.arange(0, BW) + W - BW
194
+ m_d = o_d < D
195
+ m_w = o_w >= 0
196
+
197
+ if HAS_WEIGHT:
198
+ p_x = tl.make_block_ptr(x + bos * D, (T, D), (D, 1), (i_t * BT, i_d * BD), (BT, BD), (1, 0))
199
+ b_x = tl.load(p_x, boundary_check=(0, 1))
200
+ # [BD, BW]
201
+ b_w = tl.load(weight + o_d[:, None] * W + o_w, mask=m_d[:, None] & m_w, other=0)
202
+
203
+ b_dx = tl.zeros((BT, BD), dtype=tl.float32)
204
+ if HAS_BIAS:
205
+ b_db = tl.zeros((BD,), dtype=tl.float32)
206
+
207
+ if not USE_FINAL_STATE:
208
+ for i_w in tl.static_range(0, W):
209
+ p_dy = tl.make_block_ptr(dy + bos * D, (T, D), (D, 1), (i_t * BT + i_w, i_d * BD), (BT, BD), (1, 0))
210
+ # [BT, BD]
211
+ b_dy = tl.load(p_dy, boundary_check=(0, 1)).to(tl.float32)
212
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
213
+ p_y = tl.make_block_ptr(y + bos * D, (T, D), (D, 1), (i_t * BT + i_w, i_d * BD), (BT, BD), (1, 0))
214
+ b_y = tl.load(p_y, boundary_check=(0, 1)).to(tl.float32)
215
+ b_ys = tl.sigmoid(b_y)
216
+ b_dy = b_dy * b_ys * (1 + b_y * (1 - b_ys))
217
+ b_wdy = b_dy
218
+ if HAS_WEIGHT:
219
+ # [BT, BD]
220
+ b_wdy = b_wdy * tl.sum(b_w * (o_w == (W - i_w - 1)), 1)
221
+ # [BD]
222
+ b_dw = tl.sum(b_dy * b_x, 0)
223
+ tl.store(dw + i_tg * D*W + o_d * W + W - i_w - 1, b_dw.to(dw.dtype.element_ty), mask=m_d)
224
+ if HAS_BIAS and i_w == 0:
225
+ b_db += tl.sum(b_dy, 0)
226
+ b_dx += b_wdy
227
+ elif i_t * BT >= W:
228
+ # to make Triton compiler happy, we need to copy codes
229
+ for i_w in tl.static_range(0, W):
230
+ p_dy = tl.make_block_ptr(dy + bos * D, (T, D), (D, 1), (i_t * BT + i_w, i_d * BD), (BT, BD), (1, 0))
231
+ # [BT, BD]
232
+ b_dy = tl.load(p_dy, boundary_check=(0, 1)).to(tl.float32)
233
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
234
+ p_y = tl.make_block_ptr(y + bos * D, (T, D), (D, 1), (i_t * BT + i_w, i_d * BD), (BT, BD), (1, 0))
235
+ b_y = tl.load(p_y, boundary_check=(0, 1)).to(tl.float32)
236
+ b_ys = tl.sigmoid(b_y)
237
+ b_dy = b_dy * b_ys * (1 + b_y * (1 - b_ys))
238
+ b_wdy = b_dy
239
+ if HAS_WEIGHT:
240
+ # [BT, BD]
241
+ b_wdy = b_wdy * tl.sum(b_w * (o_w == (W - i_w - 1)), 1)
242
+ # [BD]
243
+ b_dw = tl.sum(b_dy * b_x, 0)
244
+ tl.store(dw + i_tg * D*W + o_d * W + W - i_w - 1, b_dw.to(dw.dtype.element_ty), mask=m_d)
245
+ if HAS_BIAS and i_w == 0:
246
+ b_db += tl.sum(b_dy, 0)
247
+ b_dx += b_wdy
248
+ else:
249
+ # which may use initial state
250
+ o_t = i_t * BT + tl.arange(0, BT)
251
+ for i_w in tl.static_range(0, W):
252
+ p_dy = tl.make_block_ptr(dy + bos * D, (T, D), (D, 1), (i_t * BT + i_w, i_d * BD), (BT, BD), (1, 0))
253
+ b_dy_shift = tl.load(p_dy, boundary_check=(0, 1)).to(tl.float32)
254
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
255
+ p_y = tl.make_block_ptr(y + bos * D, (T, D), (D, 1), (i_t * BT + i_w, i_d * BD), (BT, BD), (1, 0))
256
+ b_y_shift = tl.load(p_y, boundary_check=(0, 1)).to(tl.float32)
257
+ b_ys = tl.sigmoid(b_y_shift)
258
+ b_dy_shift = b_dy_shift * b_ys * (1 + b_y_shift * (1 - b_ys))
259
+ if HAS_WEIGHT:
260
+ # gradient comes from x:sum_t dy[t+i_w] * x[t]
261
+ b_dw = tl.sum(b_dy_shift * b_x, 0)
262
+ # index of cache:c = W - i_w + t
263
+ if USE_INITIAL_STATE:
264
+ mask_head_rows = (o_t < i_w)
265
+ # dy_head = dy[t]
266
+ b_dy_head = tl.load(dy + bos * D + o_t[:, None] * D + o_d, mask=(mask_head_rows[:, None] & m_d[None, :]),
267
+ other=0.0).to(tl.float32)
268
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
269
+ # use y[t] (not y[t+i_w])
270
+ b_y_head = tl.load(y + bos * D + o_t[:, None] * D + o_d,
271
+ mask=(mask_head_rows[:, None] & m_d[None, :]), other=0.0).to(tl.float32)
272
+ b_ys_head = tl.sigmoid(b_y_head)
273
+ b_dy_head = b_dy_head * b_ys_head * (1 + b_y_head * (1 - b_ys_head))
274
+ o_c = W - i_w + o_t
275
+ # index 0 is padding 0
276
+ mask_c = (mask_head_rows & (o_c >= 1) & (o_c < W))
277
+ b_xc = tl.load(initial_state + i_n * D * W + o_d[None, :] * W + o_c[:, None],
278
+ mask=(mask_c[:, None] & m_d[None, :]), other=0.0).to(tl.float32)
279
+ # add the gradient comes from initial_state
280
+ b_dw += tl.sum(b_dy_head * b_xc, 0)
281
+ tl.store(dw + i_tg * D * W + o_d * W + W - i_w - 1, b_dw.to(dw.dtype.element_ty), mask=m_d)
282
+
283
+ if HAS_BIAS and i_w == 0:
284
+ b_db += tl.sum(b_dy_shift, 0)
285
+ b_wdy = b_dy_shift if not HAS_WEIGHT else (b_dy_shift * tl.sum(b_w * (o_w == (W - i_w - 1)), 1))
286
+ b_dx += b_wdy
287
+
288
+ if USE_INITIAL_STATE:
289
+ p_dy0 = tl.make_block_ptr(dy + bos * D, (T, D), (D, 1), (i_t * BT, i_d * BD), (BT, BD), (1, 0))
290
+ b_dy0 = tl.load(p_dy0, boundary_check=(0, 1)).to(tl.float32)
291
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
292
+ p_y0 = tl.make_block_ptr(y + bos * D, (T, D), (D, 1), (i_t * BT, i_d * BD), (BT, BD), (1, 0))
293
+ b_y0 = tl.load(p_y0, boundary_check=(0, 1)).to(tl.float32)
294
+ b_ys0 = tl.sigmoid(b_y0)
295
+ b_dy0 = b_dy0 * b_ys0 * (1 + b_y0 * (1 - b_ys0))
296
+ # index 0 is padding 0, skip calculation
297
+ for i_w in tl.static_range(1, W):
298
+ m_rows = (o_t < i_w)
299
+ if HAS_WEIGHT:
300
+ # [BT]
301
+ w_idx_rows = i_w - 1 - o_t
302
+ # [BT, BW]
303
+ w_mask = (o_w[None, :] == w_idx_rows[:, None])
304
+ w_pick = tl.sum(b_w[None, :, :] * w_mask[:, None, :], 2)
305
+ else:
306
+ w_pick = 1.0
307
+ contrib = (b_dy0 * w_pick).to(tl.float32)
308
+ contrib = tl.where(m_rows[:, None] & m_d[None, :], contrib, 0.0)
309
+ # [BD]
310
+ b_dh0_s = tl.sum(contrib, 0)
311
+ # dh0: [NT, B, D, W]
312
+ tl.store(dh0 + i_t * B * D * W + i_n * D * W + o_d * W + i_w,
313
+ b_dh0_s.to(dh0.dtype.element_ty, fp_downcast_rounding='rtne'), mask=m_d)
314
+
315
+ if HAS_BIAS:
316
+ b_db = tl.cast(b_db, dtype=db.dtype.element_ty, fp_downcast_rounding='rtne')
317
+ tl.store(db + i_tg * D + o_d, b_db, mask=m_d)
318
+
319
+ if USE_FINAL_STATE:
320
+ if i_t * BT + BT >= T-W:
321
+ start_tok = max(0, T - (W - 1))
322
+ offset = i_t * BT + tl.arange(0, BT)
323
+ tok_idx = offset - start_tok
324
+ mask = (offset >= start_tok) & (offset < T)
325
+ w_idx = 1 + tok_idx
326
+ dht_off = i_n * D * W + o_d[None, :] * W + w_idx[:, None]
327
+ b_dht = tl.load(dht + dht_off, mask=mask[:, None] & m_d[None, :], other=0.).to(tl.float32)
328
+ b_dx += b_dht
329
+
330
+ p_dx = tl.make_block_ptr(dx + bos * D, (T, D), (D, 1), (i_t * BT, i_d * BD), (BT, BD), (1, 0))
331
+ tl.store(p_dx, tl.cast(b_dx, dtype=p_dx.dtype.element_ty, fp_downcast_rounding='rtne'), boundary_check=(0, 1))
332
+
333
+
334
+ @triton.heuristics({
335
+ 'USE_INITIAL_STATE': lambda args: args['cache'] is not None,
336
+ 'HAS_WEIGHT': lambda args: args['weight'] is not None,
337
+ 'HAS_BIAS': lambda args: args['bias'] is not None,
338
+ 'HAS_RESIDUAL': lambda args: args['residual'] is not None,
339
+ })
340
+ @triton.jit
341
+ def causal_conv1d_update_kernel(
342
+ x,
343
+ cache,
344
+ residual,
345
+ y,
346
+ weight,
347
+ bias,
348
+ D: tl.constexpr,
349
+ W: tl.constexpr,
350
+ BD: tl.constexpr,
351
+ BW: tl.constexpr,
352
+ ACTIVATION: tl.constexpr,
353
+ USE_INITIAL_STATE: tl.constexpr,
354
+ HAS_WEIGHT: tl.constexpr,
355
+ HAS_BIAS: tl.constexpr,
356
+ HAS_RESIDUAL: tl.constexpr,
357
+ ):
358
+ i_d, i_n = tl.program_id(0), tl.program_id(1)
359
+
360
+ o_d = i_d * BD + tl.arange(0, BD)
361
+ o_w = tl.arange(0, BW) + W - BW
362
+ m_d = o_d < D
363
+ m_w = o_w >= 0
364
+ m_c = o_w < W - 1
365
+
366
+ # [BD]
367
+ b_x = tl.load(x + i_n * D + o_d, mask=m_d, other=0).to(tl.float32)
368
+
369
+ if USE_INITIAL_STATE:
370
+ # shift the cache by 1 with the last one being discarded
371
+ p_cache = tl.make_block_ptr(cache + i_n * D*W, (D, W), (W, 1), (i_d * BD, W - BW + 1), (BD, BW), (1, 0))
372
+ # [BD, BW]
373
+ b_cache = tl.load(p_cache, boundary_check=(0, 1)).to(tl.float32)
374
+ b_cache = tl.where(m_c[None, :], b_cache, b_x[:, None])
375
+ else:
376
+ b_cache = tl.zeros((BD, BW), dtype=tl.float32)
377
+
378
+ if HAS_WEIGHT:
379
+ b_w = tl.load(weight + o_d[:, None] * W + o_w, mask=m_d[:, None] & m_w, other=0)
380
+ b_y = tl.sum(b_cache * b_w, 1)
381
+ else:
382
+ b_y = tl.sum(b_cache, 1)
383
+ if HAS_BIAS:
384
+ b_y += tl.load(bias + o_d, mask=m_d)
385
+
386
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
387
+ b_y = b_y * tl.sigmoid(b_y)
388
+
389
+ if HAS_RESIDUAL:
390
+ b_y += tl.load(residual + i_n * D + o_d, mask=m_d, other=0)
391
+
392
+ tl.store(y + i_n * D + o_d, tl.cast(b_y, dtype=y.dtype.element_ty, fp_downcast_rounding='rtne'), mask=m_d)
393
+
394
+ if USE_INITIAL_STATE:
395
+ b_cache = tl.cast(b_cache, dtype=cache.dtype.element_ty, fp_downcast_rounding='rtne')
396
+ # update the cache in-place
397
+ p_cache = tl.make_block_ptr(cache + i_n * D*W, (D, W), (W, 1), (i_d * BD, W - BW), (BD, BW), (1, 0))
398
+ tl.store(p_cache, b_cache, boundary_check=(0, 1))
399
+
400
+
401
+ @input_guard
402
+ def causal_conv1d_fwd(
403
+ x: torch.Tensor,
404
+ weight: torch.Tensor,
405
+ bias: torch.Tensor,
406
+ residual: torch.Tensor,
407
+ initial_state: torch.Tensor | None = None,
408
+ output_final_state: bool = False,
409
+ activation: str | None = None,
410
+ cu_seqlens: torch.Tensor | None = None,
411
+ ) -> torch.Tensor:
412
+ shape = x.shape
413
+ if x.shape[-1] != weight.shape[0]:
414
+ x = rearrange(x, 'b t ... -> b t (...)')
415
+ B, T, D, W = *x.shape, weight.shape[1]
416
+ BT = min(64, triton.next_power_of_2(triton.cdiv(max(16, B*T), get_multiprocessor_count(x.device.index))))
417
+ BW = triton.next_power_of_2(W)
418
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
419
+ NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT)
420
+ NB = triton.cdiv(B*T, 1024)
421
+
422
+ y = torch.empty_like(x)
423
+ def grid(meta): return (triton.cdiv(D, meta['BD']), NT, B)
424
+ causal_conv1d_fwd_kernel[grid](
425
+ x=x,
426
+ y=y,
427
+ weight=weight,
428
+ bias=bias,
429
+ residual=residual,
430
+ cu_seqlens=cu_seqlens,
431
+ initial_state=initial_state,
432
+ chunk_indices=chunk_indices,
433
+ B=B,
434
+ T=T,
435
+ D=D,
436
+ W=W,
437
+ BT=BT,
438
+ BW=BW,
439
+ NB=NB,
440
+ ACTIVATION=activation,
441
+ )
442
+ final_state = None
443
+ if output_final_state:
444
+ final_state = causal_conv1d_update_states(
445
+ x=x,
446
+ state_len=W,
447
+ initial_state=initial_state,
448
+ cu_seqlens=cu_seqlens,
449
+ )
450
+ return y.view(shape), final_state
451
+
452
+
453
+ def causal_conv1d_bwd(
454
+ x: torch.Tensor,
455
+ dy: torch.Tensor,
456
+ dht: torch.Tensor,
457
+ weight: torch.Tensor | None = None,
458
+ bias: torch.Tensor | None = None,
459
+ residual: torch.Tensor | None = None,
460
+ initial_state: torch.Tensor | None = None,
461
+ activation: str | None = None,
462
+ cu_seqlens: torch.Tensor | None = None,
463
+ ):
464
+ shape = x.shape
465
+ if x.shape[-1] != weight.shape[0]:
466
+ x = rearrange(x, 'b t ... -> b t (...)')
467
+ B, T, D = x.shape
468
+ W = weight.shape[1] if weight is not None else None
469
+ BT = min(64, triton.next_power_of_2(triton.cdiv(max(16, B*T), get_multiprocessor_count(x.device.index))))
470
+ BW = triton.next_power_of_2(W)
471
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
472
+ NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT)
473
+ NB = triton.cdiv(B*T, 1024)
474
+
475
+ y = None
476
+ if activation is not None:
477
+ y, _ = causal_conv1d_fwd(
478
+ x=x,
479
+ weight=weight,
480
+ bias=bias,
481
+ residual=None,
482
+ initial_state=initial_state,
483
+ activation=None,
484
+ cu_seqlens=cu_seqlens,
485
+ output_final_state=False,
486
+ )
487
+ dx = torch.empty_like(x)
488
+ dw = weight.new_empty(B*NT, *weight.shape, dtype=torch.float) if weight is not None else None
489
+ db = bias.new_empty(B*NT, *bias.shape, dtype=torch.float) if bias is not None else None
490
+ dr = dy if residual is not None else None
491
+ dh0 = initial_state.new_zeros(min(NT, triton.cdiv(W, BT)), *initial_state.shape) if initial_state is not None else None
492
+
493
+ def grid(meta): return (triton.cdiv(D, meta['BD']), NT, B)
494
+ causal_conv1d_bwd_kernel[grid](
495
+ x=x,
496
+ y=y,
497
+ weight=weight,
498
+ initial_state=initial_state,
499
+ dh0=dh0,
500
+ dht=dht,
501
+ dy=dy,
502
+ dx=dx,
503
+ dw=dw,
504
+ db=db,
505
+ cu_seqlens=cu_seqlens,
506
+ chunk_indices=chunk_indices,
507
+ B=B,
508
+ T=T,
509
+ D=D,
510
+ W=W,
511
+ BT=BT,
512
+ BW=BW,
513
+ NB=NB,
514
+ ACTIVATION=activation,
515
+ )
516
+ if weight is not None:
517
+ dw = dw.sum(0).to(weight)
518
+ if bias is not None:
519
+ db = db.sum(0).to(bias)
520
+ if initial_state is not None:
521
+ dh0 = dh0.sum(0, dtype=torch.float32).to(initial_state)
522
+
523
+ return dx.view(shape), dw, db, dr, dh0
524
+
525
+
526
+ @triton.heuristics({
527
+ 'USE_INITIAL_STATE': lambda args: args['initial_state'] is not None,
528
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
529
+ })
530
+ @triton.jit
531
+ def causal_conv1d_states_fwd_kernel(
532
+ x,
533
+ initial_state,
534
+ final_state,
535
+ cu_seqlens,
536
+ T,
537
+ D,
538
+ W,
539
+ BD: tl.constexpr,
540
+ BW: tl.constexpr,
541
+ USE_INITIAL_STATE: tl.constexpr,
542
+ IS_VARLEN: tl.constexpr,
543
+ ):
544
+ i_d, i_n = tl.program_id(0), tl.program_id(1)
545
+ if IS_VARLEN:
546
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
547
+ T = eos - bos
548
+ else:
549
+ bos, eos = (i_n * T).to(tl.int64), (i_n * T + T).to(tl.int64)
550
+
551
+ o_t = eos - BW + tl.arange(0, BW)
552
+ o_d = i_d * BD + tl.arange(0, BD)
553
+ o_w = W - BW + tl.arange(0, BW)
554
+ m_t = (o_t >= tl.maximum(bos, eos - W))
555
+ m_d = o_d < D
556
+ m_w = (o_w >= 0) & (o_w < W)
557
+
558
+ b_x = tl.load(x + o_t * D + o_d[:, None], mask=(m_t & m_d[:, None]), other=0)
559
+ if USE_INITIAL_STATE:
560
+ if T < BW:
561
+ o_c = W - (BW - T) + tl.arange(0, BW)
562
+ m_c = (o_c >= 0) & (o_c < W)
563
+ b_cache = tl.load(initial_state + i_n * D*W + o_d[:, None] * W + o_c, mask=m_d[:, None] & m_c, other=0)
564
+ b_x += b_cache
565
+
566
+ tl.store(final_state + i_n * D*W + o_d[:, None] * W + o_w, b_x, mask=m_d[:, None] & m_w)
567
+
568
+
569
+ @input_guard
570
+ def causal_conv1d_update_states(
571
+ x: torch.Tensor,
572
+ state_len: int,
573
+ initial_state: torch.Tensor | None = None,
574
+ cu_seqlens: torch.Tensor | None = None,
575
+ ) -> torch.Tensor:
576
+ B, T, D, W = *x.shape, state_len
577
+ N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
578
+
579
+ final_state = torch.empty(N, D, W, dtype=x.dtype, device=x.device)
580
+ BD = min(triton.next_power_of_2(D), 256)
581
+ BW = triton.next_power_of_2(W)
582
+ grid = (triton.cdiv(D, BD), N)
583
+ causal_conv1d_states_fwd_kernel[grid](
584
+ x=x,
585
+ initial_state=initial_state,
586
+ final_state=final_state,
587
+ cu_seqlens=cu_seqlens,
588
+ T=T,
589
+ D=D,
590
+ W=W,
591
+ BW=BW,
592
+ BD=BD,
593
+ )
594
+ return final_state
595
+
596
+
597
+ @input_guard
598
+ def causal_conv1d_update(
599
+ x: torch.Tensor,
600
+ cache: torch.Tensor,
601
+ residual: torch.Tensor | None = None,
602
+ weight: torch.Tensor | None = None,
603
+ bias: torch.Tensor | None = None,
604
+ activation: str | None = None,
605
+ ) -> torch.Tensor:
606
+ shape = x.shape
607
+ if weight is not None and x.shape[-1] != weight.shape[0]:
608
+ x = rearrange(x, 'b t ... -> b t (...)')
609
+ *_, D = x.shape
610
+ N = x.numel() // D
611
+ W = weight.shape[1] if weight is not None else None
612
+ BD = 8
613
+ BW = triton.next_power_of_2(W)
614
+
615
+ y = torch.empty_like(x)
616
+ # NOTE: autotuning is disabled as cache is updated in-place
617
+ def grid(meta): return (triton.cdiv(D, meta['BD']), N)
618
+ causal_conv1d_update_kernel[grid](
619
+ x=x,
620
+ cache=cache,
621
+ residual=residual,
622
+ y=y,
623
+ weight=weight,
624
+ bias=bias,
625
+ D=D,
626
+ W=W,
627
+ BD=BD,
628
+ BW=BW,
629
+ ACTIVATION=activation,
630
+ num_warps=STATIC_WARPS,
631
+ )
632
+ return y.view(shape), cache
633
+
634
+
635
+ class CausalConv1dFunction(torch.autograd.Function):
636
+
637
+ @staticmethod
638
+ @input_guard
639
+ def forward(
640
+ ctx,
641
+ x: torch.Tensor,
642
+ weight: torch.Tensor | None = None,
643
+ bias: torch.Tensor | None = None,
644
+ residual: torch.Tensor | None = None,
645
+ initial_state: torch.Tensor | None = None,
646
+ output_final_state: bool | None = False,
647
+ activation: str | None = None,
648
+ cu_seqlens: torch.Tensor | None = None,
649
+ ):
650
+ ctx.activation = activation
651
+ ctx.cu_seqlens = cu_seqlens
652
+ ctx.save_for_backward(x, weight, bias, residual, initial_state)
653
+ y, final_state = causal_conv1d_fwd(
654
+ x=x,
655
+ weight=weight,
656
+ bias=bias,
657
+ residual=residual,
658
+ initial_state=initial_state,
659
+ output_final_state=output_final_state,
660
+ activation=activation,
661
+ cu_seqlens=cu_seqlens,
662
+ )
663
+ return y, final_state
664
+
665
+ @staticmethod
666
+ @input_guard
667
+ def backward(ctx, dy: torch.Tensor, dht: torch.Tensor | None = None):
668
+ x, weight, bias, residual, initial_state = ctx.saved_tensors
669
+ dx, dw, db, dr, dh0 = causal_conv1d_bwd(
670
+ x=x,
671
+ dy=dy,
672
+ dht=dht,
673
+ weight=weight,
674
+ bias=bias,
675
+ residual=residual,
676
+ initial_state=initial_state,
677
+ activation=ctx.activation,
678
+ cu_seqlens=ctx.cu_seqlens,
679
+ )
680
+ return dx, dw, db, dr, dh0, None, None, None
681
+
682
+
683
+ @input_guard
684
+ def causal_conv1d(
685
+ x: torch.Tensor,
686
+ weight: torch.Tensor | None = None,
687
+ bias: torch.Tensor | None = None,
688
+ residual: torch.Tensor | None = None,
689
+ initial_state: torch.Tensor | None = None,
690
+ output_final_state: bool | None = False,
691
+ activation: str | None = None,
692
+ backend: str | None = 'triton',
693
+ cu_seqlens: torch.Tensor | None = None,
694
+ **kwargs,
695
+ ):
696
+ """
697
+ A causal 1D convolution implementation that powers Mamba/Mamba2 and DeltaNet architectures.
698
+
699
+ When a residual connection is provided, this implements the Canon operation
700
+ described in the paper at https://papers.ssrn.com/sol3/papers.cfm?abstract_id=5240330.
701
+
702
+ Args:
703
+ x (torch.Tensor):
704
+ Input tensor of shape [B, T, D].
705
+ weight (Optional[torch.Tensor]):
706
+ Weight tensor of shape [D, W]. Default: `None`.
707
+ bias (Optional[torch.Tensor]):
708
+ Bias tensor of shape [D]. Default: `None`.
709
+ residual (Optional[torch.Tensor]):
710
+ Residual tensor of shape [B, T, D]. Default: `None`.
711
+ initial_state (Optional[torch.Tensor]):
712
+ Initial state tensor of shape [N, D, W],
713
+ where `N` is the number of sequences in the batch and `W` is the kernel size.
714
+ If provided, the initial state is used to initialize the cache. Default: `None`.
715
+ output_final_state (Optional[bool]):
716
+ Whether to output the final state of shape [N, D, W]. Default: `False`.
717
+ activation (Optional[str]):
718
+ Activations applied to output, only `swish`/`silu` or `None` (i.e., no activation) are supported.
719
+ Default: `None`.
720
+ backend (Optional[str]):
721
+ Specifies the backend to use for the convolution operation. Supported values are `'cuda'` and `'triton'`.
722
+ Default: `'triton'`.
723
+ cu_seqlens (Optional[torch.Tensor]):
724
+ Cumulative sequence lengths (optional)
725
+
726
+ Returns:
727
+ Tuple of (output, final_state).
728
+ If `output_final_state` is `False`, the final state is `None`.
729
+ """
730
+
731
+ if backend == 'triton':
732
+ y, final_state = CausalConv1dFunction.apply(
733
+ x,
734
+ weight,
735
+ bias,
736
+ residual,
737
+ initial_state,
738
+ output_final_state,
739
+ activation,
740
+ cu_seqlens,
741
+ )
742
+ return y, final_state
743
+
744
+ B, _, D, W = *x.shape, weight.shape[-1]
745
+ N = B if cu_seqlens is None else len(cu_seqlens) - 1
746
+ x = rearrange(x, 'b t d -> b d t')
747
+
748
+ # check if cu_seqlens and cache are both provided
749
+ # Sequence index for each token. Used for varlen.
750
+ # Suppose a batch consists of two sequences with lengths 3 and 4,
751
+ # seq_idx=[0, 0, 0, 1, 1, 1, 1] for this batch.
752
+ # NOTE: No need to provide this arg if `cu_seqlens` is passed.
753
+ # This arg is just for BC, and will be removed in the future.
754
+ # [B, T]
755
+ seq_idx = kwargs.get('seq_idx')
756
+ if cu_seqlens is not None and seq_idx is None:
757
+ seq_idx = prepare_sequence_ids(cu_seqlens).to(torch.int32).unsqueeze(0)
758
+
759
+ # equivalent to:
760
+ # y = _conv_forward(x, weight, bias)[..., :x.shape[-1]]
761
+ # if activation is not None:
762
+ # y = ACT2FN[activation](x)
763
+
764
+ cache, initial_state = initial_state, None
765
+ if cache is not None:
766
+ # To make causal-conv1d happy
767
+ initial_state = (
768
+ cache[:, :, -(W-1):] # [N, D, W-1]
769
+ .transpose(1, 2).contiguous() # [N, W-1, D] and stride(2)==1
770
+ .transpose(1, 2) # [N, D, W-1] and stride(1)==1
771
+ )
772
+
773
+ result = causal_conv1d_fn(
774
+ x=x,
775
+ weight=weight,
776
+ bias=bias,
777
+ activation=activation,
778
+ seq_idx=seq_idx,
779
+ initial_states=initial_state,
780
+ return_final_states=output_final_state,
781
+ )
782
+ y, final_state = result if output_final_state else (result, None)
783
+ y = rearrange(y, 'b d t -> b t d')
784
+ if output_final_state:
785
+ cache = x.new_zeros(N, D, W)
786
+ cache[:, :, -W+1:].copy_(final_state[:, :, -W+1:])
787
+ if residual is not None:
788
+ y.add_(residual)
789
+
790
+ return y, cache
791
+
792
+
793
+ class ShortConvolution(nn.Conv1d):
794
+ """Short convolution layer for efficient causal convolution operations.
795
+
796
+ This class implements a depthwise separable 1D convolution with causal padding,
797
+ designed for efficient sequence processing. It supports multiple backends (Triton/CUDA)
798
+ and optional activation functions.
799
+
800
+ Args:
801
+ hidden_size (int): Number of input/output channels (must be equal for depthwise conv)
802
+ kernel_size (int): Size of the convolution kernel
803
+ bias (bool, optional): Whether to include learnable bias. Defaults to False.
804
+ activation (Optional[str], optional): Activation function ('silu' or 'swish'). Defaults to 'silu'.
805
+ backend (Optional[str], optional): Backend implementation ('triton' or 'cuda'). Defaults to 'triton'.
806
+ device (Optional[torch.device], optional): Device to place the layer on. Defaults to None.
807
+ dtype (Optional[torch.dtype], optional): Data type for layer parameters. Defaults to None.
808
+ **kwargs: Additional keyword arguments (deprecated 'use_fast_conv1d' supported for compatibility)
809
+
810
+ Attributes:
811
+ hidden_size (int): Number of channels
812
+ activation (Optional[str]): Selected activation function
813
+ backend (str): Actual backend being used (may differ from input due to availability)
814
+
815
+ Note:
816
+ - Uses depthwise convolution (groups=hidden_size) for efficiency
817
+ - Applies causal padding (kernel_size-1) to ensure no future information leakage
818
+ - Falls back to Triton backend if CUDA backend is unavailable
819
+ """
820
+
821
+ def __init__(
822
+ self,
823
+ hidden_size: int,
824
+ kernel_size: int,
825
+ bias: bool = False,
826
+ activation: str | None = 'silu',
827
+ backend: str | None = 'triton',
828
+ device: torch.device | None = None,
829
+ dtype: torch.dtype | None = None,
830
+ **kwargs,
831
+ ):
832
+ super().__init__(
833
+ in_channels=hidden_size,
834
+ out_channels=hidden_size,
835
+ kernel_size=kernel_size,
836
+ groups=hidden_size,
837
+ bias=bias,
838
+ padding=kernel_size - 1,
839
+ device=device,
840
+ dtype=dtype,
841
+ )
842
+
843
+ self.hidden_size = hidden_size
844
+ self.activation = None
845
+
846
+ if activation is not None:
847
+ assert activation in ['silu', 'swish'], f"Activation `{activation}` not supported yet."
848
+ self.activation = activation
849
+
850
+ if 'use_fast_conv1d' in kwargs:
851
+ warnings.warn(
852
+ "The `use_fast_conv1d` parameter is deprecated and will be ignored. "
853
+ "Please use the `backend` parameter instead.",
854
+ )
855
+ import os
856
+ self.backend = os.environ.get('FLA_CONV_BACKEND', backend)
857
+ if backend not in ['cuda', 'triton']:
858
+ raise ValueError(f"Invalid backend: {backend}, must be one of ['cuda', 'triton']")
859
+ if backend == 'cuda':
860
+ if causal_conv1d_fn is None:
861
+ warnings.warn(
862
+ "The `backend` parameter is set to `cuda`, but `causal_conv1d_fn` is not available. "
863
+ "Switching to the Triton implementation instead. "
864
+ "Consider installing `causal_conv1d` to enable the CUDA backend.",
865
+ )
866
+ self.backend = 'triton'
867
+
868
+ def extra_repr(self):
869
+ s = ('{in_channels}, {out_channels}, kernel_size={kernel_size}'
870
+ ', stride={stride}')
871
+ if self.padding != (0,) * len(self.padding):
872
+ s += ', padding={padding}'
873
+ if self.dilation != (1,) * len(self.dilation):
874
+ s += ', dilation={dilation}'
875
+ if self.output_padding != (0,) * len(self.output_padding):
876
+ s += ', output_padding={output_padding}'
877
+ if self.groups != 1:
878
+ s += ', groups={groups}'
879
+ if self.bias is None:
880
+ s += ', bias=False'
881
+ if self.padding_mode != 'zeros':
882
+ s += ', padding_mode={padding_mode}'
883
+ if self.activation is not None:
884
+ s += ', activation={activation}'
885
+ s += f', backend={self.backend}'
886
+ return s.format(**self.__dict__)
887
+
888
+ def forward(
889
+ self,
890
+ x: torch.Tensor,
891
+ residual: torch.Tensor | None = None,
892
+ mask: torch.Tensor | None = None,
893
+ cache: torch.Tensor | None = None,
894
+ output_final_state: bool = False,
895
+ cu_seqlens: torch.LongTensor | None = None,
896
+ **kwargs,
897
+ ) -> tuple[torch.Tensor, torch.Tensor]:
898
+ """
899
+ Args:
900
+ x (`torch.Tensor`):
901
+ Tensor of shape `[B, T, D]`. `B` must be 1 if `seq_idx` is provided.
902
+ residual (`Optional[torch.Tensor]`):
903
+ Residual tensor of shape `[B, T, D]`. Default: `None`.
904
+ mask (`Optional[torch.Tensor]`):
905
+ Attention mask dealing with padded positions.
906
+ cache (`Optional[torch.Tensor]`):
907
+ Previous cache tensor of shape `[N, D, W]`, where `W` is the kernel size.
908
+ If provided, the cache is updated **inplace**.
909
+ output_final_state (Optional[bool]):
910
+ Whether to output the final state of shape `[N, D, W]`. Default: `False`.
911
+ cu_seqlens (Optional[torch.LongTensor]):
912
+ Cumulative sequence lengths for each batch. Used for varlen. Default: `None`.
913
+ Shape: [B+1]
914
+
915
+ Returns:
916
+ Tensor of shape `[B, T, D]`.
917
+ """
918
+
919
+ B, T, *_ = x.shape
920
+ N = B if cu_seqlens is None else len(cu_seqlens) - 1
921
+ if mask is not None:
922
+ if cu_seqlens is not None:
923
+ raise ValueError("`mask` and `cu_seqlens` cannot be provided at the same time")
924
+ x = x.mul_(mask.unsqueeze(-1))
925
+
926
+ # in decoding phase, the cache (if provided) is updated inplace
927
+ if B * T == N:
928
+ y, cache = self.step(
929
+ x=x,
930
+ residual=residual,
931
+ cache=cache,
932
+ output_final_state=output_final_state,
933
+ cu_seqlens=cu_seqlens,
934
+ )
935
+ return y, cache
936
+
937
+ # cuda backend do not support:
938
+ # 1. both `cu_seqlens` and `cache` being provided
939
+ # 2. both `cu_seqlens` and `output_final_state` being provided
940
+ if self.backend == 'cuda' and (
941
+ (cu_seqlens is not None and cache is not None) or
942
+ (cu_seqlens is not None and output_final_state)
943
+ ):
944
+ warnings.warn(
945
+ "The CUDA backend does not support both `cu_seqlens` and `cache` being provided, "
946
+ "or both `cu_seqlens` and `output_final_state` being provided. "
947
+ "Switching to the Triton backend instead. ",
948
+ stacklevel=2,
949
+ )
950
+ self.backend = 'triton'
951
+
952
+ return causal_conv1d(
953
+ x=x,
954
+ weight=rearrange(self.weight, "d 1 w -> d w"),
955
+ bias=self.bias,
956
+ residual=residual,
957
+ initial_state=cache,
958
+ output_final_state=output_final_state,
959
+ activation=self.activation,
960
+ backend=self.backend,
961
+ cu_seqlens=cu_seqlens,
962
+ **kwargs,
963
+ )
964
+
965
+ def step(
966
+ self,
967
+ x: torch.Tensor,
968
+ residual: torch.Tensor,
969
+ cache: torch.Tensor,
970
+ output_final_state: bool = False,
971
+ cu_seqlens: torch.LongTensor | None = None,
972
+ ):
973
+ B, _, D, W = *x.shape, self.kernel_size[0]
974
+ N = B if cu_seqlens is None else len(cu_seqlens) - 1
975
+ if output_final_state and cache is None:
976
+ cache = x.new_zeros(N, D, W)
977
+ # NOTE: we follow the fast mode that updates the cache in-place
978
+ if self.backend == 'triton':
979
+ return causal_conv1d_update(
980
+ x=x,
981
+ cache=cache,
982
+ residual=residual,
983
+ weight=rearrange(self.weight, "d 1 w -> d w"),
984
+ bias=self.bias,
985
+ activation=self.activation,
986
+ )
987
+
988
+ shape = x.shape
989
+ x = x.squeeze(0) if cu_seqlens is not None else x.squeeze(1)
990
+ # equivalent to:
991
+ # cache.copy_(cache.roll(shifts=-1, dims=-1))
992
+ # cache[:, :, -1] = x
993
+ # y = torch.sum(cache * rearrange(self.weight, "d 1 w -> d w"), dim=-1)
994
+ y = causal_conv1d_update_cuda(
995
+ x=x,
996
+ conv_state=cache,
997
+ weight=rearrange(self.weight, "d 1 w -> d w"),
998
+ bias=self.bias,
999
+ activation=self.activation,
1000
+ )
1001
+ y = y.view(shape)
1002
+ if residual is not None:
1003
+ y.add_(residual)
1004
+ return y, cache
1005
+
1006
+ @property
1007
+ def state_size(self) -> int:
1008
+ return self.hidden_size * self.kernel_size
1009
+
1010
+
1011
+ def fft_conv(u, k, dropout_mask, gelu=True, k_rev=None):
1012
+ seqlen = u.shape[-1]
1013
+ fft_size = 2 * seqlen
1014
+ k_f = torch.fft.rfft(k, n=fft_size) / fft_size
1015
+ if k_rev is not None:
1016
+ k_rev_f = torch.fft.rfft(k_rev, n=fft_size) / fft_size
1017
+ k_f = k_f + k_rev_f.conj()
1018
+ u_f = torch.fft.rfft(u.to(dtype=k.dtype), n=fft_size)
1019
+
1020
+ if len(u.shape) > 3:
1021
+ k_f = k_f.unsqueeze(1)
1022
+ y = torch.fft.irfft(u_f * k_f, n=fft_size, norm="forward")[..., :seqlen]
1023
+
1024
+ out = y + u
1025
+ if gelu:
1026
+ out = F.gelu(out)
1027
+ if dropout_mask is not None:
1028
+ return (out * rearrange(dropout_mask, "b H -> b H 1")).to(dtype=u.dtype)
1029
+ else:
1030
+ return out.to(dtype=u.dtype)
1031
+
1032
+
1033
+ class LongConvolution(nn.Module):
1034
+ """
1035
+ LongConvolution applies a convolution operation on the input tensor using a fixed
1036
+ filter of length max_len.
1037
+ The filter is learned during training and is applied using FFT convolution.
1038
+
1039
+ Args:
1040
+ hidden_size (int): The number of expected features in the input and output.
1041
+ max_len (int): The maximum sequence length.
1042
+
1043
+ Returns:
1044
+ y: [batch_size, seq_len, hidden_size] tensor
1045
+ """
1046
+
1047
+ def __init__(
1048
+ self,
1049
+ hidden_size: int,
1050
+ max_len: int,
1051
+ **kwargs,
1052
+ ):
1053
+ """
1054
+ Initializes the LongConvolution module.
1055
+ Args:
1056
+ hidden_size (int): The number of expected features in the input and output.
1057
+ max_len (int): The maximum sequence length.
1058
+ """
1059
+ super().__init__()
1060
+ self.hidden_size = hidden_size
1061
+ self.filter = nn.Parameter(torch.randn(self.hidden_size, max_len), requires_grad=True)
1062
+
1063
+ def forward(self, x: torch.Tensor, *args, **kwargs):
1064
+ """
1065
+ Applies the LongConvolution operation on the input tensor.
1066
+ Args:
1067
+ x: [batch_size, seq_len, hidden_size] tensor
1068
+ Returns:
1069
+ y: [batch_size, seq_len, hidden_size] tensor
1070
+ """
1071
+ x = x.transpose(1, 2)
1072
+ y = fft_conv(x, self.filter, dropout_mask=None, gelu=False)
1073
+ y = y.transpose(1, 2)
1074
+ return y.to(dtype=x.dtype)
1075
+
1076
+
1077
+ class PositionalEmbedding(nn.Module):
1078
+ def __init__(self, emb_dim: int, seq_len: int, **kwargs):
1079
+ """Complex exponential positional embeddings for implicit long convolution filters."""
1080
+ super().__init__()
1081
+
1082
+ self.seq_len = seq_len
1083
+ # The time embedding fed to the filteres is normalized so that t_f = 1
1084
+ t = torch.linspace(0, 1, self.seq_len)[None, :, None] # 1, L, 1
1085
+
1086
+ if emb_dim > 1:
1087
+ bands = (emb_dim - 1) // 2
1088
+ # To compute the right embeddings we use the "proper" linspace
1089
+ t_rescaled = torch.linspace(0, seq_len - 1, seq_len)[None, :, None]
1090
+ w = 2 * math.pi * t_rescaled / seq_len # 1, L, 1
1091
+
1092
+ f = torch.linspace(1e-4, bands - 1, bands)[None, None]
1093
+ z = torch.exp(-1j * f * w)
1094
+ z = torch.cat([t, z.real, z.imag], dim=-1)
1095
+ self.z = nn.Parameter(z, requires_grad=False)
1096
+
1097
+ def forward(self, L):
1098
+ return self.z[:, :L]
1099
+
1100
+
1101
+ class ImplicitLongConvolution(nn.Module):
1102
+ """
1103
+ Long convolution with implicit filter parameterized by an MLP.
1104
+
1105
+ Args:
1106
+ hidden_size (int):
1107
+ The number of expected features in the input and output.
1108
+ max_len (int):
1109
+ The maximum sequence length.
1110
+ d_emb (Optional[int]):
1111
+ The dimension of the positional embeddings. Must be odd and greater or equal to 3 (time, sine and cosine).
1112
+ Defaults to 3.
1113
+ d_hidden (Optional[int]):
1114
+ The number of features in the hidden layer of the MLP. Defaults to 16.
1115
+
1116
+ Attributes:
1117
+ pos_emb (`PositionalEmbedding`): The positional embedding layer.
1118
+ mlp (`nn.Sequential`): The MLP that parameterizes the implicit filter.
1119
+
1120
+ """
1121
+
1122
+ def __init__(
1123
+ self,
1124
+ hidden_size: int,
1125
+ max_len: int,
1126
+ d_emb: int = 3,
1127
+ d_hidden: int = 16,
1128
+ **kwargs,
1129
+ ):
1130
+ """
1131
+ Long convolution with implicit filter parameterized by an MLP.
1132
+
1133
+
1134
+ """
1135
+ super().__init__()
1136
+ self.hidden_size = hidden_size
1137
+ self.d_emb = d_emb
1138
+
1139
+ assert (
1140
+ d_emb % 2 != 0 and d_emb >= 3
1141
+ ), "d_emb must be odd and greater or equal to 3 (time, sine and cosine)"
1142
+ self.pos_emb = PositionalEmbedding(d_emb, max_len)
1143
+
1144
+ # final linear layer
1145
+ self.mlp = nn.Sequential(
1146
+ nn.Linear(d_emb, d_hidden),
1147
+ torch.nn.ReLU(),
1148
+ nn.Linear(d_hidden, hidden_size),
1149
+ )
1150
+
1151
+ def filter(self, seq_len: int, *args, **kwargs):
1152
+ return self.mlp(self.pos_emb(seq_len)).transpose(1, 2)
1153
+
1154
+ def forward(self, x: torch.Tensor, *args, **kwargs):
1155
+ """
1156
+ Args:
1157
+ x: [batch_size, seq_len, hidden_size] tensor
1158
+
1159
+ Returns:
1160
+ y: [batch_size, seq_len, hidden_size] tensor
1161
+ """
1162
+ x = x.transpose(1, 2)
1163
+ k = self.filter(x.shape[-1])
1164
+ y = fft_conv(x, k, dropout_mask=None, gelu=False)
1165
+
1166
+ y = y.transpose(1, 2)
1167
+ return y.to(dtype=x.dtype)
code/flash-linear-attention/fla/modules/feature_map.py ADDED
@@ -0,0 +1,298 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from __future__ import annotations
3
+
4
+ import math
5
+
6
+ import torch
7
+ import torch.nn.functional as F
8
+ from torch import nn
9
+
10
+ from fla.modules.activations import fast_gelu_impl, sigmoid, sqrelu, swish
11
+ from fla.modules.layernorm import layer_norm
12
+ from fla.utils import checkpoint
13
+
14
+
15
+ @checkpoint
16
+ def flatten_diag_outer_product(x, y):
17
+ z = torch.einsum("...i,...j->...ij", x, y)
18
+ N = z.size(-1)
19
+ indicies = torch.triu_indices(N, N)
20
+ return z[..., indicies[0], indicies[1]]
21
+
22
+
23
+ @checkpoint
24
+ def flatten_diag_outer_product_off1(x, y):
25
+ z = torch.einsum("...i,...j->...ij", x, y)
26
+ N = z.size(-1)
27
+ indicies = torch.triu_indices(N, N, 1)
28
+ indices2 = torch.arange(0, N)
29
+ return z[..., indicies[0], indicies[1]], z[..., indices2, indices2]
30
+
31
+
32
+ def is_power_of_2(n):
33
+ return (n & (n - 1) == 0) and n != 0
34
+
35
+
36
+ class HedgehogFeatureMap(nn.Module):
37
+
38
+ r"""
39
+ Hedgehog feature map as introduced in
40
+ `The Hedgehog & the Porcupine: Expressive Linear Attentions with Softmax Mimicry <https://arxiv.org/abs/2402.04347>`_
41
+ """
42
+
43
+ def __init__(
44
+ self,
45
+ head_dim: int,
46
+ ) -> HedgehogFeatureMap:
47
+ super().__init__()
48
+ # Trainable map
49
+ self.layer = nn.Linear(head_dim, head_dim)
50
+ self.init_weights_()
51
+
52
+ def init_weights_(self):
53
+ """Initialize trainable map as identity"""
54
+ with torch.no_grad():
55
+ identity = torch.eye(*self.layer.weight.shape[-2:], dtype=torch.float)
56
+ self.layer.weight.copy_(identity.to(self.layer.weight))
57
+ nn.init.zeros_(self.layer.bias)
58
+
59
+ def forward(self, x: torch.Tensor):
60
+ x = self.layer(x) # shape b, h, l, d
61
+ return torch.cat([2*x, -2*x], dim=-1).softmax(-1)
62
+
63
+
64
+ class T2RFeatureMap(nn.Module):
65
+
66
+ r"""
67
+ Simple linear mapping feature map as in
68
+ `Finetuning Pretrained Transformers into RNNs <https://arxiv.org/abs/2103.13076>`_
69
+ """
70
+
71
+ def __init__(
72
+ self,
73
+ head_dim: int,
74
+ dot_dim: int = None,
75
+ bias: bool | None = False,
76
+ ) -> T2RFeatureMap:
77
+ super().__init__()
78
+ # Trainable map
79
+ if dot_dim is None:
80
+ dot_dim = head_dim
81
+
82
+ self.head_dim = head_dim
83
+ self.dot_dim = dot_dim
84
+ self.bias = bias
85
+
86
+ self.layer = nn.Linear(head_dim, dot_dim, bias=bias)
87
+
88
+ def __repr__(self) -> str:
89
+ return f"{self.__class__.__name__}(head_dim={self.head_dim}, dot_dim={self.dot_dim}, bias={self.bias})"
90
+
91
+ def forward(self, x: torch.Tensor):
92
+ return self.layer(x).relu()
93
+
94
+
95
+ class DPFPFeatureMap(nn.Module):
96
+
97
+ r"""
98
+ Deterministic Parameter-Free Projection (DPFP) feature map in
99
+ `Linear Transformers Are Secretly Fast Weight Programmers <https://arxiv.org/abs/2102.11174>`_
100
+ """
101
+
102
+ def __init__(
103
+ self,
104
+ head_dim: int,
105
+ nu: int = 4,
106
+ ) -> DPFPFeatureMap:
107
+ super().__init__()
108
+ self.nu = nu
109
+
110
+ def forward(self, x: torch.Tensor):
111
+ x = torch.cat([x.relu(), -x.relu()], dim=-1)
112
+ x_rolled = torch.cat([x.roll(shifts=j, dims=-1) for j in range(1, self.nu+1)], dim=-1)
113
+ x_repeat = torch.cat([x] * self.nu, dim=-1)
114
+ return x_repeat * x_rolled
115
+
116
+
117
+ class HadamardFeatureMap(nn.Module):
118
+ def __init__(
119
+ self,
120
+ head_dim: int,
121
+ ) -> HadamardFeatureMap:
122
+ super().__init__()
123
+ # Trainable map
124
+ self.layer1 = nn.Linear(head_dim, head_dim)
125
+ self.layer2 = nn.Linear(head_dim, head_dim)
126
+
127
+ def forward(self, x: torch.Tensor):
128
+ return self.layer1(x) * self.layer2(x)
129
+
130
+
131
+ class LearnableOuterProductFeatureMap(nn.Module):
132
+ def __init__(
133
+ self,
134
+ head_dim: int,
135
+ feature_dim: int,
136
+ ) -> LearnableOuterProductFeatureMap:
137
+ super().__init__()
138
+ # Trainable map
139
+ self.layer1 = nn.Linear(head_dim, feature_dim, bias=False)
140
+ self.layer2 = nn.Linear(head_dim, feature_dim, bias=False)
141
+ self.normalizer = feature_dim ** -0.5
142
+
143
+ def forward(self, x: torch.Tensor):
144
+ return flatten_diag_outer_product(self.layer1(x), self.layer2(x))
145
+
146
+
147
+ class LearnablePolySketchNonNegativeFeatureMap(nn.Module):
148
+
149
+ def __init__(
150
+ self,
151
+ head_dim: int,
152
+ sketch_size: int | None = None,
153
+ degree: int | None = 2,
154
+ ) -> LearnablePolySketchNonNegativeFeatureMap:
155
+ super().__init__()
156
+
157
+ assert is_power_of_2(degree) and degree >= 2, f"The degree {degree} must be a power of 2"
158
+
159
+ self.head_dim = head_dim
160
+ self.sketch_size = sketch_size if sketch_size is not None else head_dim
161
+ self.degree = degree
162
+
163
+ self.gamma = nn.Parameter(torch.ones(head_dim))
164
+ self.beta = nn.Parameter(torch.zeros(head_dim))
165
+ # NOTE: the sketch layers defined here are quite different from the original paper
166
+ # currently we simply use linear layers without any non-linear activations
167
+ self.sketches1 = nn.ModuleList([
168
+ nn.Linear(head_dim, sketch_size, bias=False),
169
+ *[nn.Linear(sketch_size, sketch_size, bias=False) for _ in range(int(math.log2(self.degree)) - 2)],
170
+ ])
171
+ self.sketches2 = nn.ModuleList([
172
+ nn.Linear(head_dim, sketch_size, bias=False),
173
+ *[nn.Linear(sketch_size, sketch_size, bias=False) for _ in range(int(math.log2(self.degree)) - 2)],
174
+ ])
175
+
176
+ def forward(self, x: torch.Tensor):
177
+ # Section 2.1
178
+ x = layer_norm(x, self.gamma, self.beta)
179
+ # first map the input to sketch size with learnable parameters
180
+ x = self.sketches1[0](x) * self.sketches2[0](x) * self.head_dim ** -0.5
181
+ for i in range(1, int(math.log2(self.degree)) - 1):
182
+ x = self.sketches1[i](x) * self.sketches2[i](x) * self.head_dim ** -0.5
183
+ # do sketch mapping for log2(p) - 1 times in total
184
+ # do p=2 mapping to ensure non-negativity
185
+ return flatten_diag_outer_product(x, x)
186
+
187
+
188
+ class TaylorFeatureMap(nn.Module):
189
+ def __init__(
190
+ self,
191
+ head_dim: int,
192
+ ) -> TaylorFeatureMap:
193
+ super().__init__()
194
+ self.head_dim = head_dim
195
+ self.r2 = math.sqrt(2)
196
+ self.rd = math.sqrt(self.head_dim)
197
+ self.rrd = math.sqrt(self.rd)
198
+
199
+ def forward(self, x: torch.Tensor):
200
+ x2_1, x2_2 = flatten_diag_outer_product_off1(x, x)
201
+ return torch.cat([torch.ones_like(x[..., 0:1]), x / self.rrd, x2_2 / (self.rd * self.r2), x2_1 / self.rd], dim=-1)
202
+
203
+
204
+ class RebasedFeatureMap(nn.Module):
205
+
206
+ def __init__(
207
+ self,
208
+ head_dim: int,
209
+ use_gamma: bool | None = True,
210
+ use_beta: bool | None = True,
211
+ normalize: bool | None = True,
212
+ ) -> RebasedFeatureMap:
213
+ super().__init__()
214
+
215
+ self.head_dim = head_dim
216
+ self.use_gamma = use_gamma
217
+ self.use_beta = use_beta
218
+ self.normalize = normalize
219
+
220
+ self.gamma = None
221
+ self.beta = None
222
+ if use_gamma:
223
+ self.gamma = nn.Parameter(torch.ones(head_dim))
224
+ if use_beta:
225
+ self.beta = nn.Parameter(torch.zeros(head_dim))
226
+
227
+ def forward(self, x: torch.Tensor, flatten: bool | None = True):
228
+ if self.use_beta and self.use_gamma and self.normalize:
229
+ x = layer_norm(x, self.gamma, self.beta)
230
+ elif self.normalize:
231
+ x = F.layer_norm(x, (self.head_dim,), self.gamma, self.beta)
232
+ elif self.use_gamma and self.use_beta:
233
+ x = torch.addcmul(self.beta, x, self.gamma)
234
+ elif self.use_gamma:
235
+ x = x.mul(self.gamma)
236
+ else:
237
+ raise RuntimeError(f"Not supported combination of `use_gamma`, `use_beta` and `normalize`, "
238
+ f"which is currentlt set as (`{self.use_gamma}`, `{self.use_beta}`, `{self.normalize}`)")
239
+ if not flatten:
240
+ return x
241
+ x2_1, x2_2 = flatten_diag_outer_product_off1(x, x)
242
+ # rebased use learnable parameters to approximate any quadratic function
243
+ return torch.cat([x2_2 * self.head_dim ** -0.5, x2_1 * (2 / self.head_dim) ** 0.5], dim=-1)
244
+
245
+
246
+ class ReLUFeatureMap(nn.Module):
247
+
248
+ def __init__(
249
+ self,
250
+ ) -> ReLUFeatureMap:
251
+ super().__init__()
252
+
253
+ def forward(self, x: torch.Tensor):
254
+ return F.relu(x)
255
+
256
+
257
+ class SquaredReLUFeatureMap(nn.Module):
258
+
259
+ def __init__(
260
+ self,
261
+ ) -> SquaredReLUFeatureMap:
262
+ super().__init__()
263
+
264
+ def forward(self, x: torch.Tensor):
265
+ return sqrelu(x)
266
+
267
+
268
+ class GELUFeatureMap(nn.Module):
269
+
270
+ def __init__(
271
+ self,
272
+ ) -> GELUFeatureMap:
273
+ super().__init__()
274
+
275
+ def forward(self, x: torch.Tensor):
276
+ return fast_gelu_impl(x)
277
+
278
+
279
+ class SwishFeatureMap(nn.Module):
280
+
281
+ def __init__(
282
+ self,
283
+ ) -> SwishFeatureMap:
284
+ super().__init__()
285
+
286
+ def forward(self, x: torch.Tensor):
287
+ return swish(x)
288
+
289
+
290
+ class SigmoidFeatureMap(nn.Module):
291
+
292
+ def __init__(
293
+ self,
294
+ ) -> SigmoidFeatureMap:
295
+ super().__init__()
296
+
297
+ def forward(self, x: torch.Tensor):
298
+ return sigmoid(x)
code/flash-linear-attention/fla/modules/fused_bitlinear.py ADDED
@@ -0,0 +1,633 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+ # Implementations of BitLinear layer with fused LayerNorm and quantized Linear layer.
4
+ # [The Era of 1-bit LLMs: All Large Language Models are in 1.58 Bits](https://arxiv.org/abs/2402.17764)
5
+ # [Scalable MatMul-free Language Modeling](https://arxiv.org/abs/2406.02528)
6
+
7
+ # Code adapted from https://github.com/ridgerchu/matmulfreellm/
8
+
9
+ from __future__ import annotations
10
+
11
+ import math
12
+
13
+ import torch
14
+ import torch.nn as nn
15
+ import torch.nn.functional as F
16
+ import triton
17
+ import triton.language as tl
18
+
19
+ from fla.modules.layernorm import RMSNorm
20
+ from fla.utils import autotune_cache_kwargs, get_multiprocessor_count, input_guard, is_amd, require_version
21
+
22
+ NUM_WARPS_AUTOTUNE = [1, 2, 4, 8, 16] if is_amd else [1, 2, 4, 8, 16, 32]
23
+
24
+
25
+ def activation_quant(x):
26
+ """
27
+ Per-token quantization to 8 bits. No grouping is needed for quantization.
28
+
29
+ Args:
30
+ x: An activation tensor with shape [n, d].
31
+
32
+ Returns:
33
+ A quantized activation tensor with shape [n, d].
34
+ """
35
+ # Compute the scale factor
36
+ scale = 127.0 / x.abs().max(dim=-1, keepdim=True).values.clamp_(min=1e-5)
37
+ # Quantize and then de-quantize the tensor
38
+ y = (x * scale).round().clamp_(-128, 127) / scale
39
+ return y
40
+
41
+
42
+ def weight_quant(w):
43
+ """
44
+ Per-tensor quantization to 1.58 bits. No grouping is needed for quantization.
45
+
46
+ Args:
47
+ w: A weight tensor with shape [d, k].
48
+
49
+ Returns:
50
+ A quantized weight tensor with shape [d, k].
51
+ """
52
+ # Compute the scale factor
53
+ scale = 1.0 / w.abs().mean().clamp_(min=1e-5)
54
+ # Quantize and then de-quantize the tensor
55
+ u = (w * scale).round().clamp_(-1, 1) / scale
56
+ return u
57
+
58
+
59
+ @triton.autotune(
60
+ configs=[
61
+ triton.Config({}, num_warps=num_warps)
62
+ for num_warps in NUM_WARPS_AUTOTUNE
63
+ ],
64
+ key=["N", "HAS_RESIDUAL", "STORE_RESIDUAL_OUT", "IS_RMS_NORM", "HAS_BIAS"],
65
+ **autotune_cache_kwargs,
66
+ )
67
+ @triton.jit
68
+ def layer_norm_fwd_kernel_quant(
69
+ X, # pointer to the input
70
+ Y, # pointer to the output
71
+ W, # pointer to the weights
72
+ B, # pointer to the biases
73
+ RESIDUAL, # pointer to the residual
74
+ RESIDUAL_OUT, # pointer to the residual
75
+ Mean, # pointer to the mean
76
+ Rstd, # pointer to the 1/std
77
+ stride_x_row, # how much to increase the pointer when moving by 1 row
78
+ stride_y_row,
79
+ stride_res_row,
80
+ stride_res_out_row,
81
+ N, # number of columns in X
82
+ eps, # epsilon to avoid division by zero
83
+ IS_RMS_NORM: tl.constexpr,
84
+ BLOCK_N: tl.constexpr,
85
+ HAS_RESIDUAL: tl.constexpr,
86
+ STORE_RESIDUAL_OUT: tl.constexpr,
87
+ HAS_WEIGHT: tl.constexpr,
88
+ HAS_BIAS: tl.constexpr,
89
+ ):
90
+ # Map the program id to the row of X and Y it should compute.
91
+ row = tl.program_id(0)
92
+ X += row * stride_x_row
93
+ Y += row * stride_y_row
94
+ if HAS_RESIDUAL:
95
+ RESIDUAL += row * stride_res_row
96
+ if STORE_RESIDUAL_OUT:
97
+ RESIDUAL_OUT += row * stride_res_out_row
98
+ # Compute mean and variance
99
+ cols = tl.arange(0, BLOCK_N)
100
+ x = tl.load(X + cols, mask=cols < N, other=0.0).to(tl.float32)
101
+ if HAS_RESIDUAL:
102
+ residual = tl.load(RESIDUAL + cols, mask=cols < N, other=0.0).to(tl.float32)
103
+ x += residual
104
+ if STORE_RESIDUAL_OUT:
105
+ tl.store(RESIDUAL_OUT + cols, x, mask=cols < N)
106
+ if not IS_RMS_NORM:
107
+ mean = tl.sum(x, axis=0) / N
108
+ tl.store(Mean + row, mean)
109
+ xbar = tl.where(cols < N, x - mean, 0.0)
110
+ var = tl.sum(xbar * xbar, axis=0) / N
111
+ else:
112
+ xbar = tl.where(cols < N, x, 0.0)
113
+ var = tl.sum(xbar * xbar, axis=0) / N
114
+ rstd = 1 / tl.sqrt(var + eps)
115
+ tl.store(Rstd + row, rstd)
116
+ # Normalize and apply linear transformation
117
+ mask = cols < N
118
+ if HAS_WEIGHT:
119
+ w = tl.load(W + cols, mask=mask).to(tl.float32)
120
+ if HAS_BIAS:
121
+ b = tl.load(B + cols, mask=mask).to(tl.float32)
122
+ x_hat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd
123
+
124
+ y = x_hat * w if HAS_WEIGHT else x_hat
125
+ if HAS_BIAS:
126
+ y = y + b
127
+
128
+ # Aply quantization to the output
129
+ scale = 127.0 / tl.maximum(tl.max(tl.abs(y), 0), 1e-5)
130
+ # Quantize and then de-quantize the tensor
131
+ y = tl.extra.cuda.libdevice.round(y * scale)
132
+ y = tl.maximum(tl.minimum(y, 127), -128) / scale
133
+
134
+ # Write output
135
+ tl.store(Y + cols, y, mask=mask)
136
+
137
+
138
+ def layer_norm_fwd_quant(
139
+ x: torch.Tensor,
140
+ weight: torch.Tensor,
141
+ bias: torch.Tensor,
142
+ eps: float,
143
+ residual: torch.Tensor = None,
144
+ out_dtype: torch.dtype = None,
145
+ residual_dtype: torch.dtype = None,
146
+ is_rms_norm: bool = False,
147
+ ):
148
+ if residual is not None:
149
+ residual_dtype = residual.dtype
150
+ M, N = x.shape
151
+ # allocate output
152
+ y = torch.empty_like(x, dtype=x.dtype if out_dtype is None else out_dtype)
153
+ if residual is not None or (residual_dtype is not None and residual_dtype != x.dtype):
154
+ residual_out = torch.empty(M, N, device=x.device, dtype=residual_dtype)
155
+ else:
156
+ residual_out = None
157
+ mean = torch.empty((M,), dtype=torch.float32, device=x.device) if not is_rms_norm else None
158
+ rstd = torch.empty((M,), dtype=torch.float32, device=x.device)
159
+ # Less than 64KB per feature: enqueue fused kernel
160
+ MAX_FUSED_SIZE = 65536 // x.element_size()
161
+ BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(N))
162
+ if N > BLOCK_N:
163
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
164
+ # heuristics for number of warps
165
+ layer_norm_fwd_kernel_quant[(M,)](
166
+ x,
167
+ y,
168
+ weight,
169
+ bias,
170
+ residual,
171
+ residual_out,
172
+ mean,
173
+ rstd,
174
+ x.stride(0),
175
+ y.stride(0),
176
+ residual.stride(0) if residual is not None else 0,
177
+ residual_out.stride(0) if residual_out is not None else 0,
178
+ N,
179
+ eps,
180
+ is_rms_norm,
181
+ BLOCK_N,
182
+ residual is not None,
183
+ residual_out is not None,
184
+ weight is not None,
185
+ bias is not None,
186
+ )
187
+ # residual_out is None if residual is None and residual_dtype == input_dtype
188
+ return y, mean, rstd, residual_out if residual_out is not None else x
189
+
190
+
191
+ @triton.heuristics({
192
+ "RECOMPUTE_OUTPUT": lambda args: args["Y"] is not None,
193
+ })
194
+ @triton.autotune(
195
+ configs=[
196
+ triton.Config({}, num_warps=num_warps)
197
+ for num_warps in NUM_WARPS_AUTOTUNE
198
+ ],
199
+ key=["N", "HAS_DRESIDUAL", "STORE_DRESIDUAL", "IS_RMS_NORM", "HAS_BIAS"],
200
+ **autotune_cache_kwargs,
201
+ )
202
+ @triton.jit
203
+ def layer_norm_bwd_kernel(
204
+ X, # pointer to the input
205
+ W, # pointer to the weights
206
+ B, # pointer to the biases
207
+ Y, # pointer to the output to be recomputed
208
+ DY, # pointer to the output gradient
209
+ DX, # pointer to the input gradient
210
+ DW, # pointer to the partial sum of weights gradient
211
+ DB, # pointer to the partial sum of biases gradient
212
+ DRESIDUAL,
213
+ DRESIDUAL_IN,
214
+ Mean, # pointer to the mean
215
+ Rstd, # pointer to the 1/std
216
+ stride_x_row, # how much to increase the pointer when moving by 1 row
217
+ stride_y_row,
218
+ stride_dy_row,
219
+ stride_dx_row,
220
+ stride_dres_row,
221
+ stride_dres_in_row,
222
+ M, # number of rows in X
223
+ N, # number of columns in X
224
+ eps, # epsilon to avoid division by zero
225
+ rows_per_program,
226
+ IS_RMS_NORM: tl.constexpr,
227
+ BLOCK_N: tl.constexpr,
228
+ HAS_DRESIDUAL: tl.constexpr,
229
+ STORE_DRESIDUAL: tl.constexpr,
230
+ HAS_WEIGHT: tl.constexpr,
231
+ HAS_BIAS: tl.constexpr,
232
+ RECOMPUTE_OUTPUT: tl.constexpr,
233
+ ):
234
+ # Map the program id to the elements of X, DX, and DY it should compute.
235
+ row_block_id = tl.program_id(0)
236
+ row_start = row_block_id * rows_per_program
237
+ cols = tl.arange(0, BLOCK_N)
238
+ mask = cols < N
239
+ X += row_start * stride_x_row
240
+ if HAS_DRESIDUAL:
241
+ DRESIDUAL += row_start * stride_dres_row
242
+ if STORE_DRESIDUAL:
243
+ DRESIDUAL_IN += row_start * stride_dres_in_row
244
+ DY += row_start * stride_dy_row
245
+ DX += row_start * stride_dx_row
246
+ if RECOMPUTE_OUTPUT:
247
+ Y += row_start * stride_y_row
248
+ if HAS_WEIGHT:
249
+ w = tl.load(W + cols, mask=mask).to(tl.float32)
250
+ dw = tl.zeros((BLOCK_N,), dtype=tl.float32)
251
+ if RECOMPUTE_OUTPUT and HAS_BIAS:
252
+ b = tl.load(B + cols, mask=mask, other=0.0).to(tl.float32)
253
+ if HAS_BIAS:
254
+ db = tl.zeros((BLOCK_N,), dtype=tl.float32)
255
+ row_end = min((row_block_id + 1) * rows_per_program, M)
256
+ for row in range(row_start, row_end):
257
+ # Load data to SRAM
258
+ x = tl.load(X + cols, mask=mask, other=0).to(tl.float32)
259
+ dy = tl.load(DY + cols, mask=mask, other=0).to(tl.float32)
260
+ if not IS_RMS_NORM:
261
+ mean = tl.load(Mean + row)
262
+ rstd = tl.load(Rstd + row)
263
+ # Compute dx
264
+ xhat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd
265
+ xhat = tl.where(mask, xhat, 0.0)
266
+ if RECOMPUTE_OUTPUT:
267
+ y = xhat * w if HAS_WEIGHT else xhat
268
+ if HAS_BIAS:
269
+ y = y + b
270
+
271
+ # Aply quantization to the output
272
+ scale = 127.0 / tl.maximum(tl.max(tl.abs(y), 0), 1e-5)
273
+ # Quantize and then de-quantize the tensor
274
+ y = tl.extra.cuda.libdevice.round(y * scale)
275
+ y = tl.maximum(tl.minimum(y, 127), -128) / scale
276
+
277
+ tl.store(Y + cols, y, mask=mask)
278
+ wdy = dy
279
+ if HAS_WEIGHT:
280
+ wdy = dy * w
281
+ dw += dy * xhat
282
+ if HAS_BIAS:
283
+ db += dy
284
+ if not IS_RMS_NORM:
285
+ c1 = tl.sum(xhat * wdy, axis=0) / N
286
+ c2 = tl.sum(wdy, axis=0) / N
287
+ dx = (wdy - (xhat * c1 + c2)) * rstd
288
+ else:
289
+ c1 = tl.sum(xhat * wdy, axis=0) / N
290
+ dx = (wdy - xhat * c1) * rstd
291
+ if HAS_DRESIDUAL:
292
+ dres = tl.load(DRESIDUAL + cols, mask=mask, other=0).to(tl.float32)
293
+ dx += dres
294
+ # Write dx
295
+ if STORE_DRESIDUAL:
296
+ tl.store(DRESIDUAL_IN + cols, dx, mask=mask)
297
+ tl.store(DX + cols, dx, mask=mask)
298
+
299
+ X += stride_x_row
300
+ if HAS_DRESIDUAL:
301
+ DRESIDUAL += stride_dres_row
302
+ if STORE_DRESIDUAL:
303
+ DRESIDUAL_IN += stride_dres_in_row
304
+ if RECOMPUTE_OUTPUT:
305
+ Y += stride_y_row
306
+ DY += stride_dy_row
307
+ DX += stride_dx_row
308
+ if HAS_WEIGHT:
309
+ tl.store(DW + row_block_id * N + cols, dw, mask=mask)
310
+ if HAS_BIAS:
311
+ tl.store(DB + row_block_id * N + cols, db, mask=mask)
312
+
313
+
314
+ def layer_norm_bwd(
315
+ dy: torch.Tensor,
316
+ x: torch.Tensor,
317
+ weight: torch.Tensor,
318
+ bias: torch.Tensor,
319
+ eps: float,
320
+ mean: torch.Tensor,
321
+ rstd: torch.Tensor,
322
+ dresidual: torch.Tensor = None,
323
+ has_residual: bool = False,
324
+ is_rms_norm: bool = False,
325
+ x_dtype: torch.dtype = None,
326
+ recompute_output: bool = False,
327
+ ):
328
+ M, N = x.shape
329
+ # allocate output
330
+ dx = torch.empty_like(x) if x_dtype is None else torch.empty(M, N, dtype=x_dtype, device=x.device)
331
+ dresidual_in = torch.empty_like(x) if has_residual and dx.dtype != x.dtype else None
332
+ y = torch.empty(M, N, dtype=dy.dtype, device=dy.device) if recompute_output else None
333
+
334
+ # Less than 64KB per feature: enqueue fused kernel
335
+ MAX_FUSED_SIZE = 65536 // x.element_size()
336
+ BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(N))
337
+ if N > BLOCK_N:
338
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
339
+ sm_count = get_multiprocessor_count(x.device.index)
340
+ _dw = torch.empty((sm_count, N), dtype=torch.float32, device=weight.device) if weight is not None else None
341
+ _db = torch.empty((sm_count, N), dtype=torch.float32, device=bias.device) if bias is not None else None
342
+ rows_per_program = math.ceil(M / sm_count)
343
+ grid = (sm_count,)
344
+ layer_norm_bwd_kernel[grid](
345
+ x,
346
+ weight,
347
+ bias,
348
+ y,
349
+ dy,
350
+ dx,
351
+ _dw,
352
+ _db,
353
+ dresidual,
354
+ dresidual_in,
355
+ mean,
356
+ rstd,
357
+ x.stride(0),
358
+ 0 if not recompute_output else y.stride(0),
359
+ dy.stride(0),
360
+ dx.stride(0),
361
+ dresidual.stride(0) if dresidual is not None else 0,
362
+ dresidual_in.stride(0) if dresidual_in is not None else 0,
363
+ M,
364
+ N,
365
+ eps,
366
+ rows_per_program,
367
+ is_rms_norm,
368
+ BLOCK_N,
369
+ dresidual is not None,
370
+ dresidual_in is not None,
371
+ weight is not None,
372
+ bias is not None,
373
+ )
374
+ dw = _dw.sum(0).to(weight.dtype) if weight is not None else None
375
+ db = _db.sum(0).to(bias.dtype) if bias is not None else None
376
+ # Don't need to compute dresidual_in separately in this case
377
+ if has_residual and dx.dtype == x.dtype:
378
+ dresidual_in = dx
379
+ return (dx, dw, db, dresidual_in) if not recompute_output else (dx, dw, db, dresidual_in, y)
380
+
381
+
382
+ class LayerNormLinearQuantFn(torch.autograd.Function):
383
+
384
+ @staticmethod
385
+ @input_guard
386
+ def forward(
387
+ ctx,
388
+ x,
389
+ norm_weight,
390
+ norm_bias,
391
+ linear_weight,
392
+ linear_bias,
393
+ residual=None,
394
+ eps=1e-6,
395
+ prenorm=False,
396
+ residual_in_fp32=False,
397
+ is_rms_norm=False,
398
+ ):
399
+ x_shape_og = x.shape
400
+ # reshape input data into 2D tensor
401
+ x = x.reshape(-1, x.shape[-1])
402
+ if residual is not None:
403
+ assert residual.shape == x_shape_og
404
+ residual = residual.reshape(-1, residual.shape[-1])
405
+ residual_dtype = residual.dtype if residual is not None else (torch.float32 if residual_in_fp32 else None)
406
+ y, mean, rstd, residual_out = layer_norm_fwd_quant(
407
+ x,
408
+ norm_weight,
409
+ norm_bias,
410
+ eps,
411
+ residual,
412
+ out_dtype=None if not torch.is_autocast_enabled() else torch.get_autocast_gpu_dtype(),
413
+ residual_dtype=residual_dtype,
414
+ is_rms_norm=is_rms_norm,
415
+ )
416
+ y = y.reshape(x_shape_og)
417
+ dtype = torch.get_autocast_gpu_dtype() if torch.is_autocast_enabled() else y.dtype
418
+ linear_weight = weight_quant(linear_weight).to(dtype)
419
+ linear_bias = linear_bias.to(dtype) if linear_bias is not None else None
420
+ out = F.linear(y.to(linear_weight.dtype), linear_weight, linear_bias)
421
+ # We don't store y, will be recomputed in the backward pass to save memory
422
+ ctx.save_for_backward(residual_out, norm_weight, norm_bias, linear_weight, mean, rstd)
423
+ ctx.x_shape_og = x_shape_og
424
+ ctx.eps = eps
425
+ ctx.is_rms_norm = is_rms_norm
426
+ ctx.has_residual = residual is not None
427
+ ctx.prenorm = prenorm
428
+ ctx.x_dtype = x.dtype
429
+ ctx.linear_bias_is_none = linear_bias is None
430
+ return out if not prenorm else (out, residual_out.reshape(x_shape_og))
431
+
432
+ @staticmethod
433
+ @input_guard
434
+ def backward(ctx, dout, *args):
435
+ x, norm_weight, norm_bias, linear_weight, mean, rstd = ctx.saved_tensors
436
+ dout = dout.reshape(-1, dout.shape[-1])
437
+ dy = F.linear(dout, linear_weight.t())
438
+ dlinear_bias = None if ctx.linear_bias_is_none else dout.sum(0)
439
+ assert dy.shape == x.shape
440
+ if ctx.prenorm:
441
+ dresidual = args[0]
442
+ dresidual = dresidual.reshape(-1, dresidual.shape[-1])
443
+ assert dresidual.shape == x.shape
444
+ else:
445
+ dresidual = None
446
+ dx, dnorm_weight, dnorm_bias, dresidual_in, y = layer_norm_bwd(
447
+ dy,
448
+ x,
449
+ norm_weight,
450
+ norm_bias,
451
+ ctx.eps,
452
+ mean,
453
+ rstd,
454
+ dresidual,
455
+ ctx.has_residual,
456
+ ctx.is_rms_norm,
457
+ x_dtype=ctx.x_dtype,
458
+ recompute_output=True,
459
+ )
460
+ dlinear_weight = torch.einsum("bo,bi->oi", dout, y)
461
+ return (
462
+ dx.reshape(ctx.x_shape_og),
463
+ dnorm_weight,
464
+ dnorm_bias,
465
+ dlinear_weight,
466
+ dlinear_bias,
467
+ dresidual_in.reshape(ctx.x_shape_og) if ctx.has_residual else None,
468
+ None,
469
+ None,
470
+ None,
471
+ None,
472
+ )
473
+
474
+
475
+ def layer_norm_linear_quant_fn(
476
+ x,
477
+ norm_weight,
478
+ norm_bias,
479
+ linear_weight,
480
+ linear_bias,
481
+ residual=None,
482
+ eps=1e-6,
483
+ prenorm=False,
484
+ residual_in_fp32=False,
485
+ is_rms_norm=False,
486
+ ):
487
+ return LayerNormLinearQuantFn.apply(
488
+ x,
489
+ norm_weight,
490
+ norm_bias,
491
+ linear_weight,
492
+ linear_bias,
493
+ residual,
494
+ eps,
495
+ prenorm,
496
+ residual_in_fp32,
497
+ is_rms_norm,
498
+ )
499
+
500
+
501
+ def rms_norm_linear_quant(
502
+ x: torch.Tensor,
503
+ norm_weight: torch.Tensor,
504
+ norm_bias: torch.Tensor,
505
+ linear_weight: torch.Tensor,
506
+ linear_bias: torch.Tensor,
507
+ residual: torch.Tensor = None,
508
+ eps: float = 1e-5,
509
+ prenorm: bool = False,
510
+ residual_in_fp32: bool = False,
511
+ ):
512
+ return layer_norm_linear_quant_fn(
513
+ x=x,
514
+ norm_weight=norm_weight,
515
+ norm_bias=norm_bias,
516
+ linear_weight=linear_weight,
517
+ linear_bias=linear_bias,
518
+ residual=residual,
519
+ eps=eps,
520
+ prenorm=prenorm,
521
+ residual_in_fp32=residual_in_fp32,
522
+ is_rms_norm=True,
523
+ )
524
+
525
+
526
+ @require_version("triton>=3.0", "Triton >= 3.0 is required to do online quantization.")
527
+ def bit_linear(x, weight, bias=None, norm_weight=None, norm_bias=None, eps=1e-8):
528
+ """
529
+ A functional version of BitLinear that applies quantization to activations and weights.
530
+
531
+ Args:
532
+ x: Input tensor with shape [n, d].
533
+ weight: Weight tensor with shape [out_features, in_features].
534
+ bias: Bias tensor with shape [out_features] (optional).
535
+ norm_weight: Weight tensor for RMS normalization with shape [in_features].
536
+ norm_bias: Bias tensor for RMS normalization with shape [in_features].
537
+ eps: A small constant for numerical stability in normalization.
538
+
539
+ Returns:
540
+ Output tensor with shape [n, out_features].
541
+ """
542
+ return layer_norm_linear_quant_fn(
543
+ x,
544
+ norm_weight,
545
+ norm_bias,
546
+ weight,
547
+ bias,
548
+ is_rms_norm=True,
549
+ )
550
+
551
+
552
+ class BitLinear(nn.Linear):
553
+ """
554
+ A custom linear layer that applies quantization on both activations and weights.
555
+ This is primarily for training; kernel optimization is needed for efficiency in deployment.
556
+ """
557
+
558
+ def __init__(
559
+ self,
560
+ in_features: int,
561
+ out_features: int,
562
+ bias: bool = False,
563
+ norm_eps: float = 1e-8,
564
+ ):
565
+ """
566
+ Initializes the BitLinear layer.
567
+
568
+ Args:
569
+ in_features: Size of each input sample.
570
+ out_features: Size of each output sample.
571
+ bias: If set to False, the layer will not learn an additive bias. Default: True.
572
+ """
573
+ # Initialize the superclass nn.Linear with the given parameters
574
+ super().__init__(in_features, out_features, bias=bias)
575
+
576
+ self.norm = RMSNorm(in_features, eps=norm_eps)
577
+
578
+ def __repr__(self) -> str:
579
+ return f"{self.__class__.__name__}({super().extra_repr()}, norm_eps={self.norm.eps})"
580
+
581
+ def forward(self, x):
582
+ """
583
+ Overrides the forward pass to include quantization.
584
+
585
+ Args:
586
+ x: An input tensor with shape [n, d].
587
+
588
+ Returns:
589
+ An output tensor with shape [n, d].
590
+ """
591
+ # Weight tensor
592
+ w = self.weight
593
+
594
+ # Apply RMS normalization to the input
595
+ x_norm = self.norm(x)
596
+
597
+ # Apply quantization to both activations and weights
598
+ # Uses Straight-Through Estimator (STE) trick with .detach() for gradient flow
599
+ x_quant = x_norm + (activation_quant(x_norm) - x_norm).detach()
600
+ w_quant = w + (weight_quant(w) - w).detach()
601
+ # Perform linear operation with quantized values
602
+ y = F.linear(x_quant, w_quant)
603
+
604
+ return y
605
+
606
+
607
+ class FusedBitLinear(BitLinear):
608
+ """
609
+ A custom linear layer that applies quantization on both activations and weights.
610
+ This is primarily for training; kernel optimization is needed for efficiency in deployment.
611
+ """
612
+
613
+ def __init__(self, in_features, out_features, bias=False):
614
+ """
615
+ Initializes the BitLinear layer.
616
+
617
+ Args:
618
+ in_features: Size of each input sample.
619
+ out_features: Size of each output sample.
620
+ bias: If set to False, the layer will not learn an additive bias. Default: True.
621
+ """
622
+ # Initialize the superclass nn.Linear with the given parameters
623
+ super().__init__(in_features, out_features, bias=bias)
624
+
625
+ def forward(self, x):
626
+ return layer_norm_linear_quant_fn(
627
+ x,
628
+ self.norm.weight,
629
+ self.norm.bias,
630
+ self.weight,
631
+ self.bias,
632
+ is_rms_norm=True,
633
+ )
code/flash-linear-attention/fla/modules/fused_cross_entropy.py ADDED
@@ -0,0 +1,418 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ # Copyright (c) 2023, Tri Dao.
3
+
4
+ from typing import Any
5
+
6
+ import torch
7
+ import torch.nn as nn
8
+ import triton
9
+ import triton.language as tl
10
+
11
+ from fla.ops.utils.op import exp, log
12
+ from fla.utils import input_guard
13
+
14
+ # `all_gather_into_tensor` and `reduce_scatter_tensor` are new placeholders for
15
+ # `_all_gather_base` and `_reduce_scatter_base`. They require the most recent
16
+ # version of PyTorch. The following 2 lines are for backward compatibility with
17
+ # older PyTorch.
18
+ if "all_gather_into_tensor" not in dir(torch.distributed):
19
+ torch.distributed.all_gather_into_tensor = torch.distributed._all_gather_base
20
+
21
+
22
+ @triton.heuristics({
23
+ "HAS_SMOOTHING": lambda args: args["label_smoothing"] > 0.0,
24
+ })
25
+ @triton.jit
26
+ def cross_entropy_fwd_kernel(
27
+ loss_ptr, # data ptrs
28
+ lse_ptr,
29
+ z_loss_ptr,
30
+ logits_ptr,
31
+ labels_ptr,
32
+ label_smoothing,
33
+ logit_scale,
34
+ lse_square_scale,
35
+ ignore_index,
36
+ total_classes,
37
+ class_start_idx, # Useful for tensor parallel when each rank only has a subset of classes
38
+ n_cols, # shapes
39
+ n_rows,
40
+ logits_row_stride, # strides
41
+ BLOCK_SIZE: tl.constexpr,
42
+ HAS_SMOOTHING: tl.constexpr,
43
+ # if SPLIT (e.g. tensor parallel), don't include the LSE in the loss since it's not the final LSE
44
+ SPLIT: tl.constexpr,
45
+ ):
46
+ row_idx = tl.program_id(0)
47
+ col_block_idx = tl.program_id(1)
48
+ logits_ptr = logits_ptr + row_idx * logits_row_stride.to(tl.int64)
49
+ col_offsets = col_block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
50
+ label_idx = tl.load(labels_ptr + row_idx)
51
+ logits = tl.load(logits_ptr + col_offsets, mask=col_offsets < n_cols, other=-float("inf"))
52
+ logits = logits.to(tl.float32) * logit_scale
53
+ max_logits = tl.max(logits, 0)
54
+ if HAS_SMOOTHING:
55
+ sum_logits = tl.sum(tl.where(col_offsets < n_cols, logits, 0.0), 0)
56
+ lse = log(tl.sum(exp(logits - max_logits), 0)) + max_logits
57
+ tl.store(lse_ptr + col_block_idx * n_rows + row_idx, lse)
58
+ if label_idx == ignore_index:
59
+ loss = 0.0
60
+ z_loss = 0.0
61
+ else:
62
+ label_idx -= class_start_idx
63
+ if label_idx >= col_block_idx * BLOCK_SIZE and label_idx < min(
64
+ n_cols, (col_block_idx + 1) * BLOCK_SIZE,
65
+ ):
66
+ logits_label = tl.load(logits_ptr + label_idx) * logit_scale
67
+ if HAS_SMOOTHING:
68
+ loss = (
69
+ (lse if not SPLIT else 0.0)
70
+ - label_smoothing * sum_logits / total_classes
71
+ - (1 - label_smoothing) * logits_label
72
+ )
73
+ else:
74
+ loss = (lse if not SPLIT else 0.0) - logits_label
75
+ else:
76
+ # If label is out of bounds, we set the CE loss to 0.0. But we still want the label_smoothing loss
77
+ if HAS_SMOOTHING:
78
+ loss = label_smoothing * ((lse if not SPLIT else 0.0) - sum_logits / total_classes)
79
+ else:
80
+ loss = 0.0
81
+ if not SPLIT:
82
+ z_loss = lse_square_scale * lse * lse
83
+ loss += z_loss
84
+ else:
85
+ z_loss = 0.0
86
+ tl.store(loss_ptr + col_block_idx * n_rows + row_idx, loss)
87
+ if not SPLIT:
88
+ tl.store(z_loss_ptr + col_block_idx * n_rows + row_idx, z_loss)
89
+
90
+
91
+ @triton.heuristics({
92
+ "HAS_SMOOTHING": lambda args: args["label_smoothing"] > 0.0,
93
+ })
94
+ @triton.jit
95
+ def cross_entropy_bwd_kernel(
96
+ dlogits_ptr, # data ptrs
97
+ dloss_ptr,
98
+ logits_ptr,
99
+ lse_ptr,
100
+ labels_ptr,
101
+ label_smoothing,
102
+ logit_scale,
103
+ lse_square_scale,
104
+ ignore_index,
105
+ total_classes,
106
+ class_start_idx, # Useful for tensor parallel when each rank only has a subset of classes
107
+ n_cols, # shapes
108
+ logits_row_stride, # strides
109
+ dlogits_row_stride,
110
+ dloss_row_stride,
111
+ BLOCK_SIZE: tl.constexpr,
112
+ HAS_SMOOTHING: tl.constexpr,
113
+ ):
114
+ row_idx = tl.program_id(0)
115
+ col_block_idx = tl.program_id(1)
116
+ logits_ptr = logits_ptr + row_idx * logits_row_stride.to(tl.int64)
117
+ dlogits_ptr = dlogits_ptr + row_idx * dlogits_row_stride.to(tl.int64)
118
+ col_offsets = col_block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
119
+ label_idx = tl.load(labels_ptr + row_idx)
120
+ if label_idx != ignore_index:
121
+ dloss = tl.load(dloss_ptr + row_idx * dloss_row_stride)
122
+ else:
123
+ dloss = 0.0
124
+ logits = tl.load(logits_ptr + col_offsets, mask=col_offsets < n_cols, other=-float("inf")).to(
125
+ tl.float32,
126
+ ) * logit_scale
127
+ lse = tl.load(lse_ptr + row_idx)
128
+ probs = exp(logits - lse)
129
+ probs += 2.0 * lse_square_scale * lse * probs
130
+ label_idx -= class_start_idx
131
+ if HAS_SMOOTHING:
132
+ smooth_negative = label_smoothing / total_classes
133
+ probs = tl.where(col_offsets == label_idx, probs - (1 - label_smoothing), probs) - smooth_negative
134
+ else:
135
+ probs = tl.where(col_offsets == label_idx, probs - 1.0, probs)
136
+ tl.store(dlogits_ptr + col_offsets, (dloss * logit_scale) * probs, mask=col_offsets < n_cols)
137
+
138
+
139
+ def fused_cross_entropy_forward(
140
+ logits: torch.Tensor,
141
+ target: torch.Tensor,
142
+ label_smoothing: float = 0.0,
143
+ logit_scale: float = 1.0,
144
+ lse_square_scale: float = 0.0,
145
+ ignore_index: int = -100,
146
+ process_group=None,
147
+ ):
148
+ n_rows, n_cols = logits.shape
149
+ assert target.shape == (n_rows,)
150
+ world_size = 1 if process_group is None else torch.distributed.get_world_size(process_group)
151
+ total_classes = world_size * n_cols
152
+ rank = 0 if process_group is None else torch.distributed.get_rank(process_group)
153
+ class_start_idx = rank * n_cols
154
+
155
+ if logits.stride(-1) != 1:
156
+ logits = logits.contiguous()
157
+ # Set these similar to https://github.com/openai/triton/blob/main/python/tutorials/02-fused-softmax.py
158
+ MAX_BLOCK_SIZE = 64 * 1024
159
+ BLOCK_SIZE = min(triton.next_power_of_2(n_cols), MAX_BLOCK_SIZE)
160
+ num_warps = (
161
+ 4
162
+ if BLOCK_SIZE < 2048
163
+ else (8 if BLOCK_SIZE < 8192 else (16 if BLOCK_SIZE < 128 * 1024 else 32))
164
+ )
165
+ # We may split the lse computation across multiple blocks, then do a reduction
166
+ # lse(local_lse) to get the final LSE. This is faster for large n_cols (e.g., > 64k)
167
+ # where having just one thread block processing more than 64k elements is slow.
168
+ split = world_size > 1 or n_cols > MAX_BLOCK_SIZE
169
+ n_splits = (n_cols + BLOCK_SIZE - 1) // BLOCK_SIZE
170
+ loss_shape = (n_splits, n_rows) if n_splits > 1 else (n_rows,)
171
+ losses = torch.empty(*loss_shape, dtype=torch.float, device=logits.device)
172
+ lse = torch.empty(*loss_shape, dtype=torch.float, device=logits.device)
173
+ z_losses = torch.empty(*loss_shape, dtype=torch.float, device=logits.device)
174
+
175
+ cross_entropy_fwd_kernel[(n_rows, n_splits)](
176
+ losses, # data ptrs
177
+ lse,
178
+ z_losses,
179
+ logits,
180
+ target,
181
+ label_smoothing,
182
+ logit_scale,
183
+ lse_square_scale,
184
+ ignore_index,
185
+ total_classes,
186
+ class_start_idx,
187
+ n_cols, # shapes
188
+ n_rows,
189
+ logits.stride(0), # strides
190
+ BLOCK_SIZE=BLOCK_SIZE, # constants
191
+ num_warps=num_warps,
192
+ SPLIT=split,
193
+ )
194
+
195
+ if split:
196
+ # If there's no label_smoothing, if target are in the vocab of this partition, losses contains
197
+ # - predicted logit, and 0 otherwise.
198
+ # If there's label_smoothing=0.1, for target in the vocab of this partition, losses contains
199
+ # -0.9 * predicted logit - 0.1 * sum logit / total_classes.
200
+ # For target not in the vocab of this partition, losses contains
201
+ # -0.1 * sum logit / total_classes.
202
+ if n_splits > 1:
203
+ lse = torch.logsumexp(lse, dim=0)
204
+ losses = losses.sum(dim=0)
205
+ if world_size > 1:
206
+ lse_allgather = torch.empty(world_size, n_rows, dtype=lse.dtype, device=lse.device)
207
+ torch.distributed.all_gather_into_tensor(lse_allgather, lse, group=process_group)
208
+ handle_losses = torch.distributed.all_reduce(
209
+ losses, op=torch.distributed.ReduceOp.SUM, group=process_group, async_op=True,
210
+ )
211
+ lse = torch.logsumexp(lse_allgather, dim=0)
212
+ handle_losses.wait()
213
+ # After the allreduce, if there's no label_smoothing, the total losses are - predicted_logit,
214
+ # we just have to add the (global) lse.
215
+ # If there's label_smoothing=0.1, the total losses are
216
+ # -0.9 * predicted_logit - 0.1 * sum logit / total_classes.
217
+ # Again, we just have to add the (global) lse.
218
+ losses += lse
219
+ if lse_square_scale != 0.0:
220
+ z_losses = lse_square_scale * lse.square()
221
+ z_losses.masked_fill_(target == ignore_index, 0.0)
222
+ losses += z_losses
223
+ else:
224
+ z_losses = torch.zeros_like(losses)
225
+ losses.masked_fill_(target == ignore_index, 0.0)
226
+
227
+ return losses, z_losses, lse, total_classes, class_start_idx
228
+
229
+
230
+ class CrossEntropyLossFunction(torch.autograd.Function):
231
+
232
+ @staticmethod
233
+ @input_guard
234
+ def forward(
235
+ ctx,
236
+ logits,
237
+ target,
238
+ label_smoothing=0.0,
239
+ logit_scale=1.0,
240
+ lse_square_scale=0.0,
241
+ ignore_index=-100,
242
+ inplace_backward=False,
243
+ process_group=None,
244
+ ):
245
+ losses, z_losses, lse, total_classes, class_start_idx = fused_cross_entropy_forward(
246
+ logits,
247
+ target,
248
+ label_smoothing,
249
+ logit_scale,
250
+ lse_square_scale,
251
+ ignore_index,
252
+ process_group,
253
+ )
254
+ ctx.save_for_backward(logits, lse, target)
255
+ ctx.mark_non_differentiable(z_losses)
256
+ ctx.label_smoothing = label_smoothing
257
+ ctx.logit_scale = logit_scale
258
+ ctx.lse_square_scale = lse_square_scale
259
+ ctx.ignore_index = ignore_index
260
+ ctx.total_classes = total_classes
261
+ ctx.class_start_idx = class_start_idx
262
+ ctx.inplace_backward = inplace_backward
263
+
264
+ return losses, z_losses
265
+
266
+ @staticmethod
267
+ @input_guard
268
+ def backward(ctx, grad_losses, grad_z_losses):
269
+ del grad_z_losses # z_losses are only for logging.
270
+
271
+ logits, lse, target = ctx.saved_tensors
272
+ dlogits = logits if ctx.inplace_backward else torch.empty_like(logits)
273
+ n_rows, n_cols = logits.shape
274
+ BLOCK_SIZE = min(triton.next_power_of_2(n_cols), 4 * 1024)
275
+ num_warps = 4 if BLOCK_SIZE < 2048 else (8 if BLOCK_SIZE < 8192 else 16)
276
+ def grid(META): return (n_rows, triton.cdiv(n_cols, META["BLOCK_SIZE"])) # noqa
277
+ cross_entropy_bwd_kernel[grid](
278
+ dlogits, # data ptrs
279
+ grad_losses,
280
+ logits,
281
+ lse,
282
+ target,
283
+ ctx.label_smoothing,
284
+ ctx.logit_scale,
285
+ ctx.lse_square_scale,
286
+ ctx.ignore_index,
287
+ ctx.total_classes,
288
+ ctx.class_start_idx,
289
+ n_cols, # shapes
290
+ logits.stride(0), # strides
291
+ dlogits.stride(0),
292
+ grad_losses.stride(0),
293
+ BLOCK_SIZE=BLOCK_SIZE, # constants
294
+ num_warps=num_warps,
295
+ )
296
+ return dlogits, None, None, None, None, None, None, None, None
297
+
298
+
299
+ def cross_entropy_loss(
300
+ logits: torch.Tensor,
301
+ target: torch.Tensor,
302
+ label_smoothing: float = 0.0,
303
+ logit_scale: float = 1.0,
304
+ lse_square_scale: float = 0.0,
305
+ ignore_index=-100,
306
+ inplace_backward: bool = False,
307
+ process_group=None,
308
+ ) -> tuple[torch.Tensor, torch.Tensor]:
309
+ """
310
+ Arguments:
311
+ logits: [batch, vocab_size]
312
+ target: [batch,]
313
+ label_smoothing: float
314
+ logit_scale: float.
315
+ Multiply logits by this scale before calculating the loss.
316
+ lse_square_scale: float.
317
+ If > 0, we add lse_square_scale * lse(logits) ^ 2 to the loss.
318
+ This is also referred to as "z-loss".
319
+ ignore_index: int.
320
+ If target == ignore_index, the loss is set to 0.0.
321
+ inplace_backward: bool.
322
+ If True, we do the backward pass in-place by modifying the logits.
323
+ This saves memory.
324
+ process_group:
325
+ if not None, we're doing Tensor Parallel: each process is responsible for
326
+ one part of the vocab. The loss will be aggregated across processes.
327
+ Returns:
328
+ losses: [batch,], float
329
+ z_losses: [batch,], float
330
+ """
331
+ return CrossEntropyLossFunction.apply(
332
+ logits,
333
+ target,
334
+ label_smoothing,
335
+ logit_scale,
336
+ lse_square_scale,
337
+ ignore_index,
338
+ inplace_backward,
339
+ process_group,
340
+ )
341
+
342
+
343
+ class FusedCrossEntropyLoss(nn.Module):
344
+ def __init__(
345
+ self,
346
+ ignore_index: int = -100,
347
+ reduction: str = "mean",
348
+ label_smoothing: float = 0.0,
349
+ logit_scale: float = 1.0,
350
+ lse_square_scale: float = 0.0,
351
+ inplace_backward: bool = False,
352
+ process_group: Any = None,
353
+ return_z_loss: bool = False,
354
+ ):
355
+ """
356
+ Arguments:
357
+ ignore_index: int. If target == ignore_index, the loss is set to 0.0.
358
+ label_smoothing: float
359
+ lse_square_scale: float. If > 0, we add lse_square_scale * lse(logits) ^ 2 to the loss.
360
+ This is also referred to as "z-loss".
361
+ inplace_backward: bool. If True, we do the backward pass in-place by modifying the logits.
362
+ This saves memory.
363
+ process_group: if not None, we're doing Tensor Parallel: each process is responsible for
364
+ one part of the vocab. The loss will be aggregated across processes.
365
+ return_z_loss: bool. If True, we return the component of the loss contributed by
366
+ the lse_square_scale value. This value is only for logging and does not support
367
+ backprop.
368
+ """
369
+ super().__init__()
370
+ if reduction not in ["mean", "none", "sum"]:
371
+ raise NotImplementedError("Only support reduction = 'mean' or 'none' or 'sum'")
372
+ self.ignore_index = ignore_index
373
+ self.reduction = reduction
374
+ self.label_smoothing = label_smoothing
375
+ self.logit_scale = logit_scale
376
+ self.lse_square_scale = lse_square_scale
377
+ self.inplace_backward = inplace_backward
378
+ self.process_group = process_group
379
+ self.return_z_loss = return_z_loss
380
+
381
+ def forward(self, input, target):
382
+ """
383
+ Arguments:
384
+ input: (batch, vocab_size)
385
+ target: (batch,)
386
+ Returns:
387
+ losses: (batch,) if reduction is 'none', else (1,), dtype float
388
+ z_loss: (batch,) if reduction is 'none', else (1,), dtype float (if self.return_z_loss)
389
+ """
390
+ assert input.is_cuda and target.is_cuda, "Only support CUDA tensors"
391
+ loss, z_loss = cross_entropy_loss(
392
+ input,
393
+ target,
394
+ label_smoothing=self.label_smoothing,
395
+ logit_scale=self.logit_scale,
396
+ lse_square_scale=self.lse_square_scale,
397
+ ignore_index=self.ignore_index,
398
+ inplace_backward=self.inplace_backward,
399
+ process_group=self.process_group,
400
+ )
401
+ if self.reduction == "mean":
402
+ loss = loss.sum() / (target != self.ignore_index).sum()
403
+ elif self.reduction == "sum":
404
+ loss = loss.sum()
405
+ else:
406
+ loss = loss
407
+
408
+ if not self.return_z_loss:
409
+ return loss
410
+
411
+ if self.reduction == "mean":
412
+ z_loss = z_loss.sum() / (target != self.ignore_index).sum()
413
+ elif self.reduction == "sum":
414
+ z_loss = z_loss.sum()
415
+ else:
416
+ z_loss = z_loss
417
+
418
+ return loss, z_loss
code/flash-linear-attention/fla/modules/fused_kl_div.py ADDED
@@ -0,0 +1,322 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+ import triton
7
+ import triton.language as tl
8
+
9
+ from fla.ops.utils.op import exp, log
10
+ from fla.utils import input_guard, is_amd
11
+
12
+ # The hard limit of TRITON_MAX_TENSOR_NUMEL is 1048576
13
+ # https://github.com/triton-lang/triton/blob/ba42a5c68fd0505f8c42f4202d53be0f8d9a5fe0/python/triton/language/core.py#L19
14
+ # However, setting limit as 65536 as in LayerNorm tutorial is faster because of less register spilling
15
+ # The optimal maximum block size depends on your hardware, your kernel, and your dtype
16
+ MAX_FUSED_SIZE = 65536 // 2
17
+ STATIC_WARPS = 32 if not is_amd else 16
18
+
19
+
20
+ @triton.jit
21
+ def kl_div_kernel(
22
+ logits,
23
+ target_logits,
24
+ loss,
25
+ s_logits,
26
+ s_loss,
27
+ reduction: tl.constexpr,
28
+ N: tl.constexpr,
29
+ V: tl.constexpr,
30
+ BV: tl.constexpr,
31
+ ):
32
+ # https://github.com/triton-lang/triton/issues/1058
33
+ # If N*V is too large, i_n * stride will overflow out of int32, so we convert to int64
34
+ i_n = tl.program_id(0).to(tl.int64)
35
+
36
+ logits += i_n * s_logits
37
+ target_logits += i_n * s_logits
38
+
39
+ # m is the max value. use the notation from the paper
40
+ sm = float('-inf')
41
+ tm = float('-inf')
42
+ # d is the sum. use the notation from the paper
43
+ sd, td = 0.0, 0.0
44
+
45
+ NV = tl.cdiv(V, BV)
46
+ for iv in range(0, NV):
47
+ o_x = iv * BV + tl.arange(0, BV)
48
+ # for student
49
+ b_sl = tl.load(logits + o_x, mask=o_x < V, other=float('-inf'))
50
+ b_sm = tl.max(b_sl)
51
+ m_new = tl.maximum(sm, b_sm)
52
+ sd = sd * exp(sm - m_new) + tl.sum(exp(b_sl - m_new))
53
+ sm = m_new
54
+ # for teacher
55
+ b_tl = tl.load(target_logits + o_x, mask=o_x < V, other=float('-inf'))
56
+ b_tm = tl.max(b_tl)
57
+ m_new = tl.maximum(tm, b_tm)
58
+ td = td * exp(tm - m_new) + tl.sum(exp(b_tl - m_new))
59
+ tm = m_new
60
+
61
+ b_loss = 0.
62
+ # KL(y_true || y) = exp(y_true) * (log(y_true) - log(y))
63
+ for iv in range(0, NV):
64
+ o_x = iv * BV + tl.arange(0, BV)
65
+ b_sl = tl.load(logits + o_x, mask=o_x < V, other=float('-inf'))
66
+ b_tl = tl.load(target_logits + o_x, mask=o_x < V, other=float('-inf'))
67
+ b_sp_log = b_sl - sm - log(sd)
68
+ b_tp_log = b_tl - tm - log(td)
69
+ b_sp = exp(b_sp_log)
70
+ b_tp = exp(b_tp_log)
71
+ b_kl = tl.where(o_x < V, b_tp * (b_tp_log - b_sp_log), 0)
72
+ b_dl = -b_tp + b_sp
73
+ b_loss += tl.sum(b_kl)
74
+ if reduction == 'batchmean':
75
+ b_dl = b_dl / N
76
+ tl.store(logits + o_x, b_dl, mask=o_x < V)
77
+
78
+ # Normalize the loss by the number of elements if reduction is 'batchmean'
79
+ if reduction == 'batchmean':
80
+ b_loss = b_loss / N
81
+
82
+ tl.store(loss + i_n * s_loss, b_loss)
83
+
84
+
85
+ @triton.jit
86
+ def elementwise_mul_kernel(
87
+ x,
88
+ g,
89
+ N: tl.constexpr,
90
+ B: tl.constexpr,
91
+ ):
92
+ """
93
+ This function multiplies each element of the tensor pointed by x with the value pointed by g.
94
+ The multiplication is performed in-place on the tensor pointed by x.
95
+
96
+ Parameters:
97
+ x:
98
+ Pointer to the input tensor.
99
+ g:
100
+ Pointer to the gradient output value.
101
+ N (int):
102
+ The number of columns in the input tensor.
103
+ B (int):
104
+ The block size for Triton operations.
105
+ """
106
+
107
+ # Get the program ID and convert it to int64 to avoid overflow
108
+ i_x = tl.program_id(0).to(tl.int64)
109
+ o_x = i_x * B + tl.arange(0, B)
110
+
111
+ # Load the gradient output value
112
+ b_g = tl.load(g)
113
+ b_x = tl.load(x + o_x, mask=o_x < N)
114
+ tl.store(x + o_x, b_x * b_g, mask=o_x < N)
115
+
116
+
117
+ def fused_kl_div_forward(
118
+ x: torch.Tensor,
119
+ target_x: torch.Tensor,
120
+ weight: torch.Tensor,
121
+ target_weight: torch.Tensor,
122
+ reduction: str = 'batchmean',
123
+ ):
124
+ device = x.device
125
+
126
+ # ideally, we would like to achieve the same memory consumption as [N, H],
127
+ # so the expected chunk size should be:
128
+ # NC = ceil(V / H)
129
+ # C = ceil(N / NC)
130
+ # for ex: N = 4096*4, V = 32000, H = 4096 ==> NC = 8, C = ceil(N / NC) = 2048
131
+ N, H, V = *x.shape, weight.shape[0]
132
+ BV = min(MAX_FUSED_SIZE, triton.next_power_of_2(V))
133
+ # TODO: in real cases, we may need to limit the number of chunks NC to
134
+ # ensure the precisions of accumulated gradients
135
+ NC = min(8, triton.cdiv(V, H))
136
+ C = triton.next_power_of_2(triton.cdiv(N, NC))
137
+ NC = triton.cdiv(N, C)
138
+
139
+ dx = torch.zeros_like(x, device=device)
140
+ dw = torch.zeros_like(weight, device=device) if weight is not None else None
141
+ # we use fp32 for loss accumulator
142
+ loss = torch.zeros(N, dtype=torch.float32, device=device)
143
+
144
+ for ic in range(NC):
145
+ start, end = ic * C, min((ic + 1) * C, N)
146
+ # [C, N]
147
+ c_sx = x[start:end]
148
+ c_tx = target_x[start:end]
149
+ # when doing matmul, use the original precision
150
+ # [C, V]
151
+ c_sl = F.linear(c_sx, weight)
152
+ c_tl = F.linear(c_tx, target_weight)
153
+
154
+ # unreduced loss
155
+ c_loss = loss[start:end]
156
+
157
+ # Here we calculate the gradient of c_sx in place so we can save memory.
158
+ kl_div_kernel[(c_sx.shape[0],)](
159
+ logits=c_sl,
160
+ target_logits=c_tl,
161
+ loss=c_loss,
162
+ s_logits=c_sl.stride(-2),
163
+ s_loss=c_loss.stride(-1),
164
+ reduction=reduction,
165
+ N=N,
166
+ V=V,
167
+ BV=BV,
168
+ num_warps=STATIC_WARPS,
169
+ )
170
+
171
+ # gradient of logits is computed in-place by the above triton kernel and is of shape: C x V
172
+ # thus dx[start: end] should be of shape: C x H
173
+ # additionally, since we are chunking the inputs, observe that the loss and gradients are calculated only
174
+ # on `n_non_ignore` tokens. However, the gradient of the input should be calculated for all tokens.
175
+ # Thus, we need an additional scaling factor of (n_non_ignore/total) to scale the gradients.
176
+ # [C, H]
177
+
178
+ dx[start:end] = torch.mm(c_sl, weight)
179
+
180
+ if weight is not None:
181
+ torch.addmm(input=dw, mat1=c_sl.t(), mat2=c_sx, out=dw)
182
+
183
+ loss = loss.sum()
184
+ return loss, dx, dw
185
+
186
+
187
+ def fused_kl_div_backward(
188
+ do: torch.Tensor,
189
+ dx: torch.Tensor,
190
+ dw: torch.Tensor,
191
+ ):
192
+ # If cross entropy is the last layer, do is 1.0. Skip the mul to save time
193
+ if torch.ne(do, torch.tensor(1.0, device=do.device)):
194
+ # We use a Triton kernel instead of a PyTorch operation because modifying inputs in-place
195
+ # for gradient storage and backward multiple times causes anomalies with PyTorch but not with Triton.
196
+ N, H = dx.shape
197
+ B = min(MAX_FUSED_SIZE, triton.next_power_of_2(H))
198
+
199
+ elementwise_mul_kernel[(triton.cdiv(N * H, B),)](
200
+ x=dx,
201
+ g=do,
202
+ N=N*H,
203
+ B=B,
204
+ num_warps=STATIC_WARPS,
205
+ )
206
+
207
+ # handle dw
208
+ if dw is not None:
209
+ V, H = dw.shape
210
+ elementwise_mul_kernel[(triton.cdiv(V * H, B),)](
211
+ x=dw,
212
+ g=do,
213
+ N=V*H,
214
+ B=B,
215
+ num_warps=STATIC_WARPS,
216
+ )
217
+
218
+ return dx, dw
219
+
220
+
221
+ class FusedKLDivLossFunction(torch.autograd.Function):
222
+
223
+ @staticmethod
224
+ @input_guard
225
+ def forward(
226
+ ctx,
227
+ x: torch.Tensor,
228
+ target_x: torch.Tensor,
229
+ weight: torch.Tensor,
230
+ target_weight: torch.Tensor,
231
+ reduction: str,
232
+ ):
233
+ loss, dx, dw = fused_kl_div_forward(
234
+ x=x,
235
+ target_x=target_x,
236
+ weight=weight,
237
+ target_weight=target_weight,
238
+ reduction=reduction,
239
+ )
240
+ ctx.save_for_backward(dx, dw)
241
+ return loss
242
+
243
+ @staticmethod
244
+ @input_guard
245
+ def backward(ctx, do):
246
+ dx, dw = ctx.saved_tensors
247
+ dx, dw = fused_kl_div_backward(do, dx, dw)
248
+ return dx, None, dw, None, None
249
+
250
+
251
+ def fused_kl_div_loss(
252
+ x: torch.Tensor,
253
+ target_x: torch.Tensor,
254
+ weight: torch.Tensor,
255
+ target_weight: torch.Tensor,
256
+ reduction: str = 'batchmean',
257
+ ) -> tuple[torch.Tensor, torch.Tensor]:
258
+ """
259
+ Args:
260
+ x (torch.Tensor): [batch_size * seq_len, hidden_size]
261
+ target_x (torch.Tensor): [batch_size * seq_len, hidden_size]
262
+ weight (torch.Tensor): [vocab_size, hidden_size]
263
+ where `vocab_size` is the number of classes.
264
+ target_weight (torch.Tensor): [vocab_size, hidden_size]
265
+ where `vocab_size` is the number of classes.
266
+ reduction:
267
+ Specifies the reduction to apply to the output: 'batchmean'. Default: 'batchmean'.
268
+ Returns:
269
+ loss
270
+ """
271
+ return FusedKLDivLossFunction.apply(
272
+ x,
273
+ target_x,
274
+ weight,
275
+ target_weight,
276
+ reduction,
277
+ )
278
+
279
+
280
+ class FusedKLDivLoss(nn.Module):
281
+
282
+ def __init__(
283
+ self,
284
+ reduction: str = 'batchmean',
285
+ ):
286
+ """
287
+ Args:
288
+ reduction:
289
+ Specifies the reduction to apply to the output: 'batchmean'. Default: 'batchmean'.
290
+ """
291
+ super().__init__()
292
+
293
+ assert reduction in ['batchmean'], f"reduction: {reduction} is not supported"
294
+
295
+ self.reduction = reduction
296
+
297
+ def forward(
298
+ self,
299
+ x: torch.Tensor,
300
+ target_x: torch.Tensor,
301
+ weight: torch.Tensor,
302
+ target_weight: torch.Tensor,
303
+ ):
304
+ """
305
+ Args:
306
+ x (torch.Tensor): [batch_size * seq_len, hidden_size]
307
+ target_x (torch.Tensor): [batch_size * seq_len, hidden_size]
308
+ weight (torch.Tensor): [vocab_size, hidden_size]
309
+ where `vocab_size` is the number of classes.
310
+ target_weight (torch.Tensor): [vocab_size, hidden_size]
311
+ where `vocab_size` is the number of classes.
312
+ Returns:
313
+ loss
314
+ """
315
+ loss = fused_kl_div_loss(
316
+ x=x,
317
+ target_x=target_x,
318
+ weight=weight,
319
+ target_weight=target_weight,
320
+ reduction=self.reduction,
321
+ )
322
+ return loss
code/flash-linear-attention/fla/modules/fused_linear_cross_entropy.py ADDED
@@ -0,0 +1,630 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ # Code adapted from
3
+ # https://github.com/linkedin/Liger-Kernel/blob/main/src/liger_kernel/ops/fused_linear_cross_entropy.py
4
+
5
+ from functools import partial
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ import torch.nn.functional as F
10
+ import triton
11
+ import triton.language as tl
12
+ try:
13
+ from torch.distributed import DeviceMesh
14
+ except ImportError:
15
+ DeviceMesh = None
16
+ try:
17
+ from torch.distributed.tensor import Replicate, Shard, distribute_module
18
+ except ImportError:
19
+ Replicate = None
20
+ Shard = None
21
+ distribute_module = None
22
+ try:
23
+ from torch.distributed.tensor.parallel import ParallelStyle
24
+ except ImportError:
25
+ class ParallelStyle:
26
+ pass
27
+
28
+ from fla.ops.utils import logsumexp_fwd
29
+ from fla.ops.utils.op import exp
30
+ from fla.utils import input_guard, is_amd
31
+
32
+ try:
33
+ from torch.distributed.tensor import DTensor
34
+ except (ImportError, AttributeError):
35
+ DTensor = None
36
+
37
+ # The hard limit of TRITON_MAX_TENSOR_NUMEL is 1048576
38
+ # https://github.com/triton-lang/triton/blob/ba42a5c68fd0505f8c42f4202d53be0f8d9a5fe0/python/triton/language/core.py#L19
39
+ # However, setting limit as 65536 as in LayerNorm tutorial is faster because of less register spilling
40
+ # The optimal maximum block size depends on your hardware, your kernel, and your dtype
41
+ MAX_FUSED_SIZE = 65536 // 2
42
+ STATIC_WARPS = 32 if not is_amd else 16
43
+
44
+
45
+ @triton.jit
46
+ def cross_entropy_kernel(
47
+ logits,
48
+ lse,
49
+ target,
50
+ loss,
51
+ total,
52
+ ignore_index,
53
+ label_smoothing: tl.constexpr,
54
+ logit_scale: tl.constexpr,
55
+ reduction: tl.constexpr,
56
+ V: tl.constexpr,
57
+ BV: tl.constexpr,
58
+ ):
59
+ """
60
+ This kernel computes both cross entropy loss and the gradient of the input.
61
+ We only consider hard label + mean reduction for now.
62
+ Please refer to https://pytorch.org/docs/stable/generated/torch.nn.CrossEntropyLoss.html for the math.
63
+
64
+ Args:
65
+ logits:
66
+ Pointer to logits tensor.
67
+ lse:
68
+ Pointer to logsumexp tensor.
69
+ target: Pointer to target tensor.
70
+ loss:
71
+ Pointer to tensor to store the loss.
72
+ V (int):
73
+ The number of columns in the input tensor.
74
+ total (int):
75
+ The number of non-ignored classes.
76
+ ignore_index (int):
77
+ The index to ignore in the target.
78
+ label_smoothing (float):
79
+ The amount of smoothing when computing the loss, where 0.0 means no smoothing.
80
+ reduction (str):
81
+ The string for the reduction to apply
82
+ BV (int):
83
+ The block size for vocab.
84
+ """
85
+
86
+ # https://github.com/triton-lang/triton/issues/1058
87
+ # If B*T*V is too large, i_n * stride will overflow out of int32, so we convert to int64
88
+ i_n = tl.program_id(0).to(tl.int64)
89
+ NV = tl.cdiv(V, BV)
90
+
91
+ # 1. Load target first because if the target is ignore_index, we can return right away
92
+ b_y = tl.load(target + i_n)
93
+
94
+ # 2. locate the start index
95
+ logits += i_n * V
96
+
97
+ if b_y == ignore_index:
98
+ # set all x as 0
99
+ for i in range(0, V, BV):
100
+ o_v = i + tl.arange(0, BV)
101
+ tl.store(logits + o_v, 0.0, mask=o_v < V)
102
+ return
103
+
104
+ # Online softmax: 2 loads + 1 store (compared with 3 loads + 1 store for the safe softmax)
105
+ # Refer to Algorithm 3 in the paper: https://arxiv.org/pdf/1805.02867
106
+
107
+ # 3. [Online softmax] first pass: compute logsumexp
108
+ # we did this in anouter kernel
109
+ b_l = tl.load(logits + b_y) * logit_scale
110
+ b_lse = tl.load(lse + i_n)
111
+
112
+ # 4. Calculate the loss
113
+ # loss = lse - logits_l
114
+ b_loss = b_lse - b_l
115
+
116
+ # Label smoothing is a general case of normal cross entropy
117
+ # See the full derivation at https://github.com/linkedin/Liger-Kernel/pull/198#issue-2503665310
118
+ b_z = 0.0
119
+ eps = label_smoothing / V
120
+
121
+ # We need tl.debug_barrier() as mentioned in
122
+ # https://github.com/triton-lang/triton/blob/ba42a5c68fd0505f8c42f4202d53be0f8d9a5fe0/python/triton/ops/cross_entropy.py#L34
123
+ tl.debug_barrier()
124
+
125
+ # 5. [Online Softmax] Second pass: compute gradients
126
+ # For 'mean' reduction, gradients are normalized by number of non-ignored elements
127
+ # dx_y = (softmax(x_y) - 1) / N
128
+ # dx_i = softmax(x_i) / N, i != y
129
+ # For label smoothing:
130
+ # dx_i = (softmax(x_y) - label_smoothing / V) / N, i != y
131
+ # dx_y = (softmax(x_y) - label_smoothing / V - (1 - label_smoothing)) / N
132
+ # = dx_i - (1 - label_smoothing) / N
133
+ for iv in range(0, NV):
134
+ o_v = iv * BV + tl.arange(0, BV)
135
+ b_logits = tl.load(logits + o_v, mask=o_v < V, other=float('-inf')) * logit_scale
136
+ if label_smoothing > 0:
137
+ # scale X beforehand to avoid overflow
138
+ b_z += tl.sum(tl.where(o_v < V, -eps * b_logits, 0.0))
139
+ b_p = (exp(b_logits - b_lse) - eps) * logit_scale
140
+ if reduction == "mean":
141
+ b_p = b_p / total
142
+ tl.store(logits + o_v, b_p, mask=o_v < V)
143
+
144
+ tl.debug_barrier()
145
+
146
+ # Orginal loss = H(q, p), with label smoothing regularization = H(q', p) and (label_smoothing / V) = eps
147
+ # H(q', p) = (1 - label_smoothing) * H(q, p) + label_smoothing * H(u, p)
148
+ # = (1 - label_smoothing) * H(q, p) + eps * sum(logsoftmax(x_i))
149
+ # By using m (global max of xi) and d (sum of e^(xi-m)), we can simplify as:
150
+ # = (1 - label_smoothing) * H(q, p) + (-sum(x_i * eps) + label_smoothing * (m + logd))
151
+ # Refer to H(q', p) in section 7 of the paper:
152
+ # https://arxiv.org/pdf/1512.00567
153
+ # pytorch:
154
+ # https://github.com/pytorch/pytorch/blob/2981534f54d49fa3a9755c9b0855e7929c2527f0/aten/src/ATen/native/LossNLL.cpp#L516
155
+ # See full derivation at https://github.com/linkedin/Liger-Kernel/pull/198#issuecomment-2333753087
156
+ if label_smoothing > 0:
157
+ b_loss = b_loss * (1 - label_smoothing) + (b_z + label_smoothing * b_lse)
158
+
159
+ # 6. Specially handle the i==y case where `dx_y = (softmax(x_y) - (1 - label_smoothing) / N`
160
+ b_l = tl.load(logits + b_y)
161
+
162
+ # Normalize the loss by the number of non-ignored elements if reduction is "mean"
163
+ if reduction == 'mean':
164
+ b_loss = b_loss / total
165
+ b_l += (label_smoothing - 1) / total * logit_scale
166
+ else:
167
+ b_l += (label_smoothing - 1) * logit_scale
168
+
169
+ tl.store(loss + i_n, b_loss)
170
+ tl.store(logits + b_y, b_l)
171
+
172
+
173
+ @triton.jit
174
+ def elementwise_mul_kernel(
175
+ x,
176
+ g,
177
+ N: tl.constexpr,
178
+ B: tl.constexpr,
179
+ ):
180
+ """
181
+ This function multiplies each element of the tensor pointed by x with the value pointed by g.
182
+ The multiplication is performed in-place on the tensor pointed by x.
183
+
184
+ Parameters:
185
+ x:
186
+ Pointer to the input tensor.
187
+ g:
188
+ Pointer to the gradient output value.
189
+ N (int):
190
+ The number of columns in the input tensor.
191
+ B (int):
192
+ The block size for Triton operations.
193
+ """
194
+
195
+ # Get the program ID and convert it to int64 to avoid overflow
196
+ i_x = tl.program_id(0).to(tl.int64)
197
+ o_x = i_x * B + tl.arange(0, B)
198
+
199
+ # Load the gradient output value
200
+ b_g = tl.load(g)
201
+ b_x = tl.load(x + o_x, mask=o_x < N)
202
+ tl.store(x + o_x, b_x * b_g, mask=o_x < N)
203
+
204
+
205
+ def fused_linear_cross_entropy_forward(
206
+ x: torch.Tensor,
207
+ target: torch.LongTensor,
208
+ weight: torch.Tensor,
209
+ bias: torch.Tensor = None,
210
+ ignore_index: int = -100,
211
+ label_smoothing: float = 0.0,
212
+ logit_scale: float = 1.0,
213
+ num_chunks: int = 8,
214
+ reduction: str = "mean",
215
+ use_l2warp: bool = False,
216
+ l2_penalty_factor: float = 1e-4,
217
+ ):
218
+ device = x.device
219
+ # inputs have shape: [N, H]
220
+ # materialized activations will have shape: [N, V]
221
+ # the increase in memory = [N, V]
222
+ # reduction can be achieved by partitioning the number of tokens N into smaller chunks.
223
+
224
+ # ideally, we would like to achieve the same memory consumption as [N, H],
225
+ # so the expected chunk size should be:
226
+ # NC = ceil(V / H)
227
+ # C = ceil(N / NC)
228
+ # for ex: N = 4096*4, V = 32000, H = 4096 ==> NC = 8, C = ceil(N / NC) = 2048
229
+ N, H, V = *x.shape, weight.shape[0]
230
+ BV = min(MAX_FUSED_SIZE, triton.next_power_of_2(V))
231
+ # TODO: in real cases, we may need to limit the number of chunks NC to
232
+ # ensure the precisions of accumulated gradients
233
+ NC = min(num_chunks, triton.cdiv(V, H))
234
+ C = triton.next_power_of_2(triton.cdiv(N, NC))
235
+ NC = triton.cdiv(N, C)
236
+
237
+ # [N, H]
238
+ dx = torch.zeros_like(x, device=device)
239
+ # [V, H]
240
+ dw = torch.zeros_like(weight, device=device, dtype=torch.float) if weight is not None else None
241
+ # [V]
242
+ db = torch.zeros_like(bias, device=device, dtype=torch.float) if bias is not None else None
243
+ # [N]
244
+ loss = torch.zeros(N, device=device, dtype=torch.float)
245
+
246
+ total = target.ne(ignore_index).sum().item()
247
+
248
+ for ic in range(NC):
249
+ start, end = ic * C, min((ic + 1) * C, N)
250
+ # [C, N]
251
+ c_x = x[start:end]
252
+ # when doing matmul, use the original precision
253
+ # [C, V]
254
+ c_logits = F.linear(c_x, weight, bias)
255
+ c_target = target[start:end]
256
+ # [C]
257
+ # keep lse in fp32 to maintain precision
258
+ c_lse = logsumexp_fwd(c_logits, scale=logit_scale, dtype=torch.float)
259
+
260
+ # unreduced loss
261
+ c_loss = loss[start:end]
262
+ if use_l2warp:
263
+ c_maxx, c_ids = torch.max(c_logits, -1, keepdim=True)
264
+
265
+ # Here we calculate the gradient of c_logits in place so we can save memory.
266
+ cross_entropy_kernel[(c_logits.shape[0],)](
267
+ logits=c_logits,
268
+ lse=c_lse,
269
+ target=c_target,
270
+ loss=c_loss,
271
+ total=total,
272
+ ignore_index=ignore_index,
273
+ label_smoothing=label_smoothing,
274
+ logit_scale=logit_scale,
275
+ reduction=reduction,
276
+ V=V,
277
+ BV=BV,
278
+ num_warps=STATIC_WARPS,
279
+ )
280
+ if use_l2warp:
281
+ # a. Calculate the L2 gradient w.r.t logits (g_logits_l2)
282
+ g_logits_l2 = torch.zeros_like(c_logits)
283
+
284
+ # Normalize factor by B*T, which is the 'total' variable here
285
+ l2_factor = l2_penalty_factor / total if reduction == 'mean' else l2_penalty_factor
286
+ penalty_grad = c_maxx * l2_factor
287
+ g_logits_l2.scatter_(-1, c_ids, penalty_grad)
288
+
289
+ # b. Backpropagate g_logits_l2 to get its effect on dx, dw, db
290
+ # and add it to the main gradients.
291
+ # Total_dx = CE_dx + L2_dx
292
+ # Total_dw = CE_dw + L2_dw
293
+ # Total_db = CE_db + L2_db
294
+ if weight is not None:
295
+ dw.add_(g_logits_l2.t() @ c_x)
296
+ if bias is not None:
297
+ db.add_(g_logits_l2.sum(0))
298
+ # The dx contribution must be added to the final dx calculation
299
+ dx_l2_contribution = torch.mm(g_logits_l2, weight)
300
+ else:
301
+ dx_l2_contribution = 0.0
302
+
303
+ # gradient of logits is computed in-place by the above triton kernel and is of shape: C x V
304
+ # thus dx should be of shape: C x H
305
+ dx[start:end] = torch.mm(c_logits, weight) + dx_l2_contribution
306
+
307
+ # keep dw in fp32 to maintain precision
308
+ if weight is not None:
309
+ dw += c_logits.t() @ c_x
310
+
311
+ if bias is not None:
312
+ torch.add(input=db, other=c_logits.sum(0), out=db)
313
+
314
+ loss = loss.sum()
315
+ if dw is not None:
316
+ dw = dw.to(weight)
317
+ if db is not None:
318
+ db = db.to(bias)
319
+ return loss, dx, dw, db
320
+
321
+
322
+ def fused_linear_cross_entropy_backward(
323
+ do: torch.Tensor,
324
+ dx: torch.Tensor,
325
+ dw: torch.Tensor,
326
+ db: torch.Tensor,
327
+ ):
328
+ # If cross entropy is the last layer, do is 1.0. Skip the mul to save time
329
+ if torch.ne(do, torch.tensor(1.0, device=do.device)):
330
+ # We use a Triton kernel instead of a PyTorch operation because modifying inputs in-place
331
+ # for gradient storage and backward multiple times causes anomalies with PyTorch but not with Triton.
332
+ N, H = dx.shape
333
+ B = min(MAX_FUSED_SIZE, triton.next_power_of_2(H))
334
+
335
+ elementwise_mul_kernel[(triton.cdiv(N * H, B),)](
336
+ x=dx,
337
+ g=do,
338
+ N=N*H,
339
+ B=B,
340
+ num_warps=STATIC_WARPS,
341
+ )
342
+
343
+ # handle dw
344
+ if dw is not None:
345
+ V, H = dw.shape
346
+ elementwise_mul_kernel[(triton.cdiv(V * H, B),)](
347
+ x=dw,
348
+ g=do,
349
+ N=V*H,
350
+ B=B,
351
+ num_warps=STATIC_WARPS,
352
+ )
353
+
354
+ if db is not None:
355
+ V = db.shape[0]
356
+ elementwise_mul_kernel[(triton.cdiv(V, B),)](
357
+ x=db,
358
+ g=do,
359
+ N=V,
360
+ B=B,
361
+ num_warps=STATIC_WARPS,
362
+ )
363
+ return dx, dw, db
364
+
365
+
366
+ class FusedLinearCrossEntropyFunction(torch.autograd.Function):
367
+
368
+ @staticmethod
369
+ @input_guard
370
+ def forward(
371
+ ctx,
372
+ x: torch.Tensor,
373
+ target: torch.LongTensor,
374
+ weight: torch.Tensor,
375
+ bias: torch.Tensor = None,
376
+ ignore_index: int = -100,
377
+ label_smoothing: float = 0.0,
378
+ logit_scale: float = 1.0,
379
+ num_chunks: int = 8,
380
+ reduction: str = "mean",
381
+ use_l2warp: bool = False,
382
+ l2_penalty_factor: float = 1e-4,
383
+ ):
384
+ """
385
+ Fusing the last linear layer with cross-entropy loss
386
+ Reference: https://github.com/mgmalek/efficient_cross_entropy
387
+
388
+ Handle the forward and backward pass of the final linear layer via cross-entropy loss by avoiding
389
+ the materialization of the large logits tensor. Since Cross Entropy Loss is the last layer, we can
390
+ compute the gradient at the forward pass. By doing so, we don't have to store the x and target
391
+ for the backward pass.
392
+
393
+ x (torch.Tensor): [batch_size * seq_len, hidden_size]
394
+ target (torch.LongTensor): [batch_size * seq_len]
395
+ where each value is in [0, vocab_size).
396
+ weight (torch.Tensor): [vocab_size, hidden_size]
397
+ where `vocab_size` is the number of classes.
398
+ bias (Optional[torch.Tensor]): [vocab_size]
399
+ where `vocab_size` is the number of classes.
400
+ ignore_index:
401
+ the index to ignore in the target.
402
+ label_smoothing:
403
+ the amount of smoothing when computing the loss, where 0.0 means no smoothing.
404
+ logit_scale: float = 1.0,
405
+ A scaling factor applied to the logits. Default: 1.0
406
+ num_chunks: int
407
+ The number of chunks to split the input tensor into for processing.
408
+ This can help optimize memory usage and computation speed.
409
+ Default: 8
410
+ reduction:
411
+ Specifies the reduction to apply to the output: 'mean' | 'sum'.
412
+ 'mean': the weighted mean of the output is taken,
413
+ 'sum': the output will be summed.
414
+ Default: 'mean'.
415
+ use_l2warp: bool = False,
416
+ Whether to use L2 regularization on the logits to prevent overconfidence.
417
+ Default: False
418
+ l2_penalty_factor: float = 1e-4,
419
+ """
420
+ loss, dx, dw, db = fused_linear_cross_entropy_forward(
421
+ x,
422
+ target,
423
+ weight,
424
+ bias,
425
+ ignore_index,
426
+ label_smoothing,
427
+ logit_scale,
428
+ num_chunks,
429
+ reduction,
430
+ use_l2warp,
431
+ l2_penalty_factor,
432
+ )
433
+ # downcast to dtype and store for backward
434
+ ctx.save_for_backward(
435
+ dx.detach(),
436
+ dw.detach() if weight is not None else None,
437
+ db.detach() if bias is not None else None,
438
+ )
439
+ return loss
440
+
441
+ @staticmethod
442
+ @input_guard
443
+ def backward(ctx, do):
444
+ dx, dw, db = ctx.saved_tensors
445
+ dx, dw, db = fused_linear_cross_entropy_backward(do, dx, dw, db)
446
+ return dx, None, dw, db, None, None, None, None, None, None, None
447
+
448
+
449
+ def fused_linear_cross_entropy_loss(
450
+ x: torch.Tensor,
451
+ target: torch.LongTensor,
452
+ weight: torch.Tensor,
453
+ bias: torch.Tensor = None,
454
+ ignore_index: int = -100,
455
+ label_smoothing: float = 0.0,
456
+ logit_scale: float = 1.0,
457
+ num_chunks: int = 8,
458
+ reduction: str = "mean",
459
+ use_l2warp: bool = False,
460
+ l2_penalty_factor: float = 1e-4,
461
+ ) -> tuple[torch.Tensor, torch.Tensor]:
462
+ """
463
+ Args:
464
+ x (torch.Tensor): [batch_size * seq_len, hidden_size]
465
+ target (torch.LongTensor): [batch_size * seq_len]
466
+ where each value is in [0, vocab_size).
467
+ weight (torch.Tensor): [vocab_size, hidden_size]
468
+ where `vocab_size` is the number of classes.
469
+ bias (Optional[torch.Tensor]): [vocab_size]
470
+ where `vocab_size` is the number of classes.
471
+ ignore_index: int.
472
+ If target == ignore_index, the loss is set to 0.0.
473
+ label_smoothing: float
474
+ logit_scale: float
475
+ A scaling factor applied to the logits. Default: 1.0
476
+ num_chunks: int
477
+ The number of chunks to split the input tensor into for processing.
478
+ This can help optimize memory usage and computation speed.
479
+ Default: 8
480
+ reduction:
481
+ Specifies the reduction to apply to the output: 'mean' | 'sum'.
482
+ 'mean': the weighted mean of the output is taken,
483
+ 'sum': the output will be summed.
484
+ Default: 'mean'.
485
+ Returns:
486
+ losses: [batch,], float
487
+ """
488
+ return FusedLinearCrossEntropyFunction.apply(
489
+ x,
490
+ target,
491
+ weight,
492
+ bias,
493
+ ignore_index,
494
+ label_smoothing,
495
+ logit_scale,
496
+ num_chunks,
497
+ reduction,
498
+ use_l2warp,
499
+ l2_penalty_factor,
500
+ )
501
+
502
+
503
+ class FusedLinearCrossEntropyLoss(nn.Module):
504
+
505
+ def __init__(
506
+ self,
507
+ ignore_index: int = -100,
508
+ label_smoothing: float = 0.0,
509
+ logit_scale: float = 1.0,
510
+ num_chunks: int = 8,
511
+ reduction: str = "mean",
512
+ use_l2warp: bool = False,
513
+ l2_penalty_factor: float = 1e-4,
514
+ ):
515
+ """
516
+ Args:
517
+ ignore_index: int.
518
+ If target == ignore_index, the loss is set to 0.0.
519
+ label_smoothing: float
520
+ logit_scale: float
521
+ A scaling factor applied to the logits. Default: 1.0
522
+ num_chunks: int
523
+ The number of chunks to split the input tensor into for processing.
524
+ This can help optimize memory usage and computation speed.
525
+ Default: 8
526
+ reduction:
527
+ Specifies the reduction to apply to the output: 'mean' | 'sum'.
528
+ 'mean': the weighted mean of the output is taken,
529
+ 'sum': the output will be summed.
530
+ Default: 'mean'.
531
+ """
532
+ super().__init__()
533
+
534
+ assert reduction in ["mean", "sum"], f"reduction: {reduction} is not supported"
535
+
536
+ self.ignore_index = ignore_index
537
+ self.label_smoothing = label_smoothing
538
+ self.logit_scale = logit_scale
539
+ self.num_chunks = num_chunks
540
+ self.reduction = reduction
541
+ self.use_l2warp = use_l2warp
542
+ self.l2_penalty_factor = l2_penalty_factor
543
+
544
+ @torch.compiler.disable
545
+ def forward(
546
+ self,
547
+ x: torch.Tensor,
548
+ target: torch.LongTensor,
549
+ weight: torch.Tensor,
550
+ bias: torch.Tensor | None = None,
551
+ ):
552
+ """
553
+ Args:
554
+ x (torch.Tensor): [batch_size, seq_len, hidden_size]
555
+ target (torch.LongTensor): [batch_size, seq_len]
556
+ where each value is in [0, V).
557
+ weight (torch.Tensor): [vocab_size, hidden_size]
558
+ where `vocab_size` is the number of classes.
559
+ bias (Optional[torch.Tensor]): [vocab_size]
560
+ where `vocab_size` is the number of classes.
561
+ Returns:
562
+ loss
563
+ """
564
+ loss = fused_linear_cross_entropy_loss(
565
+ x.view(-1, x.shape[-1]),
566
+ target.view(-1),
567
+ weight=weight,
568
+ bias=bias,
569
+ ignore_index=self.ignore_index,
570
+ label_smoothing=self.label_smoothing,
571
+ logit_scale=self.logit_scale,
572
+ num_chunks=self.num_chunks,
573
+ reduction=self.reduction,
574
+ use_l2warp=self.use_l2warp,
575
+ l2_penalty_factor=self.l2_penalty_factor,
576
+ )
577
+ return loss
578
+
579
+
580
+ class LinearLossParallel(ParallelStyle):
581
+ def __init__(
582
+ self,
583
+ *,
584
+ sequence_dim: int = 1,
585
+ use_local_output: bool = False,
586
+ ):
587
+ super().__init__()
588
+
589
+ self.sequence_sharding = (Shard(sequence_dim),)
590
+ self.use_local_output = use_local_output
591
+
592
+ @staticmethod
593
+ def _prepare_input_fn(sequence_sharding, mod, inputs, device_mesh):
594
+ x, target, weight, bias = inputs
595
+
596
+ if not isinstance(x, DTensor):
597
+ # assume the input passed in already sharded on the sequence dim and create the DTensor
598
+ x = DTensor.from_local(x, device_mesh, sequence_sharding)
599
+ if x.placements != sequence_sharding:
600
+ x = x.redistribute(placements=sequence_sharding, async_op=True)
601
+ if not isinstance(target, DTensor):
602
+ target = DTensor.from_local(target, device_mesh, [Replicate()])
603
+ if target.placements != sequence_sharding:
604
+ target = target.redistribute(placements=sequence_sharding, async_op=True)
605
+
606
+ if not isinstance(weight, DTensor):
607
+ weight = DTensor.from_local(weight, device_mesh, [Replicate()])
608
+ if weight.placements != [Replicate()]:
609
+ # we replicate the weight/bias in FLCE
610
+ weight = weight.redistribute(placements=[Replicate()], async_op=True)
611
+
612
+ if bias is not None and not isinstance(bias, DTensor):
613
+ bias = DTensor.from_local(bias, device_mesh, [Replicate()])
614
+ if bias is not None and bias.placements != [Replicate()]:
615
+ bias = bias.redistribute(placements=[Replicate()], async_op=True)
616
+
617
+ return x.to_local(), target.to_local(), weight.to_local(), bias.to_local() if bias is not None else bias
618
+
619
+ @staticmethod
620
+ def _prepare_output_fn(use_local_output, mod, outputs, device_mesh):
621
+ return outputs.to_local() if use_local_output else outputs
622
+
623
+ def _apply(self, module: nn.Module, device_mesh: DeviceMesh) -> nn.Module:
624
+ return distribute_module(
625
+ module,
626
+ device_mesh,
627
+ partition_fn=None,
628
+ input_fn=partial(self._prepare_input_fn, self.sequence_sharding),
629
+ output_fn=partial(self._prepare_output_fn, self.use_local_output),
630
+ )
code/flash-linear-attention/fla/modules/fused_norm_gate.py ADDED
@@ -0,0 +1,1245 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ import torch.nn.functional as F
10
+ import triton
11
+ import triton.language as tl
12
+
13
+ from fla.utils import autotune_cache_kwargs, get_multiprocessor_count, input_guard
14
+
15
+
16
+ @triton.heuristics({
17
+ 'STORE_RESIDUAL_OUT': lambda args: args['residual_out'] is not None,
18
+ 'HAS_RESIDUAL': lambda args: args['residual'] is not None,
19
+ 'HAS_WEIGHT': lambda args: args['w'] is not None,
20
+ 'HAS_BIAS': lambda args: args['b'] is not None,
21
+ })
22
+ @triton.autotune(
23
+ configs=[
24
+ triton.Config({'BT': BT}, num_warps=num_warps)
25
+ for BT in [16, 32, 64]
26
+ for num_warps in [4, 8, 16]
27
+ ],
28
+ key=['D', 'NB', 'IS_RMS_NORM', 'STORE_RESIDUAL_OUT', 'HAS_RESIDUAL', 'HAS_WEIGHT'],
29
+ **autotune_cache_kwargs,
30
+ )
31
+ @triton.jit
32
+ def layer_norm_gated_fwd_kernel(
33
+ x, # pointer to the input
34
+ g, # pointer to the gate
35
+ y, # pointer to the output
36
+ w, # pointer to the weights
37
+ b, # pointer to the biases
38
+ residual, # pointer to the residual
39
+ residual_out, # pointer to the residual
40
+ mean, # pointer to the mean
41
+ rstd, # pointer to the 1/std
42
+ eps, # epsilon to avoid division by zero
43
+ T, # number of rows in x
44
+ D: tl.constexpr, # number of columns in x
45
+ BT: tl.constexpr,
46
+ BD: tl.constexpr,
47
+ NB: tl.constexpr,
48
+ ACTIVATION: tl.constexpr,
49
+ IS_RMS_NORM: tl.constexpr,
50
+ STORE_RESIDUAL_OUT: tl.constexpr,
51
+ HAS_RESIDUAL: tl.constexpr,
52
+ HAS_WEIGHT: tl.constexpr,
53
+ HAS_BIAS: tl.constexpr,
54
+ ):
55
+ i_t = tl.program_id(0)
56
+
57
+ o_d = tl.arange(0, BD)
58
+ m_d = o_d < D
59
+
60
+ p_x = tl.make_block_ptr(x, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
61
+ b_x = tl.load(p_x, boundary_check=(0, 1)).to(tl.float32)
62
+ if HAS_RESIDUAL:
63
+ p_res = tl.make_block_ptr(residual, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
64
+ b_x += tl.load(p_res, boundary_check=(0, 1)).to(tl.float32)
65
+ if STORE_RESIDUAL_OUT:
66
+ p_res_out = tl.make_block_ptr(residual_out, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
67
+ tl.store(p_res_out, b_x.to(p_res_out.dtype.element_ty), boundary_check=(0, 1))
68
+ if not IS_RMS_NORM:
69
+ b_mean = tl.sum(b_x, axis=1) / D
70
+ p_mean = tl.make_block_ptr(mean, (T,), (1,), (i_t * BT,), (BT,), (0,))
71
+ tl.store(p_mean, b_mean.to(p_mean.dtype.element_ty), boundary_check=(0,))
72
+ b_xbar = tl.where(m_d[None, :], b_x - b_mean[:, None], 0.0)
73
+ b_var = tl.sum(b_xbar * b_xbar, axis=1) / D
74
+ else:
75
+ b_xbar = tl.where(m_d[None, :], b_x, 0.0)
76
+ b_var = tl.sum(b_xbar * b_xbar, axis=1) / D
77
+ b_rstd = 1 / tl.sqrt(b_var + eps)
78
+
79
+ p_rstd = tl.make_block_ptr(rstd, (T,), (1,), (i_t * BT,), (BT,), (0,))
80
+ tl.store(p_rstd, b_rstd.to(p_rstd.dtype.element_ty), boundary_check=(0,))
81
+
82
+ if HAS_WEIGHT:
83
+ b_w = tl.load(w + o_d, mask=m_d).to(tl.float32)
84
+ if HAS_BIAS:
85
+ b_b = tl.load(b + o_d, mask=m_d).to(tl.float32)
86
+ b_x_hat = (b_x - b_mean[:, None]) * b_rstd[:, None] if not IS_RMS_NORM else b_x * b_rstd[:, None]
87
+ b_y = b_x_hat * b_w[None, :] if HAS_WEIGHT else b_x_hat
88
+ if HAS_BIAS:
89
+ b_y = b_y + b_b[None, :]
90
+
91
+ # swish/sigmoid output gate
92
+ p_g = tl.make_block_ptr(g, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
93
+ b_g = tl.load(p_g, boundary_check=(0, 1)).to(tl.float32)
94
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
95
+ b_y = b_y * b_g * tl.sigmoid(b_g)
96
+ elif ACTIVATION == 'sigmoid':
97
+ b_y = b_y * tl.sigmoid(b_g)
98
+
99
+ # Write output
100
+ p_y = tl.make_block_ptr(y, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
101
+ tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1))
102
+
103
+
104
+ @triton.heuristics({
105
+ 'STORE_RESIDUAL_OUT': lambda args: args['residual_out'] is not None,
106
+ 'HAS_RESIDUAL': lambda args: args['residual'] is not None,
107
+ 'HAS_WEIGHT': lambda args: args['w'] is not None,
108
+ 'HAS_BIAS': lambda args: args['b'] is not None,
109
+ })
110
+ @triton.autotune(
111
+ configs=[
112
+ triton.Config({}, num_warps=num_warps)
113
+ for num_warps in [2, 4, 8, 16]
114
+ ],
115
+ key=['D', 'IS_RMS_NORM', 'STORE_RESIDUAL_OUT', 'HAS_RESIDUAL', 'HAS_WEIGHT'],
116
+ **autotune_cache_kwargs,
117
+ )
118
+ @triton.jit
119
+ def layer_norm_gated_fwd_kernel1(
120
+ x, # pointer to the input
121
+ g, # pointer to the gate
122
+ y, # pointer to the output
123
+ w, # pointer to the weights
124
+ b, # pointer to the biases
125
+ residual, # pointer to the residual
126
+ residual_out, # pointer to the residual
127
+ mean, # pointer to the mean
128
+ rstd, # pointer to the 1/std
129
+ eps, # epsilon to avoid division by zero
130
+ D: tl.constexpr, # number of columns in x
131
+ BD: tl.constexpr,
132
+ ACTIVATION: tl.constexpr,
133
+ IS_RMS_NORM: tl.constexpr,
134
+ STORE_RESIDUAL_OUT: tl.constexpr,
135
+ HAS_RESIDUAL: tl.constexpr,
136
+ HAS_WEIGHT: tl.constexpr,
137
+ HAS_BIAS: tl.constexpr,
138
+ ):
139
+ i_t = tl.program_id(0)
140
+ x += i_t * D
141
+ y += i_t * D
142
+ g += i_t * D
143
+ if HAS_RESIDUAL:
144
+ residual += i_t * D
145
+ if STORE_RESIDUAL_OUT:
146
+ residual_out += i_t * D
147
+
148
+ o_d = tl.arange(0, BD)
149
+ m_d = o_d < D
150
+ b_x = tl.load(x + o_d, mask=m_d, other=0.0).to(tl.float32)
151
+ if HAS_RESIDUAL:
152
+ b_x += tl.load(residual + o_d, mask=m_d, other=0.0).to(tl.float32)
153
+ if STORE_RESIDUAL_OUT:
154
+ tl.store(residual_out + o_d, b_x, mask=m_d)
155
+ if not IS_RMS_NORM:
156
+ b_mean = tl.sum(b_x, axis=0) / D
157
+ tl.store(mean + i_t, b_mean)
158
+ b_xbar = tl.where(m_d, b_x - b_mean, 0.0)
159
+ b_var = tl.sum(b_xbar * b_xbar, axis=0) / D
160
+ else:
161
+ b_xbar = tl.where(m_d, b_x, 0.0)
162
+ b_var = tl.sum(b_xbar * b_xbar, axis=0) / D
163
+ b_rstd = 1 / tl.sqrt(b_var + eps)
164
+ tl.store(rstd + i_t, b_rstd)
165
+
166
+ if HAS_WEIGHT:
167
+ b_w = tl.load(w + o_d, mask=m_d).to(tl.float32)
168
+ if HAS_BIAS:
169
+ b_b = tl.load(b + o_d, mask=m_d).to(tl.float32)
170
+ b_x_hat = (b_x - b_mean) * b_rstd if not IS_RMS_NORM else b_x * b_rstd
171
+ b_y = b_x_hat * b_w if HAS_WEIGHT else b_x_hat
172
+ if HAS_BIAS:
173
+ b_y = b_y + b_b
174
+
175
+ # swish/sigmoid output gate
176
+ b_g = tl.load(g + o_d, mask=m_d, other=0.0).to(tl.float32)
177
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
178
+ b_y = b_y * b_g * tl.sigmoid(b_g)
179
+ elif ACTIVATION == 'sigmoid':
180
+ b_y = b_y * tl.sigmoid(b_g)
181
+
182
+ # Write output
183
+ tl.store(y + o_d, b_y, mask=m_d)
184
+
185
+
186
+ @triton.heuristics({
187
+ 'HAS_DRESIDUAL': lambda args: args['dresidual'] is not None,
188
+ 'HAS_WEIGHT': lambda args: args['w'] is not None,
189
+ 'HAS_BIAS': lambda args: args['b'] is not None,
190
+ 'RECOMPUTE_OUTPUT': lambda args: args['y'] is not None,
191
+ })
192
+ @triton.autotune(
193
+ configs=[
194
+ triton.Config({'BT': BT}, num_warps=num_warps)
195
+ for BT in [16, 32, 64]
196
+ for num_warps in [4, 8, 16]
197
+ ],
198
+ key=['D', 'NB', 'IS_RMS_NORM', 'HAS_DRESIDUAL', 'HAS_WEIGHT'],
199
+ **autotune_cache_kwargs,
200
+ )
201
+ @triton.jit
202
+ def layer_norm_gated_bwd_kernel(
203
+ x, # pointer to the input
204
+ g, # pointer to the gate
205
+ w, # pointer to the weights
206
+ b, # pointer to the biases
207
+ y, # pointer to the output to be recomputed
208
+ dy, # pointer to the output gradient
209
+ dx, # pointer to the input gradient
210
+ dg, # pointer to the gate gradient
211
+ dw, # pointer to the partial sum of weights gradient
212
+ db, # pointer to the partial sum of biases gradient
213
+ dresidual,
214
+ dresidual_in,
215
+ mean,
216
+ rstd,
217
+ T,
218
+ BS,
219
+ D: tl.constexpr,
220
+ BT: tl.constexpr,
221
+ BD: tl.constexpr,
222
+ NB: tl.constexpr,
223
+ ACTIVATION: tl.constexpr,
224
+ IS_RMS_NORM: tl.constexpr,
225
+ STORE_DRESIDUAL: tl.constexpr,
226
+ HAS_DRESIDUAL: tl.constexpr,
227
+ HAS_WEIGHT: tl.constexpr,
228
+ HAS_BIAS: tl.constexpr,
229
+ RECOMPUTE_OUTPUT: tl.constexpr,
230
+ ):
231
+ i_s = tl.program_id(0)
232
+ o_d = tl.arange(0, BD)
233
+ m_d = o_d < D
234
+ if HAS_WEIGHT:
235
+ b_w = tl.load(w + o_d, mask=m_d).to(tl.float32)
236
+ b_dw = tl.zeros((BT, BD), dtype=tl.float32)
237
+ if HAS_BIAS:
238
+ b_b = tl.load(b + o_d, mask=m_d, other=0.0).to(tl.float32)
239
+ b_db = tl.zeros((BT, BD), dtype=tl.float32)
240
+
241
+ T = min(i_s * BS + BS, T)
242
+ for i_t in range(i_s * BS, T, BT):
243
+ p_x = tl.make_block_ptr(x, (T, D), (D, 1), (i_t, 0), (BT, BD), (1, 0))
244
+ p_g = tl.make_block_ptr(g, (T, D), (D, 1), (i_t, 0), (BT, BD), (1, 0))
245
+ p_dy = tl.make_block_ptr(dy, (T, D), (D, 1), (i_t, 0), (BT, BD), (1, 0))
246
+ p_dx = tl.make_block_ptr(dx, (T, D), (D, 1), (i_t, 0), (BT, BD), (1, 0))
247
+ p_dg = tl.make_block_ptr(dg, (T, D), (D, 1), (i_t, 0), (BT, BD), (1, 0))
248
+ # [BT, BD]
249
+ b_x = tl.load(p_x, boundary_check=(0, 1)).to(tl.float32)
250
+ b_g = tl.load(p_g, boundary_check=(0, 1)).to(tl.float32)
251
+ b_dy = tl.load(p_dy, boundary_check=(0, 1)).to(tl.float32)
252
+
253
+ if not IS_RMS_NORM:
254
+ p_mean = tl.make_block_ptr(mean, (T,), (1,), (i_t,), (BT,), (0,))
255
+ b_mean = tl.load(p_mean, boundary_check=(0,))
256
+ p_rstd = tl.make_block_ptr(rstd, (T,), (1,), (i_t,), (BT,), (0,))
257
+ b_rstd = tl.load(p_rstd, boundary_check=(0,))
258
+ # Compute dx
259
+ b_xhat = (b_x - b_mean[:, None]) * b_rstd[:, None] if not IS_RMS_NORM else b_x * b_rstd[:, None]
260
+ b_xhat = tl.where(m_d[None, :], b_xhat, 0.0)
261
+
262
+ b_y = b_xhat * b_w[None, :] if HAS_WEIGHT else b_xhat
263
+ if HAS_BIAS:
264
+ b_y = b_y + b_b[None, :]
265
+ if RECOMPUTE_OUTPUT:
266
+ p_y = tl.make_block_ptr(y, (T, D), (D, 1), (i_t, 0), (BT, BD), (1, 0))
267
+ tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1))
268
+
269
+ b_sigmoid_g = tl.sigmoid(b_g)
270
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
271
+ b_dg = b_dy * b_y * (b_sigmoid_g + b_g * b_sigmoid_g * (1 - b_sigmoid_g))
272
+ b_dy = b_dy * b_g * b_sigmoid_g
273
+ elif ACTIVATION == 'sigmoid':
274
+ b_dg = b_dy * b_y * b_sigmoid_g * (1 - b_sigmoid_g)
275
+ b_dy = b_dy * b_sigmoid_g
276
+ b_wdy = b_dy
277
+
278
+ if HAS_WEIGHT or HAS_BIAS:
279
+ m_t = (i_t + tl.arange(0, BT)) < T
280
+ if HAS_WEIGHT:
281
+ b_wdy = b_dy * b_w
282
+ b_dw += tl.where(m_t[:, None], b_dy * b_xhat, 0.0)
283
+ if HAS_BIAS:
284
+ b_db += tl.where(m_t[:, None], b_dy, 0.0)
285
+ if not IS_RMS_NORM:
286
+ b_c1 = tl.sum(b_xhat * b_wdy, axis=1) / D
287
+ b_c2 = tl.sum(b_wdy, axis=1) / D
288
+ b_dx = (b_wdy - (b_xhat * b_c1[:, None] + b_c2[:, None])) * b_rstd[:, None]
289
+ else:
290
+ b_c1 = tl.sum(b_xhat * b_wdy, axis=1) / D
291
+ b_dx = (b_wdy - b_xhat * b_c1[:, None]) * b_rstd[:, None]
292
+ if HAS_DRESIDUAL:
293
+ p_dres = tl.make_block_ptr(dresidual, (T, D), (D, 1), (i_t, 0), (BT, BD), (1, 0))
294
+ b_dres = tl.load(p_dres, boundary_check=(0, 1)).to(tl.float32)
295
+ b_dx += b_dres
296
+ # Write dx
297
+ if STORE_DRESIDUAL:
298
+ p_dres_in = tl.make_block_ptr(dresidual_in, (T, D), (D, 1), (i_t, 0), (BT, BD), (1, 0))
299
+ tl.store(p_dres_in, b_dx.to(p_dres_in.dtype.element_ty), boundary_check=(0, 1))
300
+
301
+ tl.store(p_dx, b_dx.to(p_dx.dtype.element_ty), boundary_check=(0, 1))
302
+ tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), boundary_check=(0, 1))
303
+
304
+ if HAS_WEIGHT:
305
+ tl.store(dw + i_s * D + o_d, tl.sum(b_dw, axis=0), mask=m_d)
306
+ if HAS_BIAS:
307
+ tl.store(db + i_s * D + o_d, tl.sum(b_db, axis=0), mask=m_d)
308
+
309
+
310
+ @triton.heuristics({
311
+ 'HAS_DRESIDUAL': lambda args: args['dresidual'] is not None,
312
+ 'HAS_WEIGHT': lambda args: args['w'] is not None,
313
+ 'HAS_BIAS': lambda args: args['b'] is not None,
314
+ 'RECOMPUTE_OUTPUT': lambda args: args['y'] is not None,
315
+ })
316
+ @triton.autotune(
317
+ configs=[
318
+ triton.Config({}, num_warps=num_warps)
319
+ for num_warps in [2, 4, 8, 16]
320
+ ],
321
+ key=['D', 'IS_RMS_NORM', 'STORE_DRESIDUAL', 'HAS_DRESIDUAL', 'HAS_WEIGHT'],
322
+ **autotune_cache_kwargs,
323
+ )
324
+ @triton.jit
325
+ def layer_norm_gated_bwd_kernel1(
326
+ x, # pointer to the input
327
+ g, # pointer to the gate
328
+ w, # pointer to the weights
329
+ b, # pointer to the biases
330
+ y, # pointer to the output to be recomputed
331
+ dy, # pointer to the output gradient
332
+ dx, # pointer to the input gradient
333
+ dg, # pointer to the gate gradient
334
+ dw, # pointer to the partial sum of weights gradient
335
+ db, # pointer to the partial sum of biases gradient
336
+ dresidual,
337
+ dresidual_in,
338
+ mean,
339
+ rstd,
340
+ T,
341
+ BS,
342
+ D: tl.constexpr,
343
+ BD: tl.constexpr,
344
+ ACTIVATION: tl.constexpr,
345
+ IS_RMS_NORM: tl.constexpr,
346
+ STORE_DRESIDUAL: tl.constexpr,
347
+ HAS_DRESIDUAL: tl.constexpr,
348
+ HAS_WEIGHT: tl.constexpr,
349
+ HAS_BIAS: tl.constexpr,
350
+ RECOMPUTE_OUTPUT: tl.constexpr,
351
+ ):
352
+ i_s = tl.program_id(0)
353
+ o_d = tl.arange(0, BD)
354
+ mask = o_d < D
355
+ x += i_s * BS * D
356
+ g += i_s * BS * D
357
+ if HAS_DRESIDUAL:
358
+ dresidual += i_s * BS * D
359
+ if STORE_DRESIDUAL:
360
+ dresidual_in += i_s * BS * D
361
+ dy += i_s * BS * D
362
+ dx += i_s * BS * D
363
+ dg += i_s * BS * D
364
+ if RECOMPUTE_OUTPUT:
365
+ y += i_s * BS * D
366
+ if HAS_WEIGHT:
367
+ b_w = tl.load(w + o_d, mask=mask).to(tl.float32)
368
+ b_dw = tl.zeros((BD,), dtype=tl.float32)
369
+ if HAS_BIAS:
370
+ b_b = tl.load(b + o_d, mask=mask, other=0.0).to(tl.float32)
371
+ b_db = tl.zeros((BD,), dtype=tl.float32)
372
+
373
+ for i_t in range(i_s * BS, min(i_s * BS + BS, T)):
374
+ # Load data to SRAM
375
+ b_x = tl.load(x + o_d, mask=mask, other=0).to(tl.float32)
376
+ b_g = tl.load(g + o_d, mask=mask, other=0).to(tl.float32)
377
+ b_dy = tl.load(dy + o_d, mask=mask, other=0).to(tl.float32)
378
+
379
+ if not IS_RMS_NORM:
380
+ b_mean = tl.load(mean + i_t)
381
+ b_rstd = tl.load(rstd + i_t)
382
+ # Compute dx
383
+ b_xhat = (b_x - b_mean) * b_rstd if not IS_RMS_NORM else b_x * b_rstd
384
+ b_xhat = tl.where(mask, b_xhat, 0.0)
385
+
386
+ b_y = b_xhat * b_w if HAS_WEIGHT else b_xhat
387
+ if HAS_BIAS:
388
+ b_y = b_y + b_b
389
+ if RECOMPUTE_OUTPUT:
390
+ tl.store(y + o_d, b_y, mask=mask)
391
+
392
+ b_sigmoid_g = tl.sigmoid(b_g)
393
+ if ACTIVATION == 'swish' or ACTIVATION == 'silu':
394
+ b_dg = b_dy * b_y * (b_sigmoid_g + b_g * b_sigmoid_g * (1 - b_sigmoid_g))
395
+ b_dy = b_dy * b_g * b_sigmoid_g
396
+ elif ACTIVATION == 'sigmoid':
397
+ b_dg = b_dy * b_y * b_sigmoid_g * (1 - b_sigmoid_g)
398
+ b_dy = b_dy * b_sigmoid_g
399
+ b_wdy = b_dy
400
+ if HAS_WEIGHT:
401
+ b_wdy = b_dy * b_w
402
+ b_dw += b_dy * b_xhat
403
+ if HAS_BIAS:
404
+ b_db += b_dy
405
+ if not IS_RMS_NORM:
406
+ b_c1 = tl.sum(b_xhat * b_wdy, axis=0) / D
407
+ b_c2 = tl.sum(b_wdy, axis=0) / D
408
+ b_dx = (b_wdy - (b_xhat * b_c1 + b_c2)) * b_rstd
409
+ else:
410
+ b_c1 = tl.sum(b_xhat * b_wdy, axis=0) / D
411
+ b_dx = (b_wdy - b_xhat * b_c1) * b_rstd
412
+ if HAS_DRESIDUAL:
413
+ b_dres = tl.load(dresidual + o_d, mask=mask, other=0).to(tl.float32)
414
+ b_dx += b_dres
415
+ # Write dx
416
+ if STORE_DRESIDUAL:
417
+ tl.store(dresidual_in + o_d, b_dx, mask=mask)
418
+ tl.store(dx + o_d, b_dx, mask=mask)
419
+ tl.store(dg + o_d, b_dg, mask=mask)
420
+
421
+ x += D
422
+ g += D
423
+ if HAS_DRESIDUAL:
424
+ dresidual += D
425
+ if STORE_DRESIDUAL:
426
+ dresidual_in += D
427
+ if RECOMPUTE_OUTPUT:
428
+ y += D
429
+ dy += D
430
+ dx += D
431
+ dg += D
432
+ if HAS_WEIGHT:
433
+ tl.store(dw + i_s * D + o_d, b_dw, mask=mask)
434
+ if HAS_BIAS:
435
+ tl.store(db + i_s * D + o_d, b_db, mask=mask)
436
+
437
+
438
+ def layer_norm_gated_fwd(
439
+ x: torch.Tensor,
440
+ g: torch.Tensor,
441
+ weight: torch.Tensor,
442
+ bias: torch.Tensor,
443
+ activation: str = 'swish',
444
+ eps: float = 1e-5,
445
+ residual: torch.Tensor = None,
446
+ out_dtype: torch.dtype = None,
447
+ residual_dtype: torch.dtype = None,
448
+ is_rms_norm: bool = False,
449
+ ):
450
+ if residual is not None:
451
+ residual_dtype = residual.dtype
452
+ T, D = x.shape
453
+ if residual is not None:
454
+ assert residual.shape == (T, D)
455
+ if weight is not None:
456
+ assert weight.shape == (D,)
457
+ if bias is not None:
458
+ assert bias.shape == (D,)
459
+ # allocate output
460
+ y = torch.empty_like(x, dtype=x.dtype if out_dtype is None else out_dtype)
461
+ if residual is not None or (residual_dtype is not None and residual_dtype != x.dtype):
462
+ residual_out = torch.empty(T, D, device=x.device, dtype=residual_dtype)
463
+ else:
464
+ residual_out = None
465
+ mean = torch.empty((T,), dtype=torch.float, device=x.device) if not is_rms_norm else None
466
+ rstd = torch.empty((T,), dtype=torch.float, device=x.device)
467
+ # Less than 64KB per feature: enqueue fused kernel
468
+ MAX_FUSED_SIZE = 65536 // x.element_size()
469
+ BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
470
+ if D > BD:
471
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
472
+ # heuristics for number of warps
473
+
474
+ if D <= 512:
475
+ NB = triton.cdiv(T, 2048)
476
+ def grid(meta): return (triton.cdiv(T, meta['BT']),)
477
+ layer_norm_gated_fwd_kernel[grid](
478
+ x=x,
479
+ g=g,
480
+ y=y,
481
+ w=weight,
482
+ b=bias,
483
+ residual=residual,
484
+ residual_out=residual_out,
485
+ mean=mean,
486
+ rstd=rstd,
487
+ eps=eps,
488
+ T=T,
489
+ D=D,
490
+ BD=BD,
491
+ NB=NB,
492
+ ACTIVATION=activation,
493
+ IS_RMS_NORM=is_rms_norm,
494
+ )
495
+ else:
496
+ layer_norm_gated_fwd_kernel1[(T,)](
497
+ x=x,
498
+ g=g,
499
+ y=y,
500
+ w=weight,
501
+ b=bias,
502
+ residual=residual,
503
+ residual_out=residual_out,
504
+ mean=mean,
505
+ rstd=rstd,
506
+ eps=eps,
507
+ D=D,
508
+ BD=BD,
509
+ ACTIVATION=activation,
510
+ IS_RMS_NORM=is_rms_norm,
511
+ )
512
+ # residual_out is None if residual is None and residual_dtype == input_dtype
513
+ return y, mean, rstd, residual_out if residual_out is not None else x
514
+
515
+
516
+ def layer_norm_gated_bwd(
517
+ dy: torch.Tensor,
518
+ x: torch.Tensor,
519
+ g: torch.Tensor,
520
+ weight: torch.Tensor,
521
+ bias: torch.Tensor,
522
+ activation: str = 'swish',
523
+ eps: float = 1e-5,
524
+ mean: torch.Tensor = None,
525
+ rstd: torch.Tensor = None,
526
+ dresidual: torch.Tensor = None,
527
+ has_residual: bool = False,
528
+ is_rms_norm: bool = False,
529
+ x_dtype: torch.dtype = None,
530
+ recompute_output: bool = False,
531
+ ):
532
+ T, D = x.shape
533
+ assert dy.shape == (T, D)
534
+ if dresidual is not None:
535
+ assert dresidual.shape == (T, D)
536
+ if weight is not None:
537
+ assert weight.shape == (D,)
538
+ if bias is not None:
539
+ assert bias.shape == (D,)
540
+ # allocate output
541
+ dx = torch.empty_like(x) if x_dtype is None else torch.empty(T, D, dtype=x_dtype, device=x.device)
542
+ dg = torch.empty_like(g) if x_dtype is None else torch.empty(T, D, dtype=x_dtype, device=x.device)
543
+ dresidual_in = torch.empty_like(x) if has_residual and dx.dtype != x.dtype else None
544
+ y = torch.empty(T, D, dtype=dy.dtype, device=dy.device) if recompute_output else None
545
+
546
+ # Less than 64KB per feature: enqueue fused kernel
547
+ MAX_FUSED_SIZE = 65536 // x.element_size()
548
+ BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
549
+ if D > BD:
550
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
551
+ NS = get_multiprocessor_count(x.device.index)
552
+ BS = math.ceil(T / NS)
553
+
554
+ dw = torch.empty((NS, D), dtype=torch.float, device=weight.device) if weight is not None else None
555
+ db = torch.empty((NS, D), dtype=torch.float, device=bias.device) if bias is not None else None
556
+ grid = (NS,)
557
+
558
+ if D <= 512:
559
+ NB = triton.cdiv(T, 2048)
560
+ layer_norm_gated_bwd_kernel[grid](
561
+ x=x,
562
+ g=g,
563
+ w=weight,
564
+ b=bias,
565
+ y=y,
566
+ dy=dy,
567
+ dx=dx,
568
+ dg=dg,
569
+ dw=dw,
570
+ db=db,
571
+ dresidual=dresidual,
572
+ dresidual_in=dresidual_in,
573
+ mean=mean,
574
+ rstd=rstd,
575
+ T=T,
576
+ D=D,
577
+ BS=BS,
578
+ BD=BD,
579
+ NB=NB,
580
+ ACTIVATION=activation,
581
+ IS_RMS_NORM=is_rms_norm,
582
+ STORE_DRESIDUAL=dresidual_in is not None,
583
+ )
584
+ else:
585
+ layer_norm_gated_bwd_kernel1[grid](
586
+ x=x,
587
+ g=g,
588
+ w=weight,
589
+ b=bias,
590
+ y=y,
591
+ dy=dy,
592
+ dx=dx,
593
+ dg=dg,
594
+ dw=dw,
595
+ db=db,
596
+ dresidual=dresidual,
597
+ dresidual_in=dresidual_in,
598
+ mean=mean,
599
+ rstd=rstd,
600
+ T=T,
601
+ D=D,
602
+ BS=BS,
603
+ BD=BD,
604
+ ACTIVATION=activation,
605
+ IS_RMS_NORM=is_rms_norm,
606
+ STORE_DRESIDUAL=dresidual_in is not None,
607
+ )
608
+ dw = dw.sum(0).to(weight.dtype) if weight is not None else None
609
+ db = db.sum(0).to(bias.dtype) if bias is not None else None
610
+ # Don't need to compute dresidual_in separately in this case
611
+ if has_residual and dx.dtype == x.dtype:
612
+ dresidual_in = dx
613
+ return (dx, dg, dw, db, dresidual_in) if not recompute_output else (dx, dg, dw, db, dresidual_in, y)
614
+
615
+
616
+ class LayerNormGatedFunction(torch.autograd.Function):
617
+
618
+ @staticmethod
619
+ @input_guard
620
+ def forward(
621
+ ctx,
622
+ x: torch.Tensor,
623
+ g: torch.Tensor,
624
+ weight: torch.Tensor,
625
+ bias: torch.Tensor,
626
+ activation: str,
627
+ residual: torch.Tensor | None = None,
628
+ eps: float = 1e-6,
629
+ prenorm: bool = False,
630
+ residual_in_fp32: bool = False,
631
+ is_rms_norm: bool = False,
632
+ ):
633
+ x_shape_og = x.shape
634
+ g_shape_og = g.shape
635
+ # reshape input data into 2D tensor
636
+ x = x.reshape(-1, x.shape[-1])
637
+ g = g.reshape(-1, g.shape[-1])
638
+ if residual is not None:
639
+ assert residual.shape == x_shape_og
640
+ residual = residual.reshape(-1, residual.shape[-1])
641
+ residual_dtype = (
642
+ residual.dtype
643
+ if residual is not None
644
+ else (torch.float if residual_in_fp32 else None)
645
+ )
646
+ y, mean, rstd, residual_out = layer_norm_gated_fwd(
647
+ x=x,
648
+ g=g,
649
+ weight=weight,
650
+ bias=bias,
651
+ activation=activation,
652
+ eps=eps,
653
+ residual=residual,
654
+ residual_dtype=residual_dtype,
655
+ is_rms_norm=is_rms_norm,
656
+ )
657
+ ctx.save_for_backward(residual_out, g, weight, bias, mean, rstd)
658
+ ctx.x_shape_og = x_shape_og
659
+ ctx.g_shape_og = g_shape_og
660
+ ctx.activation = activation
661
+ ctx.eps = eps
662
+ ctx.is_rms_norm = is_rms_norm
663
+ ctx.has_residual = residual is not None
664
+ ctx.prenorm = prenorm
665
+ ctx.x_dtype = x.dtype
666
+ y = y.reshape(x_shape_og)
667
+ return y if not prenorm else (y, residual_out.reshape(x_shape_og))
668
+
669
+ @staticmethod
670
+ @input_guard
671
+ def backward(ctx, dy, *args):
672
+ x, g, weight, bias, mean, rstd = ctx.saved_tensors
673
+ dy = dy.reshape(-1, dy.shape[-1])
674
+ assert dy.shape == x.shape
675
+ if ctx.prenorm:
676
+ dresidual = args[0]
677
+ dresidual = dresidual.reshape(-1, dresidual.shape[-1])
678
+ assert dresidual.shape == x.shape
679
+ else:
680
+ dresidual = None
681
+ dx, dg, dw, db, dres_in = layer_norm_gated_bwd(
682
+ dy=dy,
683
+ x=x,
684
+ g=g,
685
+ weight=weight,
686
+ bias=bias,
687
+ activation=ctx.activation,
688
+ eps=ctx.eps,
689
+ mean=mean,
690
+ rstd=rstd,
691
+ dresidual=dresidual,
692
+ has_residual=ctx.has_residual,
693
+ is_rms_norm=ctx.is_rms_norm,
694
+ x_dtype=ctx.x_dtype,
695
+ )
696
+ return (
697
+ dx.reshape(ctx.x_shape_og),
698
+ dg.reshape(ctx.g_shape_og),
699
+ dw,
700
+ db,
701
+ None,
702
+ dres_in.reshape(ctx.x_shape_og) if ctx.has_residual else None,
703
+ None,
704
+ None,
705
+ None,
706
+ None,
707
+ )
708
+
709
+
710
+ class LayerNormGatedLinearFunction(torch.autograd.Function):
711
+
712
+ @staticmethod
713
+ @input_guard
714
+ def forward(
715
+ ctx,
716
+ x: torch.Tensor,
717
+ g: torch.Tensor,
718
+ norm_weight: torch.Tensor,
719
+ norm_bias: torch.Tensor,
720
+ linear_weight: torch.Tensor,
721
+ linear_bias: torch.Tensor,
722
+ residual: torch.Tensor | None = None,
723
+ eps: float = 1e-6,
724
+ prenorm: bool = False,
725
+ residual_in_fp32: bool = False,
726
+ is_rms_norm: bool = False,
727
+ ):
728
+ x_shape_og = x.shape
729
+ g_shape_og = g.shape
730
+ # reshape input data into 2D tensor
731
+ x = x.reshape(-1, x.shape[-1])
732
+ g = g.reshape(-1, g.shape[-1])
733
+ if residual is not None:
734
+ assert residual.shape == x_shape_og
735
+ residual = residual.reshape(-1, residual.shape[-1])
736
+ residual_dtype = (
737
+ residual.dtype
738
+ if residual is not None
739
+ else (torch.float if residual_in_fp32 else None)
740
+ )
741
+ y, mean, rstd, residual_out = layer_norm_gated_fwd(
742
+ x=x,
743
+ g=g,
744
+ weight=norm_weight,
745
+ bias=norm_bias,
746
+ eps=eps,
747
+ residual=residual,
748
+ residual_dtype=residual_dtype,
749
+ is_rms_norm=is_rms_norm,
750
+ )
751
+ y = y.reshape(x_shape_og)
752
+ dtype = torch.get_autocast_gpu_dtype() if torch.is_autocast_enabled() else y.dtype
753
+ linear_weight = linear_weight.to(dtype)
754
+ linear_bias = linear_bias.to(dtype) if linear_bias is not None else None
755
+ out = F.linear(y.to(linear_weight.dtype), linear_weight, linear_bias)
756
+ # We don't store y, will be recomputed in the backward pass to save memory
757
+ ctx.save_for_backward(residual_out, g, norm_weight, norm_bias, linear_weight, mean, rstd)
758
+ ctx.x_shape_og = x_shape_og
759
+ ctx.g_shape_og = g_shape_og
760
+ ctx.eps = eps
761
+ ctx.is_rms_norm = is_rms_norm
762
+ ctx.has_residual = residual is not None
763
+ ctx.prenorm = prenorm
764
+ ctx.x_dtype = x.dtype
765
+ ctx.linear_bias_is_none = linear_bias is None
766
+ return out if not prenorm else (out, residual_out.reshape(x_shape_og))
767
+
768
+ @staticmethod
769
+ @input_guard
770
+ def backward(ctx, dout, *args):
771
+ x, g, norm_weight, norm_bias, linear_weight, mean, rstd = ctx.saved_tensors
772
+ dout = dout.reshape(-1, dout.shape[-1])
773
+ dy = F.linear(dout, linear_weight.t())
774
+ dlinear_bias = None if ctx.linear_bias_is_none else dout.sum(0)
775
+ assert dy.shape == x.shape
776
+ if ctx.prenorm:
777
+ dresidual = args[0]
778
+ dresidual = dresidual.reshape(-1, dresidual.shape[-1])
779
+ assert dresidual.shape == x.shape
780
+ else:
781
+ dresidual = None
782
+ dx, dg, dnorm_weight, dnorm_bias, dres_in, y = layer_norm_gated_bwd(
783
+ dy=dy,
784
+ x=x,
785
+ g=g,
786
+ weight=norm_weight,
787
+ bias=norm_bias,
788
+ eps=ctx.eps,
789
+ mean=mean,
790
+ rstd=rstd,
791
+ dresidual=dresidual,
792
+ has_residual=ctx.has_residual,
793
+ is_rms_norm=ctx.is_rms_norm,
794
+ x_dtype=ctx.x_dtype,
795
+ recompute_output=True,
796
+ )
797
+ dlinear_weight = torch.einsum("bo,bi->oi", dout, y)
798
+ return (
799
+ dx.reshape(ctx.x_shape_og),
800
+ dg.reshape(ctx.g_shape_og),
801
+ dnorm_weight,
802
+ dnorm_bias,
803
+ dlinear_weight,
804
+ dlinear_bias,
805
+ dres_in.reshape(ctx.x_shape_og) if ctx.has_residual else None,
806
+ None,
807
+ None,
808
+ None,
809
+ None,
810
+ )
811
+
812
+
813
+ def layer_norm_gated(
814
+ x: torch.Tensor,
815
+ g: torch.Tensor,
816
+ weight: torch.Tensor,
817
+ bias: torch.Tensor,
818
+ activation: str = 'swish',
819
+ residual: torch.Tensor | None = None,
820
+ prenorm: bool = False,
821
+ residual_in_fp32: bool = False,
822
+ eps: float = 1e-6,
823
+ ):
824
+ return LayerNormGatedFunction.apply(
825
+ x,
826
+ g,
827
+ weight,
828
+ bias,
829
+ activation,
830
+ residual,
831
+ eps,
832
+ prenorm,
833
+ residual_in_fp32,
834
+ False,
835
+ )
836
+
837
+
838
+ def rms_norm_gated(
839
+ x: torch.Tensor,
840
+ g: torch.Tensor,
841
+ weight: torch.Tensor,
842
+ bias: torch.Tensor,
843
+ activation: str = 'swish',
844
+ residual: torch.Tensor | None = None,
845
+ prenorm: bool = False,
846
+ residual_in_fp32: bool = False,
847
+ eps: float = 1e-6,
848
+ ):
849
+ return LayerNormGatedFunction.apply(
850
+ x,
851
+ g,
852
+ weight,
853
+ bias,
854
+ activation,
855
+ residual,
856
+ eps,
857
+ prenorm,
858
+ residual_in_fp32,
859
+ True,
860
+ )
861
+
862
+
863
+ def layer_norm_swish_gate_linear(
864
+ x: torch.Tensor,
865
+ g: torch.Tensor,
866
+ norm_weight: torch.Tensor,
867
+ norm_bias: torch.Tensor,
868
+ linear_weight: torch.Tensor,
869
+ linear_bias: torch.Tensor,
870
+ residual: torch.Tensor | None = None,
871
+ prenorm: bool = False,
872
+ residual_in_fp32: bool = False,
873
+ eps: float = 1e-6,
874
+ ):
875
+ return LayerNormGatedLinearFunction.apply(
876
+ x,
877
+ g,
878
+ norm_weight,
879
+ norm_bias,
880
+ linear_weight,
881
+ linear_bias,
882
+ residual,
883
+ eps,
884
+ prenorm,
885
+ residual_in_fp32,
886
+ False,
887
+ )
888
+
889
+
890
+ def rms_norm_swish_gate_linear(
891
+ x,
892
+ g: torch.Tensor,
893
+ norm_weight: torch.Tensor,
894
+ norm_bias: torch.Tensor,
895
+ linear_weight: torch.Tensor,
896
+ linear_bias: torch.Tensor,
897
+ residual: torch.Tensor | None = None,
898
+ prenorm: bool = False,
899
+ residual_in_fp32: bool = False,
900
+ eps: float = 1e-6,
901
+ ):
902
+ return LayerNormGatedLinearFunction.apply(
903
+ x,
904
+ g,
905
+ norm_weight,
906
+ norm_bias,
907
+ linear_weight,
908
+ linear_bias,
909
+ residual,
910
+ eps,
911
+ prenorm,
912
+ residual_in_fp32,
913
+ True,
914
+ )
915
+
916
+
917
+ class FusedLayerNormGated(nn.Module):
918
+
919
+ def __init__(
920
+ self,
921
+ hidden_size: int,
922
+ elementwise_affine: bool = True,
923
+ bias: bool = False,
924
+ activation: str = 'swish',
925
+ eps: float = 1e-5,
926
+ device: torch.device | None = None,
927
+ dtype: torch.dtype | None = None,
928
+ ) -> FusedLayerNormGated:
929
+ factory_kwargs = {"device": device, "dtype": dtype}
930
+ super().__init__()
931
+
932
+ self.hidden_size = hidden_size
933
+ self.elementwise_affine = elementwise_affine
934
+ self.eps = eps
935
+ self.activation = activation
936
+
937
+ if self.activation not in ['swish', 'silu', 'sigmoid']:
938
+ raise ValueError(f"Unsupported activation: {self.activation}")
939
+
940
+ self.register_parameter("weight", None)
941
+ self.register_parameter("bias", None)
942
+ if elementwise_affine:
943
+ self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
944
+ if bias:
945
+ self.bias = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
946
+
947
+ self.reset_parameters()
948
+
949
+ def reset_parameters(self):
950
+ if self.elementwise_affine:
951
+ nn.init.ones_(self.weight)
952
+ if self.bias is not None:
953
+ nn.init.zeros_(self.bias)
954
+
955
+ def __repr__(self) -> str:
956
+ s = f"{self.__class__.__name__}({self.hidden_size}"
957
+ if not self.elementwise_affine:
958
+ s += f", elementwise_affine={self.elementwise_affine}"
959
+ s += f", eps={self.eps}"
960
+ s += f", activation={self.activation}"
961
+ s += ")"
962
+ return s
963
+
964
+ def forward(
965
+ self,
966
+ x: torch.Tensor,
967
+ g: torch.Tensor,
968
+ residual: torch.Tensor | None = None,
969
+ prenorm: bool = False,
970
+ residual_in_fp32: bool = False,
971
+ ) -> torch.Tensor:
972
+ return layer_norm_gated(
973
+ x,
974
+ g,
975
+ self.weight,
976
+ self.bias,
977
+ self.activation,
978
+ residual=residual,
979
+ eps=self.eps,
980
+ prenorm=prenorm,
981
+ residual_in_fp32=residual_in_fp32,
982
+ )
983
+
984
+
985
+ class FusedRMSNormGated(nn.Module):
986
+
987
+ def __init__(
988
+ self,
989
+ hidden_size: int,
990
+ elementwise_affine: bool = True,
991
+ eps: float = 1e-5,
992
+ activation: str = 'swish',
993
+ device: torch.device | None = None,
994
+ dtype: torch.dtype | None = None,
995
+ ) -> FusedRMSNormGated:
996
+ factory_kwargs = {"device": device, "dtype": dtype}
997
+ super().__init__()
998
+
999
+ self.hidden_size = hidden_size
1000
+ self.elementwise_affine = elementwise_affine
1001
+ self.eps = eps
1002
+ self.activation = activation
1003
+
1004
+ if self.activation not in ['swish', 'silu', 'sigmoid']:
1005
+ raise ValueError(f"Unsupported activation: {self.activation}")
1006
+
1007
+ if elementwise_affine:
1008
+ self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
1009
+ else:
1010
+ self.register_parameter("weight", None)
1011
+ self.register_parameter("bias", None)
1012
+
1013
+ self.reset_parameters()
1014
+
1015
+ def reset_parameters(self):
1016
+ if self.elementwise_affine:
1017
+ nn.init.ones_(self.weight)
1018
+
1019
+ def __repr__(self) -> str:
1020
+ s = f"{self.__class__.__name__}({self.hidden_size}"
1021
+ if not self.elementwise_affine:
1022
+ s += f", elementwise_affine={self.elementwise_affine}"
1023
+ s += f", eps={self.eps}"
1024
+ s += f", activation={self.activation}"
1025
+ s += ")"
1026
+ return s
1027
+
1028
+ def forward(
1029
+ self,
1030
+ x: torch.Tensor,
1031
+ g: torch.Tensor,
1032
+ residual: torch.Tensor | None = None,
1033
+ prenorm: bool = False,
1034
+ residual_in_fp32: bool = False,
1035
+ ) -> torch.Tensor:
1036
+ return rms_norm_gated(
1037
+ x,
1038
+ g,
1039
+ self.weight,
1040
+ self.bias,
1041
+ self.activation,
1042
+ residual=residual,
1043
+ eps=self.eps,
1044
+ prenorm=prenorm,
1045
+ residual_in_fp32=residual_in_fp32,
1046
+ )
1047
+
1048
+
1049
+ class FusedLayerNormSwishGate(FusedLayerNormGated):
1050
+
1051
+ def __init__(
1052
+ self,
1053
+ hidden_size: int,
1054
+ elementwise_affine: bool = True,
1055
+ bias: bool = False,
1056
+ eps: float = 1e-5,
1057
+ device: torch.device | None = None,
1058
+ dtype: torch.dtype | None = None,
1059
+ ) -> FusedLayerNormSwishGate:
1060
+ super().__init__(
1061
+ hidden_size=hidden_size,
1062
+ elementwise_affine=elementwise_affine,
1063
+ bias=bias,
1064
+ eps=eps,
1065
+ device=device,
1066
+ dtype=dtype,
1067
+ )
1068
+
1069
+
1070
+ class FusedRMSNormSwishGate(FusedRMSNormGated):
1071
+
1072
+ def __init__(
1073
+ self,
1074
+ hidden_size: int,
1075
+ elementwise_affine: bool = True,
1076
+ eps: float = 1e-5,
1077
+ device: torch.device | None = None,
1078
+ dtype: torch.dtype | None = None,
1079
+ ) -> FusedRMSNormSwishGate:
1080
+ super().__init__(
1081
+ hidden_size=hidden_size,
1082
+ elementwise_affine=elementwise_affine,
1083
+ eps=eps,
1084
+ device=device,
1085
+ dtype=dtype,
1086
+ )
1087
+
1088
+
1089
+ class FusedLayerNormGatedLinear(nn.Module):
1090
+
1091
+ def __init__(
1092
+ self,
1093
+ hidden_size: int,
1094
+ elementwise_affine: bool = True,
1095
+ eps: float = 1e-5,
1096
+ device: torch.device | None = None,
1097
+ dtype: torch.dtype | None = None,
1098
+ ) -> FusedLayerNormGatedLinear:
1099
+ factory_kwargs = {"device": device, "dtype": dtype}
1100
+ super().__init__()
1101
+
1102
+ self.hidden_size = hidden_size
1103
+ self.elementwise_affine = elementwise_affine
1104
+ self.eps = eps
1105
+
1106
+ if elementwise_affine:
1107
+ self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
1108
+ else:
1109
+ self.register_parameter("weight", None)
1110
+ self.register_parameter("bias", None)
1111
+
1112
+ self.reset_parameters()
1113
+
1114
+ def reset_parameters(self):
1115
+ if self.elementwise_affine:
1116
+ nn.init.ones_(self.weight)
1117
+
1118
+ def __repr__(self) -> str:
1119
+ s = f"{self.__class__.__name__}({self.hidden_size}"
1120
+ if not self.elementwise_affine:
1121
+ s += f", elementwise_affine={self.elementwise_affine}"
1122
+ s += f", eps={self.eps}"
1123
+ s += ")"
1124
+ return s
1125
+
1126
+ def forward(
1127
+ self,
1128
+ x: torch.Tensor,
1129
+ g: torch.Tensor,
1130
+ weight: torch.Tensor | None = None,
1131
+ bias: torch.Tensor | None = None,
1132
+ residual: torch.Tensor | None = None,
1133
+ prenorm: bool = False,
1134
+ residual_in_fp32: bool = False,
1135
+ ) -> torch.Tensor:
1136
+ return layer_norm_swish_gate_linear(
1137
+ x,
1138
+ g,
1139
+ self.weight,
1140
+ self.bias,
1141
+ weight,
1142
+ bias,
1143
+ residual=residual,
1144
+ eps=self.eps,
1145
+ prenorm=prenorm,
1146
+ residual_in_fp32=residual_in_fp32,
1147
+ )
1148
+
1149
+
1150
+ class FusedLayerNormSwishGateLinear(FusedLayerNormGatedLinear):
1151
+
1152
+ def __init__(
1153
+ self,
1154
+ hidden_size: int,
1155
+ elementwise_affine: bool = True,
1156
+ eps: float = 1e-5,
1157
+ device: torch.device | None = None,
1158
+ dtype: torch.dtype | None = None,
1159
+ ) -> FusedLayerNormSwishGateLinear:
1160
+ super().__init__(
1161
+ hidden_size=hidden_size,
1162
+ elementwise_affine=elementwise_affine,
1163
+ eps=eps,
1164
+ device=device,
1165
+ dtype=dtype,
1166
+ )
1167
+
1168
+
1169
+ class FusedRMSNormGatedLinear(nn.Module):
1170
+
1171
+ def __init__(
1172
+ self,
1173
+ hidden_size,
1174
+ elementwise_affine: bool = True,
1175
+ eps: float = 1e-5,
1176
+ device: torch.device | None = None,
1177
+ dtype: torch.dtype | None = None,
1178
+ ) -> FusedRMSNormGatedLinear:
1179
+ factory_kwargs = {"device": device, "dtype": dtype}
1180
+ super().__init__()
1181
+
1182
+ self.hidden_size = hidden_size
1183
+ self.elementwise_affine = elementwise_affine
1184
+ self.eps = eps
1185
+
1186
+ self.register_parameter("weight", None)
1187
+ self.register_parameter("bias", None)
1188
+ if elementwise_affine:
1189
+ self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
1190
+
1191
+ self.reset_parameters()
1192
+
1193
+ def reset_parameters(self):
1194
+ if self.elementwise_affine:
1195
+ nn.init.ones_(self.weight)
1196
+
1197
+ def __repr__(self) -> str:
1198
+ s = f"{self.__class__.__name__}({self.hidden_size}"
1199
+ if not self.elementwise_affine:
1200
+ s += f", elementwise_affine={self.elementwise_affine}"
1201
+ s += f", eps={self.eps}"
1202
+ s += ")"
1203
+ return s
1204
+
1205
+ def forward(
1206
+ self,
1207
+ x: torch.Tensor,
1208
+ g: torch.Tensor,
1209
+ weight: torch.Tensor | None = None,
1210
+ bias: torch.Tensor | None = None,
1211
+ residual: torch.Tensor | None = None,
1212
+ prenorm: bool = False,
1213
+ residual_in_fp32: bool = False,
1214
+ ) -> torch.Tensor:
1215
+ return rms_norm_swish_gate_linear(
1216
+ x,
1217
+ g,
1218
+ self.weight,
1219
+ self.bias,
1220
+ weight,
1221
+ bias,
1222
+ residual=residual,
1223
+ eps=self.eps,
1224
+ prenorm=prenorm,
1225
+ residual_in_fp32=residual_in_fp32,
1226
+ )
1227
+
1228
+
1229
+ class FusedRMSNormSwishGateLinear(FusedRMSNormGatedLinear):
1230
+
1231
+ def __init__(
1232
+ self,
1233
+ hidden_size: int,
1234
+ elementwise_affine: bool = True,
1235
+ eps: float = 1e-5,
1236
+ device: torch.device | None = None,
1237
+ dtype: torch.dtype | None = None,
1238
+ ) -> FusedRMSNormSwishGateLinear:
1239
+ super().__init__(
1240
+ hidden_size=hidden_size,
1241
+ elementwise_affine=elementwise_affine,
1242
+ eps=eps,
1243
+ device=device,
1244
+ dtype=dtype,
1245
+ )
code/flash-linear-attention/fla/modules/grpo.py ADDED
@@ -0,0 +1,412 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # modified from https://github.com/mdy666/mdy_triton/blob/e0a856347bd988e05e0152332bba35f1d33c5b1f/others/grpo/grpo_loss.ipynb
2
+ # XHS ID: blueeeee
3
+
4
+ # https://github.com/huggingface/trl/blob/main/trl/trainer/grpo_trainer.py
5
+ """
6
+ # Get the per-token log probabilities for the completions for the model and the reference model
7
+ def _get_per_token_logps(self, model, input_ids, attention_mask, logits_to_keep):
8
+ # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded
9
+ logits = model(input_ids=input_ids, attention_mask=attention_mask, logits_to_keep=logits_to_keep + 1).logits
10
+ logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred
11
+
12
+ input_ids = input_ids[:, -logits_to_keep:]
13
+ # For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves.
14
+ # See https://github.com/huggingface/trl/issues/2770
15
+ logits = logits[:, -logits_to_keep:]
16
+ return selective_log_softmax(logits, input_ids) # compute logprobs for the input tokens
17
+
18
+ def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
19
+ if return_outputs:
20
+ raise ValueError("The GRPOTrainer does not support returning outputs")
21
+ # Compute the per-token log probabilities for the model
22
+
23
+ prompt_ids, prompt_mask = inputs["prompt_ids"], inputs["prompt_mask"]
24
+ completion_ids, completion_mask = inputs["completion_ids"], inputs["completion_mask"]
25
+ input_ids = torch.cat([prompt_ids, completion_ids], dim=1)
26
+ attention_mask = torch.cat([prompt_mask, completion_mask], dim=1)
27
+ logits_to_keep = completion_ids.size(1) # we only need to compute the logits for the completion tokens
28
+
29
+ per_token_logps = self._get_per_token_logps(model, input_ids, attention_mask, logits_to_keep)
30
+
31
+ # Compute the KL divergence between the model and the reference model
32
+ ref_per_token_logps = inputs["ref_per_token_logps"]
33
+ per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1
34
+
35
+ # x - x.detach() allows for preserving gradients from x
36
+ advantages = inputs["advantages"]
37
+ per_token_loss = torch.exp(per_token_logps - per_token_logps.detach()) * advantages.unsqueeze(1)
38
+ per_token_loss = -(per_token_loss - self.beta * per_token_kl)
39
+ loss = ((per_token_loss * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean()
40
+
41
+ # Log the metrics
42
+ completion_length = self.accelerator.gather_for_metrics(completion_mask.sum(1)).float().mean().item()
43
+ self._metrics["completion_length"].append(completion_length)
44
+
45
+ mean_kl = ((per_token_kl * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean()
46
+ self._metrics["kl"].append(self.accelerator.gather_for_metrics(mean_kl).mean().item())
47
+
48
+ return loss
49
+ """
50
+
51
+
52
+ import torch
53
+ import triton
54
+ import triton.language as tl
55
+
56
+ from fla.ops.utils.op import exp, log
57
+ from fla.utils import autotune_cache_kwargs, input_guard, is_amd
58
+
59
+ NUM_WARPS_AUTOTUNE = [4, 8, 16] if is_amd else [4, 8, 16, 32]
60
+
61
+
62
+ @triton.autotune(
63
+ configs=[
64
+ triton.Config({'BLOCK_SIZE': BLOCK_SIZE}, num_warps=NUM_WARPS, num_stages=NUM_STAGES)
65
+ for BLOCK_SIZE in [1024, 2048, 4096, 8192]
66
+ for NUM_WARPS in NUM_WARPS_AUTOTUNE
67
+ for NUM_STAGES in [1, 2, 4]
68
+ ],
69
+ key=['B', 'N'],
70
+ **autotune_cache_kwargs,
71
+ )
72
+ @triton.jit
73
+ def grpo_fwd_kernel(
74
+ logits_ptr,
75
+ ref_logp_ptr,
76
+ input_ids_ptr,
77
+ advantages_ptr,
78
+ completion_mask_ptr,
79
+ loss_ptr,
80
+ lse_ptr,
81
+ beta,
82
+ save_kl: tl.constexpr,
83
+ B,
84
+ M,
85
+ N,
86
+ L,
87
+ start_idx,
88
+ BLOCK_SIZE: tl.constexpr,
89
+ ):
90
+ row_idx = tl.program_id(0)
91
+
92
+ off_b = row_idx // L
93
+ N = tl.cast(N, tl.int64)
94
+
95
+ loss_ptr += row_idx
96
+
97
+ completion_mask_ptr += row_idx
98
+ not_skip = tl.load(completion_mask_ptr).to(tl.int1)
99
+ if not_skip == 1:
100
+ ref_logp_ptr += row_idx
101
+ lse_ptr += row_idx
102
+ advantages_ptr += off_b
103
+ logits_ptr += N * (row_idx + off_b)
104
+ input_ids_ptr += row_idx + (off_b+1) * start_idx
105
+ base_cols = tl.arange(0, BLOCK_SIZE)
106
+
107
+ m_i = -float("inf")
108
+ l_i = 0.0
109
+ for start_n in tl.range(0, N, BLOCK_SIZE):
110
+ cols = start_n + base_cols
111
+ mask = cols < N
112
+ logits = tl.load(logits_ptr+cols, mask=mask, other=-float('inf')).to(tl.float32)
113
+ m_ij = tl.max(logits)
114
+ new_m_i = tl.maximum(m_i, m_ij)
115
+ l_i = l_i * exp(m_i - new_m_i) + tl.sum(exp(logits - new_m_i))
116
+ m_i = new_m_i
117
+ lse = log(l_i) + m_i
118
+
119
+ idx = tl.load(input_ids_ptr)
120
+ x = tl.load(logits_ptr+idx).to(tl.float32)
121
+ advantage = tl.load(advantages_ptr).to(tl.float32)
122
+ ref_logp = tl.load(ref_logp_ptr)
123
+ logp = x - lse
124
+ diff = ref_logp - logp
125
+ kl = exp(diff) - diff - 1
126
+ loss = kl * beta - advantage
127
+
128
+ tl.store(loss_ptr, loss.to(loss_ptr.dtype.element_ty))
129
+ tl.store(lse_ptr, lse.to(lse_ptr.dtype.element_ty))
130
+ if save_kl:
131
+ tl.store(loss_ptr+M, kl.to(loss_ptr.dtype.element_ty))
132
+ else:
133
+ # store 0
134
+ tl.store(loss_ptr, 0.0)
135
+ if save_kl:
136
+ tl.store(loss_ptr+M, 0.0)
137
+
138
+
139
+ @triton.autotune(
140
+ configs=[
141
+ triton.Config({}, num_warps=NUM_WARPS, num_stages=NUM_STAGES)
142
+ for NUM_WARPS in [32]
143
+ for NUM_STAGES in [4]
144
+ ],
145
+ key=['B', 'N'],
146
+ **autotune_cache_kwargs,
147
+ )
148
+ @triton.jit
149
+ def grpo_bwd_kernel(
150
+ dloss_ptr,
151
+ dlogits_ptr,
152
+ logits_ptr,
153
+ ref_logp_ptr,
154
+ input_ids_ptr,
155
+ advantages_ptr,
156
+ completion_mask_ptr,
157
+ lse_ptr,
158
+ beta,
159
+ B,
160
+ N,
161
+ L,
162
+ start_idx,
163
+ BLOCK_SIZE: tl.constexpr,
164
+ ):
165
+
166
+ row_idx = tl.program_id(0) # B*L
167
+ off_b = row_idx // L
168
+
169
+ N = tl.cast(N, tl.int64)
170
+
171
+ dlogits_ptr += N * (row_idx + off_b)
172
+ base_cols = tl.arange(0, BLOCK_SIZE)
173
+ completion_mask_ptr += row_idx
174
+ not_skip = tl.load(completion_mask_ptr).to(tl.int1)
175
+
176
+ if not_skip == 1:
177
+ lse_ptr += row_idx
178
+ dloss_ptr += row_idx
179
+ advantages_ptr += off_b
180
+ ref_logp_ptr += row_idx
181
+ logits_ptr += N * (row_idx + off_b)
182
+ input_ids_ptr += row_idx + (off_b+1) * start_idx
183
+ dloss = tl.load(dloss_ptr).to(tl.float32)
184
+ lse = tl.load(lse_ptr).to(tl.float32)
185
+ idx = tl.load(input_ids_ptr)
186
+ x = tl.load(logits_ptr+idx).to(tl.float32)
187
+ advantage = tl.load(advantages_ptr).to(tl.float32)
188
+ ref_logp = tl.load(ref_logp_ptr)
189
+ # Need for in-place grad.
190
+ tl.debug_barrier()
191
+ logp = x - lse
192
+
193
+ dlogp = (beta * (-1.0 * exp(ref_logp - logp) + 1)
194
+ - advantage) * dloss
195
+
196
+ for start_n in tl.range(0, N, BLOCK_SIZE):
197
+ cols = start_n + base_cols
198
+ mask = cols < N
199
+ logits = tl.load(logits_ptr+cols, mask=mask, other=-float('inf')).to(tl.float32)
200
+ probs = exp(logits - lse)
201
+ dlogits = tl.where(cols == idx, 1-probs, -probs) * dlogp
202
+
203
+ tl.store(dlogits_ptr+cols, dlogits.to(dlogits_ptr.dtype.element_ty), mask=mask)
204
+ else:
205
+ dlogits = tl.zeros((BLOCK_SIZE,), dtype=tl.float32)
206
+ for start_n in tl.range(0, N, BLOCK_SIZE):
207
+ cols = start_n + base_cols
208
+ mask = cols < N
209
+
210
+ tl.store(dlogits_ptr+cols, dlogits.to(dlogits_ptr.dtype.element_ty), mask=mask)
211
+
212
+
213
+ class GrpoLoss(torch.autograd.Function):
214
+
215
+ @input_guard
216
+ @staticmethod
217
+ def forward(ctx, logits, ref_logp, input_ids, advantages, beta, completion_mask, save_kl, inplace=True):
218
+ ctx.input_shape = logits.shape
219
+ B, L_ADD_1, N = ctx.input_shape
220
+ L = L_ADD_1 - 1
221
+ M = B * L
222
+ input_ids_start_index = input_ids.size(1) - L
223
+
224
+ if not save_kl:
225
+ loss = torch.empty(B, L, device=logits.device, dtype=torch.float32)
226
+ else:
227
+ loss = torch.empty(B*2, L, device=logits.device, dtype=torch.float32)
228
+
229
+ lse = torch.empty(B, L, device=logits.device, dtype=torch.float32)
230
+
231
+ if completion_mask is None:
232
+ completion_mask = torch.ones(B, L, device=logits.device, dtype=torch.int32)
233
+ else:
234
+ loss[:B].masked_fill_(completion_mask.logical_not(), 0.0)
235
+
236
+ grpo_fwd_kernel[(M,)](
237
+ logits_ptr=logits,
238
+ ref_logp_ptr=ref_logp,
239
+ input_ids_ptr=input_ids,
240
+ advantages_ptr=advantages,
241
+ completion_mask_ptr=completion_mask,
242
+ loss_ptr=loss,
243
+ lse_ptr=lse,
244
+ beta=beta,
245
+ save_kl=save_kl,
246
+ B=B, M=M, N=N, L=L,
247
+ start_idx=input_ids_start_index,
248
+ )
249
+ ctx.beta = beta
250
+ ctx.save_for_backward(lse, logits, input_ids, advantages, completion_mask)
251
+ ctx.ref_logp = ref_logp
252
+ ctx.inplace = inplace
253
+ return loss
254
+
255
+ @input_guard
256
+ @staticmethod
257
+ def backward(ctx, dloss):
258
+ # The grad of logits comes from two parts, the reward part and the kl part
259
+ lse, logits, input_ids, advantages, completion_mask = ctx.saved_tensors
260
+ inplace = ctx.inplace
261
+ B, L_ADD_1, N = ctx.input_shape
262
+ L = L_ADD_1 - 1
263
+ M = B * L
264
+
265
+ input_ids_start_index = input_ids.size(1) - L
266
+
267
+ # B, L_ADD_1, N
268
+ dlogits = logits if inplace else torch.empty_like(logits)
269
+ BN = min(65536, triton.next_power_of_2(N))
270
+
271
+ grpo_bwd_kernel[(M,)](
272
+ dloss_ptr=dloss,
273
+ dlogits_ptr=dlogits,
274
+ logits_ptr=logits,
275
+ ref_logp_ptr=ctx.ref_logp,
276
+ input_ids_ptr=input_ids,
277
+ advantages_ptr=advantages,
278
+ completion_mask_ptr=completion_mask,
279
+ lse_ptr=lse,
280
+ beta=ctx.beta,
281
+ B=B, N=N, L=L,
282
+ BLOCK_SIZE=BN,
283
+ start_idx=input_ids_start_index,
284
+ )
285
+ # The last token in the completion is not used in the loss computation
286
+ # and therefore its gradient should be set to 0
287
+ dlogits[:, -1, :].fill_(0.0)
288
+ return dlogits.view(*ctx.input_shape), None, None, None, None, None, None, None
289
+
290
+
291
+ def fused_grpo_loss(logits, ref_logp, input_ids, advantages,
292
+ beta=0.1, completion_mask=None, save_kl=False, inplace=False) -> torch.Tensor:
293
+ '''
294
+ compute grpo loss, save memory(no addition usage) and fast speed(6X for A800)
295
+
296
+ Args:
297
+ logtits: Tensor, [B, L+1, vocab_size], the origin output of model, it's not logits[:, :-1]
298
+ ref_logp: Tensor, [B, L], the origin output of model, it's not ref_logits[:, :-1]
299
+ input_ids: Tensor, [B, K+L], it's prompt_completion_id, it contains the prompt ids and output ids
300
+ advantages: Tensor, [B], the advantages of each prompt
301
+ beta: float, the weight of kl loss
302
+ completion_mask: Tensor, loss mask
303
+ save_kl: bool, if true will save kl
304
+
305
+ Retutn:
306
+ loss: Tensor, [B, L], the loss of grpo, it contains the advantage part and kl part
307
+
308
+ NOTE: logits(ref_logits) is computed by these steps
309
+ logits_to_keep = completion_ids.size(1)
310
+
311
+ def get_per_token_logits(model, input_ids, attention_mask, logits_to_keep):
312
+ # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded
313
+ logits = model(
314
+ input_ids=input_ids, attention_mask=attention_mask, logits_to_keep=logits_to_keep + 1
315
+ ).logits
316
+ return logits
317
+
318
+ logits = get_per_token_logits(model, prompt_completion_ids, attention_mask, logits_to_keep)
319
+ '''
320
+ out = GrpoLoss.apply(logits, ref_logp, input_ids, advantages, beta, completion_mask, save_kl, inplace)
321
+ if not save_kl:
322
+ return out
323
+ else:
324
+ return out.chunk(2, axis=0)
325
+
326
+
327
+ def grpo_loss_torch(logits, ref_logp, input_ids, advantages, beta=0.1, completion_mask=None, save_kl=False):
328
+ def get_log_probs(logits, input_ids):
329
+ per_token_logps = []
330
+ for logits_row, input_ids_row in zip(logits, input_ids[:, -logits.size(1):], strict=False):
331
+ log_probs = logits_row.log_softmax(dim=-1)
332
+ token_log_prob = torch.gather(log_probs, dim=1, index=input_ids_row.unsqueeze(1)).squeeze(1)
333
+ per_token_logps.append(token_log_prob)
334
+ return torch.stack(per_token_logps)
335
+
336
+ logits = logits[:, :-1]
337
+ per_token_logps = get_log_probs(logits, input_ids)
338
+ ref_per_token_logps = ref_logp
339
+ per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1
340
+
341
+ per_token_loss = torch.exp(per_token_logps - per_token_logps.detach()) * advantages.unsqueeze(1)
342
+ per_token_loss = -(per_token_loss - beta * per_token_kl)
343
+ if completion_mask is not None:
344
+ per_token_loss *= completion_mask
345
+ if save_kl:
346
+ per_token_kl *= completion_mask
347
+ return per_token_loss if not save_kl else (per_token_loss, per_token_kl)
348
+
349
+
350
+ @torch.compile(fullgraph=True)
351
+ def grpo_loss_with_old_logps(
352
+ logps: torch.Tensor,
353
+ ref_logps: torch.Tensor,
354
+ old_logps: torch.Tensor,
355
+ pad_mask: torch.Tensor,
356
+ logits_to_keep: int,
357
+ rewards: torch.Tensor,
358
+ beta: float = 0.2,
359
+ epsilon: float = 0.2,
360
+ ):
361
+ """
362
+ Compute the GRPO (Group Relative Policy Optimization) loss.
363
+
364
+ Args:
365
+ logps (torch.Tensor): [Batch, Token_length] Log probabilities of the current policy.
366
+ ref_logps (torch.Tensor):[Batch, Token_length] Log probabilities of the reference policy.
367
+ old_logps (torch.Tensor): [Batch, Token_length] Log probabilities of the old policy.
368
+ completion_ids (torch.Tensor): [Batch, Token_length] Completion token IDs (bool).
369
+ pad_token_id: Pad token ID.
370
+ logits_to_keep (int): Number of logits to keep for masking.
371
+ rewards (torch.Tensor): [Batch] Rewards for each generation.
372
+ beta (float) = 0.2: A hyperparameter for weighting the KL divergence term.
373
+ epsilon (float) = 0.2: An float hyperparameter for clipping the importance weights.
374
+
375
+ Returns:
376
+ torch.Tensor: The computed GRPO loss.
377
+ """
378
+ B = logps.shape[0]
379
+ assert B > 1, "Batch * Num generations should be greater than 1"
380
+
381
+ rewards_shaped = rewards.view(-1, B) # B,num_generations
382
+ advantages = (rewards_shaped - rewards_shaped.mean(dim=1, keepdim=True)) / \
383
+ (rewards_shaped.std(dim=1, keepdim=True) + 1e-8)
384
+ advantages = advantages.view(-1) # B*num_generations
385
+ # Calculate the per - token KL divergence
386
+ per_token_kl = torch.exp(ref_logps - logps) - (ref_logps - logps) - 1
387
+
388
+ # Calculate the ratio of probabilities (importance weights)
389
+ # Importance weights are calculated as exp(log_pi_theta - log_pi_theta_old)
390
+ importance_weights = torch.exp(logps - old_logps)
391
+
392
+ # Clip the importance weights to the range [1 - epsilon, 1 + epsilon]
393
+ importance_weights_clipped = torch.clamp(importance_weights, 1 - epsilon, 1 + epsilon)
394
+
395
+ # Create a completion mask. It checks which positions are valid based on logits_to_keep
396
+ completion_mask = torch.arange(logits_to_keep, device=logps.device)[None, :] >= 0
397
+
398
+ # Combine the completion mask and padding mask
399
+ completion_mask = completion_mask & pad_mask # Ensure matching shape
400
+
401
+ # Add an extra dimension to advantages to match the shape for element - wise multiplication
402
+ advantages = advantages.unsqueeze(1)
403
+
404
+ # Calculate the per - token loss. It takes the minimum of the unclipped and clipped importance weights
405
+ # and subtracts the KL divergence term weighted by beta, then multiplies by the completion mask
406
+ token_loss = -(torch.min(advantages * importance_weights, advantages *
407
+ importance_weights_clipped) - beta * per_token_kl) * completion_mask
408
+
409
+ # Calculate the final loss by summing the token losses and normalizing by the number of valid tokens
410
+ loss = -token_loss.sum() / completion_mask.sum()
411
+
412
+ return loss
code/flash-linear-attention/fla/modules/l2norm.py ADDED
@@ -0,0 +1,287 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import torch.nn as nn
6
+ import triton
7
+ import triton.language as tl
8
+
9
+ from fla.utils import autotune_cache_kwargs, input_guard, is_amd
10
+
11
+ BT_LIST = [8, 16, 32, 64, 128]
12
+ NUM_WARPS_AUTOTUNE = [1, 2, 4, 8, 16] if is_amd else [1, 2, 4, 8, 16, 32]
13
+
14
+
15
+ @triton.autotune(
16
+ configs=[
17
+ triton.Config({}, num_warps=num_warps)
18
+ for num_warps in NUM_WARPS_AUTOTUNE
19
+ ],
20
+ key=['D'],
21
+ **autotune_cache_kwargs,
22
+ )
23
+ @triton.jit
24
+ def l2norm_fwd_kernel1(
25
+ x,
26
+ y,
27
+ rstd,
28
+ eps,
29
+ D,
30
+ BD: tl.constexpr,
31
+ ):
32
+ i_t = tl.program_id(0)
33
+ x += i_t * D
34
+ y += i_t * D
35
+ # Compute mean and variance
36
+ cols = tl.arange(0, BD)
37
+ mask = cols < D
38
+
39
+ b_x = tl.load(x + cols, mask=mask, other=0.0).to(tl.float32)
40
+ b_rstd = 1 / tl.sqrt(tl.sum(b_x * b_x) + eps)
41
+ b_y = b_x * b_rstd
42
+ tl.store(y + cols, b_y, mask=mask)
43
+ tl.store(rstd + i_t, b_rstd)
44
+
45
+
46
+ @triton.autotune(
47
+ configs=[
48
+ triton.Config({}, num_warps=num_warps)
49
+ for num_warps in NUM_WARPS_AUTOTUNE
50
+ ],
51
+ key=['D'],
52
+ **autotune_cache_kwargs,
53
+ )
54
+ @triton.jit
55
+ def l2norm_bwd_kernel1(
56
+ y,
57
+ rstd,
58
+ dy,
59
+ dx,
60
+ eps,
61
+ D,
62
+ BD: tl.constexpr,
63
+ ):
64
+ i_t = tl.program_id(0)
65
+ y += i_t * D
66
+ dx += i_t * D
67
+ dy += i_t * D
68
+
69
+ cols = tl.arange(0, BD)
70
+ mask = cols < D
71
+ b_y = tl.load(y + cols, mask=mask, other=0.0).to(tl.float32)
72
+ b_rstd = tl.load(rstd + i_t).to(tl.float32)
73
+ b_dy = tl.load(dy + cols, mask=mask, other=0.0).to(tl.float32)
74
+ b_dx = b_dy * b_rstd - tl.sum(b_dy * b_y) * b_y * b_rstd
75
+ tl.store(dx + cols, b_dx, mask=mask)
76
+
77
+
78
+ @triton.autotune(
79
+ configs=[
80
+ triton.Config({'BT': BT}, num_warps=num_warps)
81
+ for num_warps in [1, 2, 4, 8, 16]
82
+ for BT in BT_LIST
83
+ ],
84
+ key=['D', 'NB'],
85
+ **autotune_cache_kwargs,
86
+ )
87
+ @triton.jit
88
+ def l2norm_fwd_kernel(
89
+ x,
90
+ y,
91
+ rstd,
92
+ eps,
93
+ T: tl.constexpr,
94
+ D: tl.constexpr,
95
+ BD: tl.constexpr,
96
+ NB: tl.constexpr,
97
+ BT: tl.constexpr,
98
+ ):
99
+ i_t = tl.program_id(0)
100
+ p_x = tl.make_block_ptr(x, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
101
+ p_y = tl.make_block_ptr(y, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
102
+ p_rstd = tl.make_block_ptr(rstd, (T,), (1,), (i_t * BT,), (BT,), (0,))
103
+
104
+ b_x = tl.load(p_x, boundary_check=(0, 1)).to(tl.float32)
105
+ b_rstd = 1 / tl.sqrt(tl.sum(b_x * b_x, 1) + eps)
106
+ b_y = b_x * b_rstd[:, None]
107
+
108
+ tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1))
109
+ tl.store(p_rstd, b_rstd.to(p_rstd.dtype.element_ty), boundary_check=(0,))
110
+
111
+
112
+ @triton.autotune(
113
+ configs=[
114
+ triton.Config({'BT': BT}, num_warps=num_warps)
115
+ for num_warps in [1, 2, 4, 8, 16]
116
+ for BT in BT_LIST
117
+ ],
118
+ key=['D', 'NB'],
119
+ **autotune_cache_kwargs,
120
+ )
121
+ @triton.jit
122
+ def l2norm_bwd_kernel(
123
+ y,
124
+ rstd,
125
+ dy,
126
+ dx,
127
+ eps,
128
+ T: tl.constexpr,
129
+ D: tl.constexpr,
130
+ BD: tl.constexpr,
131
+ NB: tl.constexpr,
132
+ BT: tl.constexpr,
133
+ ):
134
+ i_t = tl.program_id(0)
135
+ p_y = tl.make_block_ptr(y, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
136
+ p_rstd = tl.make_block_ptr(rstd, (T,), (1,), (i_t * BT,), (BT,), (0,))
137
+ p_dy = tl.make_block_ptr(dy, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
138
+ p_dx = tl.make_block_ptr(dx, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
139
+
140
+ b_y = tl.load(p_y, boundary_check=(0, 1)).to(tl.float32)
141
+ b_rstd = tl.load(p_rstd, boundary_check=(0,)).to(tl.float32)
142
+ b_dy = tl.load(p_dy, boundary_check=(0, 1)).to(tl.float32)
143
+ b_dx = b_dy * b_rstd[:, None] - tl.sum(b_dy * b_y, 1)[:, None] * b_y * b_rstd[:, None]
144
+ tl.store(p_dx, b_dx.to(p_dx.dtype.element_ty), boundary_check=(0, 1))
145
+
146
+
147
+ def l2norm_fwd(
148
+ x: torch.Tensor,
149
+ eps: float = 1e-6,
150
+ output_dtype: torch.dtype | None = None,
151
+ ):
152
+ x_shape_og = x.shape
153
+ x = x.view(-1, x.shape[-1])
154
+ # allocate output
155
+ if output_dtype is None:
156
+ y = torch.empty_like(x)
157
+ else:
158
+ y = torch.empty_like(x, dtype=output_dtype)
159
+ assert y.stride(-1) == 1
160
+ T, D = x.shape[0], x.shape[-1]
161
+ # Less than 64KB per feature: enqueue fused kernel
162
+ MAX_FUSED_SIZE = 65536 // x.element_size()
163
+ BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
164
+ if D > BD:
165
+ raise RuntimeError("This layer doesn't support feature dim >= 64KB.")
166
+
167
+ rstd = torch.empty((T,), dtype=torch.float32, device=x.device)
168
+ if D <= 512:
169
+ NB = triton.cdiv(T, 2048)
170
+ def grid(meta): return (triton.cdiv(T, meta['BT']), )
171
+ l2norm_fwd_kernel[grid](
172
+ x=x,
173
+ y=y,
174
+ rstd=rstd,
175
+ eps=eps,
176
+ T=T,
177
+ D=D,
178
+ BD=BD,
179
+ NB=NB,
180
+ )
181
+ else:
182
+ l2norm_fwd_kernel1[(T,)](
183
+ x=x,
184
+ y=y,
185
+ rstd=rstd,
186
+ eps=eps,
187
+ D=D,
188
+ BD=BD,
189
+ )
190
+ return y.view(x_shape_og), rstd.view(x_shape_og[:-1])
191
+
192
+
193
+ def l2norm_bwd(
194
+ y: torch.Tensor,
195
+ rstd: torch.Tensor,
196
+ dy: torch.Tensor,
197
+ eps: float = 1e-6,
198
+ ):
199
+ y_shape_og = y.shape
200
+ y = y.view(-1, dy.shape[-1])
201
+ dy = dy.view(-1, dy.shape[-1])
202
+ assert dy.shape == y.shape
203
+ # allocate output
204
+ dx = torch.empty_like(y)
205
+ T, D = y.shape[0], y.shape[-1]
206
+ # Less than 64KB per feature: enqueue fused kernel
207
+ MAX_FUSED_SIZE = 65536 // y.element_size()
208
+ BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
209
+ if D > BD:
210
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
211
+
212
+ if D <= 512:
213
+ NB = triton.cdiv(T, 2048)
214
+ def grid(meta): return (triton.cdiv(T, meta['BT']), )
215
+ l2norm_bwd_kernel[grid](
216
+ y=y,
217
+ rstd=rstd,
218
+ dy=dy,
219
+ dx=dx,
220
+ eps=eps,
221
+ T=T,
222
+ D=D,
223
+ BD=BD,
224
+ NB=NB,
225
+ )
226
+ else:
227
+ l2norm_bwd_kernel1[(T,)](
228
+ y=y,
229
+ rstd=rstd,
230
+ dy=dy,
231
+ dx=dx,
232
+ eps=eps,
233
+ D=D,
234
+ BD=BD,
235
+ )
236
+
237
+ return dx.view(y_shape_og)
238
+
239
+
240
+ class L2NormFunction(torch.autograd.Function):
241
+
242
+ @staticmethod
243
+ @input_guard
244
+ def forward(
245
+ ctx,
246
+ x,
247
+ eps=1e-6,
248
+ output_dtype=None,
249
+ ):
250
+ y, rstd = l2norm_fwd(x, eps, output_dtype)
251
+ ctx.eps = eps
252
+ ctx.x_dtype = x.dtype
253
+ ctx.save_for_backward(y, rstd)
254
+ return y
255
+
256
+ @staticmethod
257
+ @input_guard
258
+ def backward(ctx, dy):
259
+ y, rstd = ctx.saved_tensors
260
+ dx = l2norm_bwd(y, rstd, dy, ctx.eps)
261
+ return dx, None, None
262
+
263
+
264
+ def l2norm(
265
+ x: torch.Tensor,
266
+ eps: float = 1e-6,
267
+ output_dtype: torch.dtype | None = None,
268
+ ) -> torch.Tensor:
269
+ return L2NormFunction.apply(x, eps, output_dtype)
270
+
271
+
272
+ l2_norm = l2norm
273
+
274
+
275
+ class L2Norm(nn.Module):
276
+
277
+ def __init__(
278
+ self,
279
+ eps: float = 1e-6,
280
+ output_dtype: torch.dtype | None = None,
281
+ ):
282
+ super().__init__()
283
+ self.eps = eps
284
+ self.output_dtype = output_dtype
285
+
286
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
287
+ return l2norm(x, self.eps, self.output_dtype)
code/flash-linear-attention/fla/modules/l2warp.py ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import torch
3
+
4
+
5
+ class L2Wrap(torch.autograd.Function):
6
+ r"""
7
+ This class of penalty prevents the model from becoming overconfident,
8
+ thereby mitigating precision loss in BF16.
9
+
10
+ This version is memory-optimized by not storing the full logits tensor.
11
+ """
12
+ @staticmethod
13
+ def forward(ctx, loss, logits, l2_penalty_factor=1e-4):
14
+ """
15
+ Forward pass for L2 penalty.
16
+ Args:
17
+ loss (torch.Tensor): The loss tensor.
18
+ logits (torch.Tensor): Shape[B, T, V] The logits tensor.
19
+ l2_penalty_factor (float): The factor for L2 penalty.
20
+ """
21
+ maxx, ids = torch.max(logits, dim=-1, keepdim=True)
22
+ ctx.logits_shape = logits.shape
23
+ factor = l2_penalty_factor / (logits.shape[0] * logits.shape[1])
24
+ maxx = maxx * factor
25
+ ctx.save_for_backward(maxx, ids)
26
+ return loss
27
+
28
+ @staticmethod
29
+ def backward(ctx, grad_output):
30
+ maxx, ids = ctx.saved_tensors
31
+ glogits = torch.zeros(ctx.logits_shape, device=grad_output.device,
32
+ dtype=grad_output.dtype)
33
+ glogits.scatter_(-1, ids, maxx)
34
+ return grad_output, glogits, None
35
+
36
+
37
+ l2_warp = L2Wrap.apply
code/flash-linear-attention/fla/modules/layernorm.py ADDED
@@ -0,0 +1,1444 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+ # Copyright (c) 2023, Tri Dao
4
+ # https://github.com/state-spaces/mamba/blob/fb7b5310fa865dbd62aa059b1e26f2b431363e2a/mamba_ssm/ops/triton/layernorm.py
5
+ # Implement residual + layer_norm / rms_norm.
6
+
7
+ # Based on the Triton LayerNorm tutorial: https://triton-lang.org/main/getting-started/tutorials/05-layer-norm.html
8
+ # For the backward pass, we keep weight_grad and bias_grad in registers and accumulate.
9
+ # This is faster for dimensions up to 8k, but after that it's much slower due to register spilling.
10
+ # The models we train have hidden dim up to 8k anyway (e.g. Llama 70B), so this is fine.
11
+
12
+ from __future__ import annotations
13
+
14
+ from functools import partial
15
+
16
+ import torch
17
+ import torch.nn as nn
18
+ import torch.nn.functional as F
19
+ import triton
20
+ import triton.language as tl
21
+ from einops import rearrange
22
+ try:
23
+ from torch.distributed import DeviceMesh
24
+ from torch.distributed.tensor import Replicate, Shard, distribute_module
25
+ from torch.distributed.tensor.parallel import ParallelStyle
26
+ except ImportError:
27
+ DeviceMesh = None
28
+ Replicate = None
29
+ Shard = None
30
+ distribute_module = None
31
+ class ParallelStyle:
32
+ pass
33
+
34
+ from fla.utils import autotune_cache_kwargs, get_multiprocessor_count, input_guard
35
+
36
+ try:
37
+ from torch.distributed.tensor import DTensor
38
+ except (ImportError, AttributeError):
39
+ DTensor = None
40
+
41
+
42
+ def layer_norm_ref(
43
+ x: torch.Tensor,
44
+ weight: torch.Tensor,
45
+ bias: torch.Tensor,
46
+ residual: torch.Tensor = None,
47
+ eps: float = 1e-5,
48
+ prenorm: bool = False,
49
+ upcast: bool = False,
50
+ ):
51
+ dtype = x.dtype
52
+ if upcast:
53
+ weight = weight.float()
54
+ bias = bias.float() if bias is not None else None
55
+ if upcast:
56
+ x = x.float()
57
+ residual = residual.float() if residual is not None else residual
58
+ if residual is not None:
59
+ x = (x + residual).to(x.dtype)
60
+ out = F.layer_norm(x.to(weight.dtype), x.shape[-1:], weight=weight, bias=bias, eps=eps).to(
61
+ dtype,
62
+ )
63
+ return out if not prenorm else (out, x)
64
+
65
+
66
+ def rms_norm_ref(
67
+ x: torch.Tensor,
68
+ weight: torch.Tensor,
69
+ bias: torch.Tensor,
70
+ residual: torch.Tensor = None,
71
+ eps: float = 1e-5,
72
+ prenorm: bool = False,
73
+ upcast: bool = False,
74
+ ):
75
+ dtype = x.dtype
76
+ if upcast:
77
+ weight = weight.float()
78
+ bias = bias.float() if bias is not None else None
79
+ if upcast:
80
+ x = x.float()
81
+ residual = residual.float() if residual is not None else residual
82
+ if residual is not None:
83
+ x = (x + residual).to(x.dtype)
84
+ rstd = 1 / torch.sqrt((x.square()).mean(dim=-1, keepdim=True) + eps)
85
+ out = (x * rstd * weight) + bias if bias is not None else (x * rstd * weight)
86
+ out = out.to(dtype)
87
+ return out if not prenorm else (out, x)
88
+
89
+
90
+ def group_norm_ref(
91
+ x: torch.Tensor,
92
+ weight: torch.Tensor,
93
+ bias: torch.Tensor,
94
+ num_groups: int,
95
+ residual: torch.Tensor = None,
96
+ eps: float = 1e-5,
97
+ is_rms_norm: bool = False,
98
+ prenorm: bool = False,
99
+ upcast: bool = False,
100
+ ):
101
+ dtype = x.dtype
102
+ if upcast:
103
+ weight = weight.float()
104
+ bias = bias.float() if bias is not None else None
105
+ if upcast:
106
+ x = x.float()
107
+ residual = residual.float() if residual is not None else residual
108
+ if residual is not None:
109
+ x = (x + residual).to(x.dtype)
110
+ residual = x
111
+ x, weight = [
112
+ rearrange(data, "... (g d) -> ... g d", g=num_groups) for data in (x, weight)
113
+ ]
114
+ if bias is not None:
115
+ bias = rearrange(bias, '... (g d) -> ... g d', g=num_groups)
116
+ if not is_rms_norm:
117
+ mean = x.mean(dim=-1, keepdim=True)
118
+ x = x - mean
119
+ rstd = 1 / torch.sqrt((x.square()).mean(dim=-1, keepdim=True) + eps)
120
+ out = (x * rstd * weight) + bias if bias is not None else (x * rstd * weight)
121
+ out = rearrange(out, "... g d -> ... (g d)")
122
+ out = out.to(dtype)
123
+ return out if not prenorm else (out, residual)
124
+
125
+
126
+ class GroupNormRef(nn.Module):
127
+
128
+ def __init__(
129
+ self,
130
+ num_groups: int,
131
+ hidden_size: int,
132
+ elementwise_affine: bool = True,
133
+ bias: bool = False,
134
+ eps: float = 1e-5,
135
+ is_rms_norm: bool = False,
136
+ ) -> GroupNormRef:
137
+ super().__init__()
138
+
139
+ if hidden_size % num_groups != 0:
140
+ raise ValueError('num_channels must be divisible by num_groups')
141
+
142
+ self.num_groups = num_groups
143
+ self.hidden_size = hidden_size
144
+ self.elementwise_affine = elementwise_affine
145
+ self.eps = eps
146
+ self.is_rms_norm = is_rms_norm
147
+
148
+ self.register_parameter("weight", None)
149
+ self.register_parameter("bias", None)
150
+ if elementwise_affine:
151
+ self.weight = nn.Parameter(torch.empty(hidden_size))
152
+ if bias:
153
+ self.bias = nn.Parameter(torch.empty(hidden_size))
154
+
155
+ self.reset_parameters()
156
+
157
+ def reset_parameters(self):
158
+ if self.elementwise_affine:
159
+ nn.init.ones_(self.weight)
160
+ if self.bias is not None:
161
+ nn.init.zeros_(self.bias)
162
+
163
+ def __repr__(self) -> str:
164
+ s = f"{self.__class__.__name__}({self.num_groups}, {self.hidden_size}"
165
+ if not self.elementwise_affine:
166
+ s += f", elementwise_affine={self.elementwise_affine}"
167
+ if self.is_rms_norm:
168
+ s += f", is_rms_norm={self.is_rms_norm}"
169
+ s += f", eps={self.eps}"
170
+ s += ")"
171
+ return s
172
+
173
+ def forward(self, x, residual=None, prenorm=False):
174
+ return group_norm_ref(
175
+ x,
176
+ self.weight,
177
+ self.bias,
178
+ num_groups=self.num_groups,
179
+ residual=residual,
180
+ eps=self.eps,
181
+ is_rms_norm=self.is_rms_norm,
182
+ prenorm=prenorm,
183
+ upcast=True,
184
+ )
185
+
186
+
187
+ @triton.autotune(
188
+ configs=[
189
+ triton.Config({'BT': BT}, num_warps=num_warps)
190
+ for BT in [32, 64, 128]
191
+ for num_warps in [2, 4, 8]
192
+ ],
193
+ key=['D', 'NB', 'HAS_RESIDUAL', 'STORE_RESIDUAL_OUT', 'IS_RMS_NORM'],
194
+ **autotune_cache_kwargs,
195
+ )
196
+ @triton.jit
197
+ def layer_norm_fwd_kernel(
198
+ x, # pointer to the input
199
+ y, # pointer to the output
200
+ w, # pointer to the weights
201
+ b, # pointer to the biases
202
+ res, # pointer to the res
203
+ res_out, # pointer to the res
204
+ mean, # pointer to the mean
205
+ rstd, # pointer to the 1/std
206
+ eps, # epsilon to avoid division by zero
207
+ T,
208
+ G: tl.constexpr,
209
+ D: tl.constexpr,
210
+ BT: tl.constexpr,
211
+ BD: tl.constexpr,
212
+ NB: tl.constexpr,
213
+ IS_RMS_NORM: tl.constexpr,
214
+ HAS_RESIDUAL: tl.constexpr,
215
+ STORE_RESIDUAL_OUT: tl.constexpr,
216
+ HAS_WEIGHT: tl.constexpr,
217
+ HAS_BIAS: tl.constexpr,
218
+ ):
219
+ i_t = tl.program_id(0)
220
+
221
+ o_t = i_t * BT + tl.arange(0, BT)
222
+ o_g = o_t % G
223
+ o_d = tl.arange(0, BD)
224
+ m_d = o_d < D
225
+
226
+ p_x = tl.make_block_ptr(x, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
227
+ b_x = tl.load(p_x, boundary_check=(0, 1)).to(tl.float32)
228
+ if HAS_RESIDUAL:
229
+ p_res = tl.make_block_ptr(res, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
230
+ b_x += tl.load(p_res, boundary_check=(0, 1)).to(tl.float32)
231
+ if STORE_RESIDUAL_OUT:
232
+ p_res_out = tl.make_block_ptr(res_out, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
233
+ tl.store(p_res_out, b_x.to(p_res_out.dtype.element_ty), boundary_check=(0, 1))
234
+ if not IS_RMS_NORM:
235
+ b_mean = tl.sum(b_x, axis=1) / D
236
+ p_mean = tl.make_block_ptr(mean, (T,), (1,), (i_t * BT,), (BT,), (0,))
237
+ tl.store(p_mean, b_mean.to(p_mean.dtype.element_ty), boundary_check=(0,))
238
+ b_xbar = tl.where(m_d[None, :], b_x - b_mean[:, None], 0.0)
239
+ b_var = tl.sum(b_xbar * b_xbar, axis=1) / D
240
+ else:
241
+ b_xbar = tl.where(m_d[None, :], b_x, 0.0)
242
+ b_var = tl.sum(b_xbar * b_xbar, axis=1) / D
243
+ b_rstd = 1 / tl.sqrt(b_var + eps)
244
+
245
+ p_rstd = tl.make_block_ptr(rstd, (T,), (1,), (i_t * BT,), (BT,), (0,))
246
+ tl.store(p_rstd, b_rstd.to(p_rstd.dtype.element_ty), boundary_check=(0,))
247
+
248
+ if HAS_WEIGHT:
249
+ b_w = tl.load(w + o_g[:, None] * D + o_d[None, :], mask=m_d[None, :]).to(tl.float32)
250
+ if HAS_BIAS:
251
+ b_b = tl.load(b + o_g[:, None] * D + o_d[None, :], mask=m_d[None, :]).to(tl.float32)
252
+ b_x_hat = (b_x - b_mean[:, None]) * b_rstd[:, None] if not IS_RMS_NORM else b_x * b_rstd[:, None]
253
+ b_y = b_x_hat * b_w if HAS_WEIGHT else b_x_hat
254
+ if HAS_BIAS:
255
+ b_y = b_y + b_b
256
+
257
+ # Write output
258
+ p_y = tl.make_block_ptr(y, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
259
+ tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1))
260
+
261
+
262
+ @triton.autotune(
263
+ configs=[
264
+ triton.Config({}, num_warps=num_warps)
265
+ for num_warps in [2, 4, 8, 16]
266
+ ],
267
+ key=['D', 'HAS_RESIDUAL', 'STORE_RESIDUAL_OUT', 'IS_RMS_NORM'],
268
+ **autotune_cache_kwargs,
269
+ )
270
+ @triton.jit
271
+ def layer_norm_fwd_kernel1(
272
+ x, # pointer to the input
273
+ y, # pointer to the output
274
+ w, # pointer to the weights
275
+ b, # pointer to the biases
276
+ res, # pointer to the res
277
+ res_out, # pointer to the res
278
+ mean, # pointer to the mean
279
+ rstd, # pointer to the 1/std
280
+ eps, # epsilon to avoid division by zero
281
+ G: tl.constexpr,
282
+ D: tl.constexpr,
283
+ BD: tl.constexpr,
284
+ IS_RMS_NORM: tl.constexpr,
285
+ HAS_RESIDUAL: tl.constexpr,
286
+ STORE_RESIDUAL_OUT: tl.constexpr,
287
+ HAS_WEIGHT: tl.constexpr,
288
+ HAS_BIAS: tl.constexpr,
289
+ ):
290
+ i_t = tl.program_id(0)
291
+ i_g = i_t % G
292
+
293
+ x += i_t * D
294
+ y += i_t * D
295
+ if HAS_RESIDUAL:
296
+ res += i_t * D
297
+ if STORE_RESIDUAL_OUT:
298
+ res_out += i_t * D
299
+
300
+ o_d = tl.arange(0, BD)
301
+ m_d = o_d < D
302
+ b_x = tl.load(x + o_d, mask=m_d, other=0.0).to(tl.float32)
303
+ if HAS_RESIDUAL:
304
+ b_x += tl.load(res + o_d, mask=m_d, other=0.0).to(tl.float32)
305
+ if STORE_RESIDUAL_OUT:
306
+ tl.store(res_out + o_d, b_x, mask=m_d)
307
+ if not IS_RMS_NORM:
308
+ b_mean = tl.sum(b_x, axis=0) / D
309
+ tl.store(mean + i_t, b_mean)
310
+ b_xbar = tl.where(m_d, b_x - b_mean, 0.0)
311
+ b_var = tl.sum(b_xbar * b_xbar, axis=0) / D
312
+ else:
313
+ b_xbar = tl.where(m_d, b_x, 0.0)
314
+ b_var = tl.sum(b_xbar * b_xbar, axis=0) / D
315
+ b_rstd = 1 / tl.sqrt(b_var + eps)
316
+ tl.store(rstd + i_t, b_rstd)
317
+
318
+ if HAS_WEIGHT:
319
+ b_w = tl.load(w + i_g * D + o_d, mask=m_d).to(tl.float32)
320
+ if HAS_BIAS:
321
+ b_b = tl.load(b + i_g * D + o_d, mask=m_d).to(tl.float32)
322
+ b_x_hat = (b_x - b_mean) * b_rstd if not IS_RMS_NORM else b_x * b_rstd
323
+ b_y = b_x_hat * b_w if HAS_WEIGHT else b_x_hat
324
+ if HAS_BIAS:
325
+ b_y = b_y + b_b
326
+
327
+ # Write output
328
+ tl.store(y + o_d, b_y, mask=m_d)
329
+
330
+
331
+ @triton.heuristics({
332
+ 'RECOMPUTE_OUTPUT': lambda args: args['y'] is not None,
333
+ })
334
+ @triton.autotune(
335
+ configs=[
336
+ triton.Config({'BT': BT}, num_warps=num_warps)
337
+ for BT in [32, 64]
338
+ for num_warps in [2, 4, 8]
339
+ ],
340
+ key=['D', 'NB', 'HAS_DRESIDUAL', 'STORE_DRESIDUAL', 'IS_RMS_NORM'],
341
+ **autotune_cache_kwargs,
342
+ )
343
+ @triton.jit
344
+ def layer_norm_bwd_kernel(
345
+ x, # pointer to the input
346
+ w, # pointer to the weights
347
+ b, # pointer to the biases
348
+ y, # pointer to the output to be recomputed
349
+ dy, # pointer to the output gradient
350
+ dx, # pointer to the input gradient
351
+ dw, # pointer to the partial sum of weights gradient
352
+ db, # pointer to the partial sum of biases gradient
353
+ dres,
354
+ dres_in,
355
+ mean,
356
+ rstd,
357
+ T,
358
+ G: tl.constexpr,
359
+ D: tl.constexpr,
360
+ BS: tl.constexpr,
361
+ BT: tl.constexpr,
362
+ BD: tl.constexpr,
363
+ NB: tl.constexpr,
364
+ GS: tl.constexpr,
365
+ IS_RMS_NORM: tl.constexpr,
366
+ HAS_DRESIDUAL: tl.constexpr,
367
+ STORE_DRESIDUAL: tl.constexpr,
368
+ HAS_WEIGHT: tl.constexpr,
369
+ HAS_BIAS: tl.constexpr,
370
+ RECOMPUTE_OUTPUT: tl.constexpr,
371
+ ):
372
+ i_s = tl.program_id(0)
373
+ i_g, i_sg = i_s // GS, i_s % GS
374
+
375
+ o_d = tl.arange(0, BD)
376
+ m_d = o_d < D
377
+ if HAS_WEIGHT:
378
+ b_w = tl.load(w + i_g * D + o_d, mask=m_d).to(tl.float32)
379
+ b_dw = tl.zeros((BT, BD), dtype=tl.float32)
380
+ if HAS_BIAS:
381
+ b_b = tl.load(b + i_g * D + o_d, mask=m_d, other=0.0).to(tl.float32)
382
+ b_db = tl.zeros((BT, BD), dtype=tl.float32)
383
+
384
+ T = min(i_sg * BS + BS, T // G)
385
+ for i_t in range(i_sg * BS, T, BT):
386
+ p_x = tl.make_block_ptr(x + i_g * D, (T, D), (G*D, 1), (i_t, 0), (BT, BD), (1, 0))
387
+ p_dy = tl.make_block_ptr(dy + i_g * D, (T, D), (G*D, 1), (i_t, 0), (BT, BD), (1, 0))
388
+ p_dx = tl.make_block_ptr(dx + i_g * D, (T, D), (G*D, 1), (i_t, 0), (BT, BD), (1, 0))
389
+ # [BT, BD]
390
+ b_x = tl.load(p_x, boundary_check=(0, 1)).to(tl.float32)
391
+ b_dy = tl.load(p_dy, boundary_check=(0, 1)).to(tl.float32)
392
+
393
+ if not IS_RMS_NORM:
394
+ p_mean = tl.make_block_ptr(mean + i_g, (T,), (G,), (i_t,), (BT,), (0,))
395
+ b_mean = tl.load(p_mean, boundary_check=(0,))
396
+ p_rstd = tl.make_block_ptr(rstd + i_g, (T,), (G,), (i_t,), (BT,), (0,))
397
+ b_rstd = tl.load(p_rstd, boundary_check=(0,))
398
+ # Compute dx
399
+ b_xhat = (b_x - b_mean[:, None]) * b_rstd[:, None] if not IS_RMS_NORM else b_x * b_rstd[:, None]
400
+ b_xhat = tl.where(m_d[None, :], b_xhat, 0.0)
401
+
402
+ b_y = b_xhat * b_w[None, :] if HAS_WEIGHT else b_xhat
403
+ if HAS_BIAS:
404
+ b_y = b_y + b_b[None, :]
405
+ if RECOMPUTE_OUTPUT:
406
+ p_y = tl.make_block_ptr(y + i_g * D, (T, D), (G*D, 1), (i_t, 0), (BT, BD), (1, 0))
407
+ tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1))
408
+
409
+ b_wdy = b_dy
410
+
411
+ if HAS_WEIGHT or HAS_BIAS:
412
+ m_t = (i_t + tl.arange(0, BT)) < T
413
+ if HAS_WEIGHT:
414
+ b_wdy = b_dy * b_w
415
+ b_dw += tl.where(m_t[:, None], b_dy * b_xhat, 0.0)
416
+ if HAS_BIAS:
417
+ b_db += tl.where(m_t[:, None], b_dy, 0.0)
418
+ if not IS_RMS_NORM:
419
+ b_c1 = tl.sum(b_xhat * b_wdy, axis=1) / D
420
+ b_c2 = tl.sum(b_wdy, axis=1) / D
421
+ b_dx = (b_wdy - (b_xhat * b_c1[:, None] + b_c2[:, None])) * b_rstd[:, None]
422
+ else:
423
+ b_c1 = tl.sum(b_xhat * b_wdy, axis=1) / D
424
+ b_dx = (b_wdy - b_xhat * b_c1[:, None]) * b_rstd[:, None]
425
+ if HAS_DRESIDUAL:
426
+ p_dres = tl.make_block_ptr(dres + i_g * D, (T, D), (G*D, 1), (i_t, 0), (BT, BD), (1, 0))
427
+ b_dres = tl.load(p_dres, boundary_check=(0, 1)).to(tl.float32)
428
+ b_dx += b_dres
429
+ # Write dx
430
+ if STORE_DRESIDUAL:
431
+ p_dres_in = tl.make_block_ptr(dres_in + i_g * D, (T, D), (G*D, 1), (i_t, 0), (BT, BD), (1, 0))
432
+ tl.store(p_dres_in, b_dx.to(p_dres_in.dtype.element_ty), boundary_check=(0, 1))
433
+
434
+ tl.store(p_dx, b_dx.to(p_dx.dtype.element_ty), boundary_check=(0, 1))
435
+
436
+ if HAS_WEIGHT:
437
+ tl.store(dw + i_s * D + o_d, tl.sum(b_dw, axis=0), mask=m_d)
438
+ if HAS_BIAS:
439
+ tl.store(db + i_s * D + o_d, tl.sum(b_db, axis=0), mask=m_d)
440
+
441
+
442
+ @triton.heuristics({
443
+ 'RECOMPUTE_OUTPUT': lambda args: args['y'] is not None,
444
+ })
445
+ @triton.autotune(
446
+ configs=[
447
+ triton.Config({}, num_warps=num_warps)
448
+ for num_warps in [2, 4, 8]
449
+ ],
450
+ key=['D', 'HAS_DRESIDUAL', 'STORE_DRESIDUAL', 'IS_RMS_NORM'],
451
+ **autotune_cache_kwargs,
452
+ )
453
+ @triton.jit
454
+ def layer_norm_bwd_kernel1(
455
+ x, # pointer to the input
456
+ w, # pointer to the weights
457
+ b, # pointer to the biases
458
+ y, # pointer to the output to be recomputed
459
+ dy, # pointer to the output gradient
460
+ dx, # pointer to the input gradient
461
+ dw, # pointer to the partial sum of weights gradient
462
+ db, # pointer to the partial sum of biases gradient
463
+ dres,
464
+ dres_in,
465
+ mean,
466
+ rstd,
467
+ T,
468
+ G: tl.constexpr,
469
+ D: tl.constexpr,
470
+ BS: tl.constexpr,
471
+ BD: tl.constexpr,
472
+ GS: tl.constexpr,
473
+ IS_RMS_NORM: tl.constexpr,
474
+ HAS_DRESIDUAL: tl.constexpr,
475
+ STORE_DRESIDUAL: tl.constexpr,
476
+ HAS_WEIGHT: tl.constexpr,
477
+ HAS_BIAS: tl.constexpr,
478
+ RECOMPUTE_OUTPUT: tl.constexpr,
479
+ ):
480
+ i_s = tl.program_id(0)
481
+ i_g, i_sg = i_s // GS, i_s % GS
482
+
483
+ o_d = tl.arange(0, BD)
484
+ mask = o_d < D
485
+
486
+ if HAS_WEIGHT:
487
+ b_w = tl.load(w + i_g * D + o_d, mask=mask).to(tl.float32)
488
+ b_dw = tl.zeros((BD,), dtype=tl.float32)
489
+ if RECOMPUTE_OUTPUT and HAS_BIAS:
490
+ b_b = tl.load(b + i_g * D + o_d, mask=mask, other=0.0).to(tl.float32)
491
+ if HAS_BIAS:
492
+ b_db = tl.zeros((BD,), dtype=tl.float32)
493
+
494
+ for i_t in range(i_sg * BS * G + i_g, min((i_sg * BS + BS) * G + i_g, T), G):
495
+ b_x = tl.load(x + i_t * D + o_d, mask=mask, other=0).to(tl.float32)
496
+ b_dy = tl.load(dy + i_t * D + o_d, mask=mask, other=0).to(tl.float32)
497
+
498
+ if not IS_RMS_NORM:
499
+ b_mean = tl.load(mean + i_t)
500
+ b_rstd = tl.load(rstd + i_t)
501
+ # Compute dx
502
+ b_xhat = (b_x - b_mean) * b_rstd if not IS_RMS_NORM else b_x * b_rstd
503
+ b_xhat = tl.where(mask, b_xhat, 0.0)
504
+ if RECOMPUTE_OUTPUT:
505
+ b_y = b_xhat * b_w if HAS_WEIGHT else b_xhat
506
+ if HAS_BIAS:
507
+ b_y = b_y + b_b
508
+ tl.store(y + i_t * D + o_d, b_y, mask=mask)
509
+ b_wdy = b_dy
510
+ if HAS_WEIGHT:
511
+ b_wdy = b_dy * b_w
512
+ b_dw += b_dy * b_xhat
513
+ if HAS_BIAS:
514
+ b_db += b_dy
515
+ if not IS_RMS_NORM:
516
+ b_c1 = tl.sum(b_xhat * b_wdy, axis=0) / D
517
+ b_c2 = tl.sum(b_wdy, axis=0) / D
518
+ b_dx = (b_wdy - (b_xhat * b_c1 + b_c2)) * b_rstd
519
+ else:
520
+ b_c1 = tl.sum(b_xhat * b_wdy, axis=0) / D
521
+ b_dx = (b_wdy - b_xhat * b_c1) * b_rstd
522
+ if HAS_DRESIDUAL:
523
+ b_dres = tl.load(dres + i_t * D + o_d, mask=mask, other=0).to(tl.float32)
524
+ b_dx += b_dres
525
+ # Write dx
526
+ b_dx = tl.cast(b_dx, dtype=dx.dtype.element_ty, fp_downcast_rounding='rtne')
527
+ if STORE_DRESIDUAL:
528
+ tl.store(dres_in + i_t * D + o_d, b_dx, mask=mask)
529
+ tl.store(dx + i_t * D + o_d, b_dx, mask=mask)
530
+
531
+ if HAS_WEIGHT:
532
+ tl.store(dw + i_s * D + o_d, b_dw, mask=mask)
533
+ if HAS_BIAS:
534
+ tl.store(db + i_s * D + o_d, b_db, mask=mask)
535
+
536
+
537
+ def layer_norm_fwd(
538
+ x: torch.Tensor,
539
+ weight: torch.Tensor,
540
+ bias: torch.Tensor,
541
+ eps: float = 1e-5,
542
+ residual: torch.Tensor = None,
543
+ out_dtype: torch.dtype = None,
544
+ residual_dtype: torch.dtype = None,
545
+ is_rms_norm: bool = False,
546
+ num_groups: int = 1,
547
+ ):
548
+ if residual is not None:
549
+ residual_dtype = residual.dtype
550
+ T, D, G = *x.shape, num_groups
551
+ if residual is not None:
552
+ assert residual.shape == (T, D)
553
+ if weight is not None:
554
+ assert weight.shape == (G * D,)
555
+ if bias is not None:
556
+ assert bias.shape == (G * D,)
557
+ # allocate output
558
+ y = torch.empty_like(x, dtype=x.dtype if out_dtype is None else out_dtype)
559
+ if residual is not None or (residual_dtype is not None and residual_dtype != x.dtype):
560
+ res_out = torch.empty(T, D, device=x.device, dtype=residual_dtype)
561
+ else:
562
+ res_out = None
563
+ mean = torch.empty((T,), dtype=torch.float, device=x.device) if not is_rms_norm else None
564
+ rstd = torch.empty((T,), dtype=torch.float, device=x.device)
565
+ # Less than 64KB per feature: enqueue fused kernel
566
+ MAX_FUSED_SIZE = 65536 // x.element_size()
567
+ BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
568
+ if D > BD:
569
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
570
+ # heuristics for number of warps
571
+
572
+ if D <= 512:
573
+ NB = triton.cdiv(T, 2048)
574
+ def grid(meta): return (triton.cdiv(T, meta['BT']), )
575
+ layer_norm_fwd_kernel[grid](
576
+ x,
577
+ y,
578
+ weight,
579
+ bias,
580
+ residual,
581
+ res_out,
582
+ mean,
583
+ rstd,
584
+ eps,
585
+ T=T,
586
+ G=G,
587
+ D=D,
588
+ BD=BD,
589
+ NB=NB,
590
+ IS_RMS_NORM=is_rms_norm,
591
+ HAS_RESIDUAL=residual is not None,
592
+ STORE_RESIDUAL_OUT=res_out is not None,
593
+ HAS_WEIGHT=weight is not None,
594
+ HAS_BIAS=bias is not None,
595
+ )
596
+ else:
597
+ layer_norm_fwd_kernel1[(T,)](
598
+ x,
599
+ y,
600
+ weight,
601
+ bias,
602
+ residual,
603
+ res_out,
604
+ mean,
605
+ rstd,
606
+ eps,
607
+ G=G,
608
+ D=D,
609
+ BD=BD,
610
+ IS_RMS_NORM=is_rms_norm,
611
+ HAS_RESIDUAL=residual is not None,
612
+ STORE_RESIDUAL_OUT=res_out is not None,
613
+ HAS_WEIGHT=weight is not None,
614
+ HAS_BIAS=bias is not None,
615
+ )
616
+ # res_out is None if residual is None and residual_dtype == input_dtype
617
+ return y, mean, rstd, res_out if res_out is not None else x
618
+
619
+
620
+ def layer_norm_bwd(
621
+ dy: torch.Tensor,
622
+ x: torch.Tensor,
623
+ weight: torch.Tensor,
624
+ bias: torch.Tensor,
625
+ mean: torch.Tensor = None,
626
+ rstd: torch.Tensor = None,
627
+ dres: torch.Tensor = None,
628
+ has_residual: bool = False,
629
+ is_rms_norm: bool = False,
630
+ x_dtype: torch.dtype = None,
631
+ recompute_output: bool = False,
632
+ num_groups: int = 1,
633
+ ):
634
+ T, D, G = *x.shape, num_groups
635
+ assert dy.shape == (T, D)
636
+ if dres is not None:
637
+ assert dres.shape == (T, D)
638
+ if weight is not None:
639
+ assert weight.shape == (G * D,)
640
+ if bias is not None:
641
+ assert bias.shape == (G * D,)
642
+ # allocate output
643
+ dx = torch.empty_like(x) if x_dtype is None else torch.empty(T, D, dtype=x_dtype, device=x.device)
644
+ dres_in = torch.empty_like(x) if has_residual and dx.dtype != x.dtype else None
645
+ y = torch.empty(T, D, dtype=dy.dtype, device=dy.device) if recompute_output else None
646
+
647
+ # Less than 64KB per feature: enqueue fused kernel
648
+ MAX_FUSED_SIZE = 65536 // x.element_size()
649
+ BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
650
+ if D > BD:
651
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
652
+ # each program handles one group only
653
+ NS = triton.cdiv(get_multiprocessor_count(x.device.index), G) * G
654
+ BS = triton.cdiv(T, NS)
655
+ GS = NS // G
656
+
657
+ dw = torch.empty((NS, D), dtype=torch.float, device=weight.device) if weight is not None else None
658
+ db = torch.empty((NS, D), dtype=torch.float, device=bias.device) if bias is not None else None
659
+ grid = (NS,)
660
+
661
+ if D <= 512:
662
+ NB = triton.cdiv(T, 2048)
663
+ layer_norm_bwd_kernel[grid](
664
+ x,
665
+ weight,
666
+ bias,
667
+ y,
668
+ dy,
669
+ dx,
670
+ dw,
671
+ db,
672
+ dres,
673
+ dres_in,
674
+ mean,
675
+ rstd,
676
+ T=T,
677
+ G=G,
678
+ D=D,
679
+ BS=BS,
680
+ BD=BD,
681
+ NB=NB,
682
+ GS=GS,
683
+ IS_RMS_NORM=is_rms_norm,
684
+ HAS_DRESIDUAL=dres is not None,
685
+ STORE_DRESIDUAL=dres_in is not None,
686
+ HAS_WEIGHT=weight is not None,
687
+ HAS_BIAS=bias is not None,
688
+ )
689
+ else:
690
+ layer_norm_bwd_kernel1[grid](
691
+ x,
692
+ weight,
693
+ bias,
694
+ y,
695
+ dy,
696
+ dx,
697
+ dw,
698
+ db,
699
+ dres,
700
+ dres_in,
701
+ mean,
702
+ rstd,
703
+ T=T,
704
+ G=G,
705
+ D=D,
706
+ BS=BS,
707
+ BD=BD,
708
+ GS=GS,
709
+ IS_RMS_NORM=is_rms_norm,
710
+ HAS_DRESIDUAL=dres is not None,
711
+ STORE_DRESIDUAL=dres_in is not None,
712
+ HAS_WEIGHT=weight is not None,
713
+ HAS_BIAS=bias is not None,
714
+ )
715
+ dw = dw.view(G, -1, D).sum(1).to(weight).view_as(weight) if weight is not None else None
716
+ db = db.view(G, -1, D).sum(1).to(bias).view_as(bias) if bias is not None else None
717
+ # Don't need to compute dres_in separately in this case
718
+ if has_residual and dx.dtype == x.dtype:
719
+ dres_in = dx
720
+ return (dx, dw, db, dres_in) if not recompute_output else (dx, dw, db, dres_in, y)
721
+
722
+
723
+ class LayerNormFunction(torch.autograd.Function):
724
+
725
+ @staticmethod
726
+ @input_guard
727
+ def forward(
728
+ ctx,
729
+ x,
730
+ weight,
731
+ bias,
732
+ residual: torch.Tensor = None,
733
+ eps: float = 1e-5,
734
+ prenorm: bool = False,
735
+ residual_in_fp32: bool = False,
736
+ is_rms_norm: bool = False,
737
+ num_groups: int = 1,
738
+ ):
739
+ x_shape_og = x.shape
740
+
741
+ if x.shape[-1] % num_groups != 0:
742
+ raise ValueError('num_channels must be divisible by num_groups')
743
+ # reshape input data into 2D tensor
744
+ x = x.reshape(-1, (x.shape[-1] // num_groups))
745
+ if residual is not None:
746
+ assert residual.shape == x_shape_og
747
+ residual = residual.reshape_as(x)
748
+ residual_dtype = (
749
+ residual.dtype
750
+ if residual is not None
751
+ else (torch.float32 if residual_in_fp32 else None)
752
+ )
753
+ y, mean, rstd, res_out = layer_norm_fwd(
754
+ x,
755
+ weight,
756
+ bias,
757
+ eps,
758
+ residual,
759
+ residual_dtype=residual_dtype,
760
+ is_rms_norm=is_rms_norm,
761
+ num_groups=num_groups,
762
+ )
763
+ ctx.save_for_backward(res_out, weight, bias, mean, rstd)
764
+ ctx.x_shape_og = x_shape_og
765
+ ctx.eps = eps
766
+ ctx.is_rms_norm = is_rms_norm
767
+ ctx.num_groups = num_groups
768
+ ctx.has_residual = residual is not None
769
+ ctx.prenorm = prenorm
770
+ ctx.x_dtype = x.dtype
771
+ y = y.reshape(x_shape_og)
772
+ return y if not prenorm else (y, res_out.reshape(x_shape_og))
773
+
774
+ @staticmethod
775
+ @input_guard
776
+ def backward(ctx, dy, *args):
777
+ x, weight, bias, mean, rstd = ctx.saved_tensors
778
+ dy = dy.reshape(-1, (dy.shape[-1] // ctx.num_groups))
779
+ assert dy.shape == x.shape
780
+ if ctx.prenorm:
781
+ dresidual = args[0]
782
+ dresidual = dresidual.reshape(-1, x.shape[-1])
783
+ assert dresidual.shape == x.shape
784
+ else:
785
+ dresidual = None
786
+ dx, dw, db, dresidual_in = layer_norm_bwd(
787
+ dy,
788
+ x,
789
+ weight,
790
+ bias,
791
+ mean,
792
+ rstd,
793
+ dresidual,
794
+ ctx.has_residual,
795
+ ctx.is_rms_norm,
796
+ x_dtype=ctx.x_dtype,
797
+ num_groups=ctx.num_groups,
798
+ )
799
+ return (
800
+ dx.reshape(ctx.x_shape_og),
801
+ dw,
802
+ db,
803
+ dresidual_in.reshape(ctx.x_shape_og) if ctx.has_residual else None,
804
+ None,
805
+ None,
806
+ None,
807
+ None,
808
+ None,
809
+ )
810
+
811
+
812
+ def layer_norm(
813
+ x: torch.Tensor,
814
+ weight: torch.Tensor,
815
+ bias: torch.Tensor,
816
+ residual: torch.Tensor = None,
817
+ eps: float = 1e-5,
818
+ prenorm: bool = False,
819
+ residual_in_fp32: bool = False,
820
+ is_rms_norm: bool = False,
821
+ ):
822
+ return LayerNormFunction.apply(
823
+ x,
824
+ weight,
825
+ bias,
826
+ residual,
827
+ eps,
828
+ prenorm,
829
+ residual_in_fp32,
830
+ is_rms_norm,
831
+ )
832
+
833
+
834
+ def group_norm(
835
+ x: torch.Tensor,
836
+ weight: torch.Tensor,
837
+ bias: torch.Tensor,
838
+ residual: torch.Tensor = None,
839
+ eps: float = 1e-5,
840
+ prenorm: bool = False,
841
+ residual_in_fp32: bool = False,
842
+ is_rms_norm: bool = False,
843
+ num_groups: int = 1,
844
+ ):
845
+ return LayerNormFunction.apply(
846
+ x,
847
+ weight,
848
+ bias,
849
+ residual,
850
+ eps,
851
+ prenorm,
852
+ residual_in_fp32,
853
+ is_rms_norm,
854
+ num_groups,
855
+ )
856
+
857
+
858
+ def rms_norm(
859
+ x: torch.Tensor,
860
+ weight: torch.Tensor,
861
+ bias: torch.Tensor,
862
+ residual: torch.Tensor = None,
863
+ eps: float = 1e-5,
864
+ prenorm: bool = False,
865
+ residual_in_fp32: bool = False,
866
+ ):
867
+ return LayerNormFunction.apply(
868
+ x,
869
+ weight,
870
+ bias,
871
+ residual,
872
+ eps,
873
+ prenorm,
874
+ residual_in_fp32,
875
+ True,
876
+ )
877
+
878
+
879
+ def layer_norm_linear(
880
+ x: torch.Tensor,
881
+ norm_weight: torch.Tensor,
882
+ norm_bias: torch.Tensor,
883
+ linear_weight: torch.Tensor,
884
+ linear_bias: torch.Tensor,
885
+ residual: torch.Tensor = None,
886
+ eps: float = 1e-5,
887
+ prenorm: bool = False,
888
+ residual_in_fp32: bool = False,
889
+ is_rms_norm: bool = False,
890
+ num_groups: int = 1,
891
+ ):
892
+ return LayerNormLinearFunction.apply(
893
+ x,
894
+ norm_weight,
895
+ norm_bias,
896
+ linear_weight,
897
+ linear_bias,
898
+ residual,
899
+ eps,
900
+ prenorm,
901
+ residual_in_fp32,
902
+ is_rms_norm,
903
+ num_groups,
904
+ )
905
+
906
+
907
+ def rms_norm_linear(
908
+ x: torch.Tensor,
909
+ norm_weight: torch.Tensor,
910
+ norm_bias: torch.Tensor,
911
+ linear_weight: torch.Tensor,
912
+ linear_bias: torch.Tensor,
913
+ residual: torch.Tensor = None,
914
+ eps: float = 1e-5,
915
+ prenorm: bool = False,
916
+ residual_in_fp32: bool = False,
917
+ ):
918
+ return layer_norm_linear(
919
+ x=x,
920
+ norm_weight=norm_weight,
921
+ norm_bias=norm_bias,
922
+ linear_weight=linear_weight,
923
+ linear_bias=linear_bias,
924
+ residual=residual,
925
+ eps=eps,
926
+ prenorm=prenorm,
927
+ residual_in_fp32=residual_in_fp32,
928
+ is_rms_norm=True,
929
+ )
930
+
931
+
932
+ def group_norm_linear(
933
+ x: torch.Tensor,
934
+ norm_weight: torch.Tensor,
935
+ norm_bias: torch.Tensor,
936
+ linear_weight: torch.Tensor,
937
+ linear_bias: torch.Tensor,
938
+ residual: torch.Tensor = None,
939
+ eps: float = 1e-5,
940
+ prenorm: bool = False,
941
+ residual_in_fp32: bool = False,
942
+ is_rms_norm: bool = False,
943
+ num_groups: int = 1,
944
+ ):
945
+ return layer_norm_linear(
946
+ x=x,
947
+ norm_weight=norm_weight,
948
+ norm_bias=norm_bias,
949
+ linear_weight=linear_weight,
950
+ linear_bias=linear_bias,
951
+ residual=residual,
952
+ eps=eps,
953
+ prenorm=prenorm,
954
+ residual_in_fp32=residual_in_fp32,
955
+ is_rms_norm=is_rms_norm,
956
+ num_groups=num_groups,
957
+ )
958
+
959
+
960
+ class LayerNorm(nn.Module):
961
+
962
+ def __init__(
963
+ self,
964
+ hidden_size: int,
965
+ elementwise_affine: bool = True,
966
+ bias: bool = False,
967
+ eps: float = 1e-5,
968
+ ) -> LayerNorm:
969
+ super().__init__()
970
+
971
+ self.hidden_size = hidden_size
972
+ self.elementwise_affine = elementwise_affine
973
+ self.eps = eps
974
+
975
+ self.register_parameter("weight", None)
976
+ self.register_parameter("bias", None)
977
+ if elementwise_affine:
978
+ self.weight = nn.Parameter(torch.empty(hidden_size))
979
+ if bias:
980
+ self.bias = nn.Parameter(torch.empty(hidden_size))
981
+
982
+ self.reset_parameters()
983
+
984
+ def reset_parameters(self):
985
+ if self.elementwise_affine:
986
+ nn.init.ones_(self.weight)
987
+ if self.bias is not None:
988
+ nn.init.zeros_(self.bias)
989
+
990
+ def __repr__(self) -> str:
991
+ s = f"{self.__class__.__name__}({self.hidden_size}"
992
+ if not self.elementwise_affine:
993
+ s += f", elementwise_affine={self.elementwise_affine}"
994
+ s += f", eps={self.eps}"
995
+ s += ")"
996
+ return s
997
+
998
+ def forward(self, x, residual=None, prenorm=False, residual_in_fp32=False):
999
+ return layer_norm(
1000
+ x,
1001
+ self.weight,
1002
+ self.bias,
1003
+ residual=residual,
1004
+ eps=self.eps,
1005
+ prenorm=prenorm,
1006
+ residual_in_fp32=residual_in_fp32,
1007
+ )
1008
+
1009
+
1010
+ class GroupNorm(nn.Module):
1011
+
1012
+ def __init__(
1013
+ self,
1014
+ num_groups: int,
1015
+ hidden_size: int,
1016
+ elementwise_affine: bool = True,
1017
+ bias: bool = False,
1018
+ eps: float = 1e-5,
1019
+ is_rms_norm: bool = False,
1020
+ ) -> GroupNorm:
1021
+ super().__init__()
1022
+
1023
+ if hidden_size % num_groups != 0:
1024
+ raise ValueError('num_channels must be divisible by num_groups')
1025
+
1026
+ self.num_groups = num_groups
1027
+ self.hidden_size = hidden_size
1028
+ self.elementwise_affine = elementwise_affine
1029
+ self.eps = eps
1030
+ self.is_rms_norm = is_rms_norm
1031
+
1032
+ self.register_parameter("weight", None)
1033
+ self.register_parameter("bias", None)
1034
+ if elementwise_affine:
1035
+ self.weight = nn.Parameter(torch.empty(hidden_size))
1036
+ if bias:
1037
+ self.bias = nn.Parameter(torch.empty(hidden_size))
1038
+
1039
+ self.reset_parameters()
1040
+
1041
+ def reset_parameters(self):
1042
+ if self.elementwise_affine:
1043
+ nn.init.ones_(self.weight)
1044
+ if self.bias is not None:
1045
+ nn.init.zeros_(self.bias)
1046
+
1047
+ def __repr__(self) -> str:
1048
+ s = f"{self.__class__.__name__}({self.num_groups}, {self.hidden_size}"
1049
+ if not self.elementwise_affine:
1050
+ s += f", elementwise_affine={self.elementwise_affine}"
1051
+ if self.is_rms_norm:
1052
+ s += f", is_rms_norm={self.is_rms_norm}"
1053
+ s += f", eps={self.eps}"
1054
+ s += ")"
1055
+ return s
1056
+
1057
+ def forward(self, x, residual=None, prenorm=False, residual_in_fp32=False):
1058
+ return group_norm(
1059
+ x,
1060
+ self.weight,
1061
+ self.bias,
1062
+ residual=residual,
1063
+ eps=self.eps,
1064
+ prenorm=prenorm,
1065
+ residual_in_fp32=residual_in_fp32,
1066
+ is_rms_norm=self.is_rms_norm,
1067
+ num_groups=self.num_groups,
1068
+ )
1069
+
1070
+
1071
+ class RMSNorm(nn.Module):
1072
+
1073
+ def __init__(
1074
+ self,
1075
+ hidden_size: int,
1076
+ elementwise_affine: bool = True,
1077
+ bias: bool = False,
1078
+ eps: float = 1e-5,
1079
+ ) -> RMSNorm:
1080
+ super().__init__()
1081
+
1082
+ self.hidden_size = hidden_size
1083
+ self.elementwise_affine = elementwise_affine
1084
+ self.eps = eps
1085
+
1086
+ self.register_parameter("weight", None)
1087
+ self.register_parameter("bias", None)
1088
+ if elementwise_affine:
1089
+ self.weight = nn.Parameter(torch.empty(hidden_size))
1090
+ if bias:
1091
+ self.bias = nn.Parameter(torch.empty(hidden_size))
1092
+
1093
+ self.reset_parameters()
1094
+
1095
+ def reset_parameters(self):
1096
+ if self.elementwise_affine:
1097
+ nn.init.ones_(self.weight)
1098
+ if self.bias is not None:
1099
+ nn.init.zeros_(self.bias)
1100
+
1101
+ def __repr__(self) -> str:
1102
+ s = f"{self.__class__.__name__}({self.hidden_size}"
1103
+ if not self.elementwise_affine:
1104
+ s += f", elementwise_affine={self.elementwise_affine}"
1105
+ s += f", eps={self.eps}"
1106
+ s += ")"
1107
+ return s
1108
+
1109
+ def forward(self, x, residual=None, prenorm=False, residual_in_fp32=False):
1110
+ return rms_norm(
1111
+ x,
1112
+ self.weight,
1113
+ self.bias,
1114
+ residual=residual,
1115
+ eps=self.eps,
1116
+ prenorm=prenorm,
1117
+ residual_in_fp32=residual_in_fp32,
1118
+ )
1119
+
1120
+
1121
+ class LayerNormLinearFunction(torch.autograd.Function):
1122
+
1123
+ @staticmethod
1124
+ @input_guard
1125
+ def forward(
1126
+ ctx,
1127
+ x,
1128
+ norm_weight,
1129
+ norm_bias,
1130
+ linear_weight,
1131
+ linear_bias,
1132
+ residual=None,
1133
+ eps=1e-5,
1134
+ prenorm=False,
1135
+ residual_in_fp32=False,
1136
+ is_rms_norm=False,
1137
+ num_groups=1,
1138
+ ):
1139
+ x_shape_og = x.shape
1140
+
1141
+ if x.shape[-1] % num_groups != 0:
1142
+ raise ValueError('num_channels must be divisible by num_groups')
1143
+ # reshape input data into 2D tensor
1144
+ x = x.reshape(-1, (x.shape[-1] // num_groups))
1145
+ if residual is not None:
1146
+ assert residual.shape == x_shape_og
1147
+ residual = residual.reshape_as(x)
1148
+ residual_dtype = (
1149
+ residual.dtype
1150
+ if residual is not None
1151
+ else (torch.float32 if residual_in_fp32 else None)
1152
+ )
1153
+ y, mean, rstd, res_out = layer_norm_fwd(
1154
+ x,
1155
+ norm_weight,
1156
+ norm_bias,
1157
+ eps,
1158
+ residual,
1159
+ out_dtype=None if not torch.is_autocast_enabled() else torch.get_autocast_gpu_dtype(),
1160
+ residual_dtype=residual_dtype,
1161
+ is_rms_norm=is_rms_norm,
1162
+ num_groups=num_groups,
1163
+ )
1164
+ y = y.reshape(x_shape_og)
1165
+ dtype = torch.get_autocast_gpu_dtype() if torch.is_autocast_enabled() else y.dtype
1166
+ linear_weight = linear_weight.to(dtype)
1167
+ linear_bias = linear_bias.to(dtype) if linear_bias is not None else None
1168
+ out = F.linear(y.to(linear_weight.dtype), linear_weight, linear_bias)
1169
+ # We don't store y, will be recomputed in the backward pass to save memory
1170
+ ctx.save_for_backward(res_out, norm_weight, norm_bias, linear_weight, mean, rstd)
1171
+ ctx.x_shape_og = x_shape_og
1172
+ ctx.eps = eps
1173
+ ctx.is_rms_norm = is_rms_norm
1174
+ ctx.num_groups = num_groups
1175
+ ctx.has_residual = residual is not None
1176
+ ctx.prenorm = prenorm
1177
+ ctx.x_dtype = x.dtype
1178
+ ctx.linear_bias_is_none = linear_bias is None
1179
+ return out if not prenorm else (out, res_out.reshape(x_shape_og))
1180
+
1181
+ @staticmethod
1182
+ @input_guard
1183
+ def backward(ctx, dout, *args):
1184
+ x, norm_weight, norm_bias, linear_weight, mean, rstd = ctx.saved_tensors
1185
+ dout = dout.reshape(-1, dout.shape[-1])
1186
+ dy = F.linear(dout, linear_weight.t())
1187
+ dy = dy.reshape(-1, (dy.shape[-1] // ctx.num_groups))
1188
+ dlinear_bias = None if ctx.linear_bias_is_none else dout.sum(0)
1189
+ assert dy.shape == x.shape
1190
+ if ctx.prenorm:
1191
+ dresidual = args[0]
1192
+ dresidual = dresidual.reshape(-1, x.shape[-1])
1193
+ assert dresidual.shape == x.shape
1194
+ else:
1195
+ dresidual = None
1196
+ dx, dnorm_weight, dnorm_bias, dresidual_in, y = layer_norm_bwd(
1197
+ dy,
1198
+ x,
1199
+ norm_weight,
1200
+ norm_bias,
1201
+ mean,
1202
+ rstd,
1203
+ dresidual,
1204
+ ctx.has_residual,
1205
+ ctx.is_rms_norm,
1206
+ x_dtype=ctx.x_dtype,
1207
+ recompute_output=True,
1208
+ num_groups=ctx.num_groups,
1209
+ )
1210
+ dlinear_weight = torch.einsum("bo,bi->oi", dout, y.view(-1, linear_weight.shape[-1]))
1211
+ return (
1212
+ dx.reshape(ctx.x_shape_og),
1213
+ dnorm_weight,
1214
+ dnorm_bias,
1215
+ dlinear_weight,
1216
+ dlinear_bias,
1217
+ dresidual_in.reshape(ctx.x_shape_og) if ctx.has_residual else None,
1218
+ None,
1219
+ None,
1220
+ None,
1221
+ None,
1222
+ None,
1223
+ )
1224
+
1225
+
1226
+ class LayerNormLinear(nn.Module):
1227
+
1228
+ def __init__(
1229
+ self,
1230
+ hidden_size,
1231
+ elementwise_affine: bool = True,
1232
+ bias: bool = False,
1233
+ eps: float = 1e-5,
1234
+ ) -> LayerNormLinear:
1235
+ super().__init__()
1236
+
1237
+ self.hidden_size = hidden_size
1238
+ self.elementwise_affine = elementwise_affine
1239
+ self.eps = eps
1240
+
1241
+ self.register_parameter("weight", None)
1242
+ self.register_parameter("bias", None)
1243
+ if elementwise_affine:
1244
+ self.weight = nn.Parameter(torch.empty(hidden_size))
1245
+ if bias:
1246
+ self.bias = nn.Parameter(torch.empty(hidden_size))
1247
+
1248
+ self.reset_parameters()
1249
+
1250
+ def reset_parameters(self):
1251
+ if self.elementwise_affine:
1252
+ nn.init.ones_(self.weight)
1253
+ if self.bias is not None:
1254
+ nn.init.zeros_(self.bias)
1255
+
1256
+ def __repr__(self) -> str:
1257
+ s = f"{self.__class__.__name__}({self.hidden_size}"
1258
+ if not self.elementwise_affine:
1259
+ s += f", elementwise_affine={self.elementwise_affine}"
1260
+ s += f", eps={self.eps}"
1261
+ s += ")"
1262
+ return s
1263
+
1264
+ def forward(self, x, weight, bias, residual=None, prenorm=False, residual_in_fp32=False):
1265
+ return layer_norm_linear(
1266
+ x=x,
1267
+ norm_weight=self.weight,
1268
+ norm_bias=self.bias,
1269
+ linear_weight=weight,
1270
+ linear_bias=bias,
1271
+ residual=residual,
1272
+ eps=self.eps,
1273
+ prenorm=prenorm,
1274
+ residual_in_fp32=residual_in_fp32,
1275
+ is_rms_norm=False,
1276
+ )
1277
+
1278
+
1279
+ class GroupNormLinear(nn.Module):
1280
+
1281
+ def __init__(
1282
+ self,
1283
+ num_groups: int,
1284
+ hidden_size: int,
1285
+ elementwise_affine: bool = True,
1286
+ bias: bool = False,
1287
+ eps: float = 1e-5,
1288
+ is_rms_norm: bool = False,
1289
+ ) -> GroupNormLinear:
1290
+ super().__init__()
1291
+
1292
+ if hidden_size % num_groups != 0:
1293
+ raise ValueError('num_channels must be divisible by num_groups')
1294
+
1295
+ self.num_groups = num_groups
1296
+ self.hidden_size = hidden_size
1297
+ self.elementwise_affine = elementwise_affine
1298
+ self.eps = eps
1299
+ self.is_rms_norm = is_rms_norm
1300
+
1301
+ self.register_parameter("weight", None)
1302
+ self.register_parameter("bias", None)
1303
+ if elementwise_affine:
1304
+ self.weight = nn.Parameter(torch.empty(hidden_size))
1305
+ if bias:
1306
+ self.bias = nn.Parameter(torch.empty(hidden_size))
1307
+
1308
+ self.reset_parameters()
1309
+
1310
+ def reset_parameters(self):
1311
+ if self.elementwise_affine:
1312
+ nn.init.ones_(self.weight)
1313
+ if self.bias is not None:
1314
+ nn.init.zeros_(self.bias)
1315
+
1316
+ def __repr__(self) -> str:
1317
+ s = f"{self.__class__.__name__}({self.num_groups}, {self.hidden_size}"
1318
+ if not self.elementwise_affine:
1319
+ s += f", elementwise_affine={self.elementwise_affine}"
1320
+ if self.is_rms_norm:
1321
+ s += f", is_rms_norm={self.is_rms_norm}"
1322
+ s += f", eps={self.eps}"
1323
+ s += ")"
1324
+ return s
1325
+
1326
+ def forward(self, x, weight, bias, residual=None, prenorm=False, residual_in_fp32=False):
1327
+ return layer_norm_linear(
1328
+ x=x,
1329
+ norm_weight=self.weight,
1330
+ norm_bias=self.bias,
1331
+ linear_weight=weight,
1332
+ linear_bias=bias,
1333
+ residual=residual,
1334
+ eps=self.eps,
1335
+ prenorm=prenorm,
1336
+ residual_in_fp32=residual_in_fp32,
1337
+ is_rms_norm=self.is_rms_norm,
1338
+ num_groups=self.num_groups,
1339
+ )
1340
+
1341
+
1342
+ class RMSNormLinear(nn.Module):
1343
+
1344
+ def __init__(
1345
+ self,
1346
+ hidden_size,
1347
+ elementwise_affine: bool = True,
1348
+ bias: bool = False,
1349
+ eps: float = 1e-5,
1350
+ ) -> RMSNormLinear:
1351
+ super().__init__()
1352
+
1353
+ self.hidden_size = hidden_size
1354
+ self.elementwise_affine = elementwise_affine
1355
+ self.eps = eps
1356
+
1357
+ self.register_parameter("weight", None)
1358
+ self.register_parameter("bias", None)
1359
+ if elementwise_affine:
1360
+ self.weight = nn.Parameter(torch.empty(hidden_size))
1361
+ if bias:
1362
+ self.bias = nn.Parameter(torch.empty(hidden_size))
1363
+
1364
+ self.reset_parameters()
1365
+
1366
+ def reset_parameters(self):
1367
+ if self.elementwise_affine:
1368
+ nn.init.ones_(self.weight)
1369
+ if self.bias is not None:
1370
+ nn.init.zeros_(self.bias)
1371
+
1372
+ def __repr__(self) -> str:
1373
+ s = f"{self.__class__.__name__}({self.hidden_size}"
1374
+ if not self.elementwise_affine:
1375
+ s += f", elementwise_affine={self.elementwise_affine}"
1376
+ s += f", eps={self.eps}"
1377
+ s += ")"
1378
+ return s
1379
+
1380
+ def forward(self, x, weight, bias, residual=None, prenorm=False, residual_in_fp32=False):
1381
+ return layer_norm_linear(
1382
+ x=x,
1383
+ norm_weight=self.weight,
1384
+ norm_bias=self.bias,
1385
+ linear_weight=weight,
1386
+ linear_bias=bias,
1387
+ residual=residual,
1388
+ eps=self.eps,
1389
+ prenorm=prenorm,
1390
+ residual_in_fp32=residual_in_fp32,
1391
+ is_rms_norm=True,
1392
+ )
1393
+
1394
+
1395
+ class NormParallel(ParallelStyle):
1396
+
1397
+ def __init__(self, *, sequence_dim: int = 1, use_local_output: bool = False):
1398
+ super().__init__()
1399
+ self.sequence_sharding = (Shard(sequence_dim),)
1400
+ self.use_local_output = use_local_output
1401
+
1402
+ def _replicate_module_fn(
1403
+ self, name: str, module: nn.Module, device_mesh: DeviceMesh,
1404
+ ):
1405
+ for p_name, param in module.named_parameters():
1406
+ # simple replication with fixed ones_ init from LayerNorm/RMSNorm, which allow
1407
+ # us to simply just use from_local
1408
+ replicated_param = torch.nn.Parameter(
1409
+ DTensor.from_local(param, device_mesh, [Replicate()], run_check=False),
1410
+ )
1411
+ module.register_parameter(p_name, replicated_param)
1412
+
1413
+ @staticmethod
1414
+ def _prepare_input_fn(sequence_sharding, mod, inputs, device_mesh):
1415
+ input_tensor = inputs[0]
1416
+ if isinstance(input_tensor, DTensor):
1417
+ # if the passed in input DTensor is not sharded on the sequence dim, we need to redistribute it
1418
+ if input_tensor.placements != sequence_sharding:
1419
+ input_tensor = input_tensor.redistribute(
1420
+ placements=sequence_sharding, async_op=True,
1421
+ )
1422
+ return input_tensor
1423
+ elif isinstance(input_tensor, torch.Tensor):
1424
+ # assume the input passed in already sharded on the sequence dim and create the DTensor
1425
+ return DTensor.from_local(
1426
+ input_tensor, device_mesh, sequence_sharding, run_check=False,
1427
+ )
1428
+ else:
1429
+ raise ValueError(
1430
+ f"expecting input of {mod} to be a torch.Tensor or DTensor, but got {input_tensor}",
1431
+ )
1432
+
1433
+ @staticmethod
1434
+ def _prepare_output_fn(use_local_output, mod, outputs, device_mesh):
1435
+ return outputs.to_local() if use_local_output else outputs
1436
+
1437
+ def _apply(self, module: nn.Module, device_mesh: DeviceMesh) -> nn.Module:
1438
+ return distribute_module(
1439
+ module,
1440
+ device_mesh,
1441
+ self._replicate_module_fn,
1442
+ partial(self._prepare_input_fn, self.sequence_sharding),
1443
+ partial(self._prepare_output_fn, self.use_local_output),
1444
+ )
code/flash-linear-attention/fla/modules/layernorm_gated.py ADDED
@@ -0,0 +1,527 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2024, Tri Dao.
2
+ # Based on the Triton LayerNorm tutorial: https://triton-lang.org/main/getting-started/tutorials/05-layer-norm.html
3
+ # For the backward pass, we keep weight_grad and bias_grad in registers and accumulate.
4
+ # This backward pass is faster for dimensions up to 8k, but after that it's much slower due to register spilling.
5
+ # The models we train have hidden dim up to 8k anyway (e.g. Llama 70B), so this is fine.
6
+
7
+ import math
8
+
9
+ import torch
10
+ import torch.nn as nn
11
+ import torch.nn.functional as F
12
+ import triton
13
+ import triton.language as tl
14
+ from einops import rearrange
15
+
16
+ from fla.utils import get_multiprocessor_count, input_guard
17
+
18
+
19
+ def rms_norm_ref(x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True, upcast=True):
20
+ dtype = x.dtype
21
+ weight = weight.float()
22
+ bias = bias.float() if bias is not None else None
23
+ if upcast:
24
+ x = x.float()
25
+ z = z.float() if z is not None else z
26
+ if z is not None and not norm_before_gate:
27
+ x = x * F.silu(z)
28
+ if group_size is None:
29
+ rstd = 1 / torch.sqrt((x.square()).mean(dim=-1, keepdim=True) + eps)
30
+ out = (x * rstd * weight) + bias if bias is not None else (x * rstd * weight)
31
+ else:
32
+ x_group = rearrange(x, "... (g d) -> ... g d", d=group_size)
33
+ rstd = 1 / torch.sqrt((x_group.square()).mean(dim=-1, keepdim=True) + eps)
34
+ out = rearrange(x_group * rstd, "... g d -> ... (g d)") * weight
35
+ if bias is not None:
36
+ out = out + bias
37
+ if z is not None and norm_before_gate:
38
+ out *= F.silu(z)
39
+ return out.to(dtype)
40
+
41
+
42
+ @triton.heuristics({
43
+ "HAS_BIAS": lambda args: args["B"] is not None,
44
+ "HAS_Z": lambda args: args["Z"] is not None,
45
+ })
46
+ @triton.jit
47
+ def layer_norm_fwd_kernel(
48
+ X, # pointer to the input
49
+ Y, # pointer to the output
50
+ W, # pointer to the weights
51
+ B, # pointer to the biases
52
+ Z, # pointer to the other branch
53
+ Mean, # pointer to the mean
54
+ Rstd, # pointer to the 1/std
55
+ stride_x_row, # how much to increase the pointer when moving by 1 row
56
+ stride_y_row,
57
+ stride_z_row,
58
+ M, # number of rows in X
59
+ N, # number of columns in X
60
+ eps, # epsilon to avoid division by zero
61
+ BLOCK_N: tl.constexpr,
62
+ HAS_BIAS: tl.constexpr,
63
+ HAS_Z: tl.constexpr,
64
+ NORM_BEFORE_GATE: tl.constexpr,
65
+ IS_RMS_NORM: tl.constexpr,
66
+ ):
67
+ # Map the program id to the row of X and Y it should compute.
68
+ row = tl.program_id(0)
69
+ group = tl.program_id(1)
70
+ X += row * stride_x_row + group * N
71
+ Y += row * stride_y_row + group * N
72
+ if HAS_Z:
73
+ Z += row * stride_z_row + group * N
74
+ if not IS_RMS_NORM:
75
+ Mean += group * M
76
+ Rstd += group * M
77
+ W += group * N
78
+ if HAS_BIAS:
79
+ B += group * N
80
+ # Compute mean and variance
81
+ cols = tl.arange(0, BLOCK_N)
82
+ x = tl.load(X + cols, mask=cols < N, other=0.).to(tl.float32)
83
+ if HAS_Z and not NORM_BEFORE_GATE:
84
+ z = tl.load(Z + cols, mask=cols < N).to(tl.float32)
85
+ x *= z * tl.sigmoid(z)
86
+ if not IS_RMS_NORM:
87
+ mean = tl.sum(x, axis=0) / N
88
+ tl.store(Mean + row, mean)
89
+ xbar = tl.where(cols < N, x - mean, 0.)
90
+ var = tl.sum(xbar * xbar, axis=0) / N
91
+ else:
92
+ xbar = tl.where(cols < N, x, 0.)
93
+ var = tl.sum(xbar * xbar, axis=0) / N
94
+ rstd = 1 / tl.sqrt(var + eps)
95
+ tl.store(Rstd + row, rstd)
96
+ # Normalize and apply linear transformation
97
+ mask = cols < N
98
+ w = tl.load(W + cols, mask=mask).to(tl.float32)
99
+ if HAS_BIAS:
100
+ b = tl.load(B + cols, mask=mask).to(tl.float32)
101
+ x_hat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd
102
+ y = x_hat * w + b if HAS_BIAS else x_hat * w
103
+ if HAS_Z and NORM_BEFORE_GATE:
104
+ z = tl.load(Z + cols, mask=mask).to(tl.float32)
105
+ y *= z * tl.sigmoid(z)
106
+ # Write output
107
+ tl.store(Y + cols, y, mask=mask)
108
+
109
+
110
+ def layer_norm_fwd(
111
+ x: torch.Tensor,
112
+ weight: torch.Tensor,
113
+ bias: torch.Tensor,
114
+ eps: float,
115
+ z: torch.Tensor = None,
116
+ out: torch.Tensor = None,
117
+ group_size: int = None,
118
+ norm_before_gate: bool = True,
119
+ is_rms_norm: bool = False,
120
+ ):
121
+ M, N = x.shape
122
+ if group_size is None:
123
+ group_size = N
124
+ assert N % group_size == 0
125
+ ngroups = N // group_size
126
+ assert x.stride(-1) == 1
127
+ if z is not None:
128
+ assert z.stride(-1) == 1
129
+ assert z.shape == (M, N)
130
+ assert weight.shape == (N,)
131
+ assert weight.stride(-1) == 1
132
+ if bias is not None:
133
+ assert bias.stride(-1) == 1
134
+ assert bias.shape == (N,)
135
+ # allocate output
136
+ if out is not None:
137
+ assert out.shape == x.shape
138
+ else:
139
+ out = torch.empty_like(x)
140
+ assert out.stride(-1) == 1
141
+ mean = torch.empty((ngroups * M, ), dtype=torch.float32, device=x.device) if not is_rms_norm else None
142
+ rstd = torch.empty((ngroups * M, ), dtype=torch.float32, device=x.device)
143
+ # Less than 64KB per feature: enqueue fused kernel
144
+ MAX_FUSED_SIZE = 65536 // x.element_size()
145
+ BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(group_size))
146
+ if group_size > BLOCK_N:
147
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
148
+ # heuristics for number of warps
149
+ num_warps = min(max(BLOCK_N // 256, 1), 8)
150
+ grid = (M, ngroups)
151
+ layer_norm_fwd_kernel[grid](
152
+ x,
153
+ out,
154
+ weight,
155
+ bias,
156
+ z,
157
+ mean,
158
+ rstd,
159
+ x.stride(0),
160
+ out.stride(0),
161
+ z.stride(0) if z is not None else 0,
162
+ M,
163
+ group_size,
164
+ eps,
165
+ BLOCK_N=BLOCK_N,
166
+ NORM_BEFORE_GATE=norm_before_gate,
167
+ IS_RMS_NORM=is_rms_norm,
168
+ num_warps=num_warps,
169
+ )
170
+ return out, mean, rstd
171
+
172
+
173
+ @triton.heuristics({
174
+ "HAS_BIAS": lambda args: args["B"] is not None,
175
+ "HAS_Z": lambda args: args["Z"] is not None,
176
+ "RECOMPUTE_OUTPUT": lambda args: args["Y"] is not None,
177
+ })
178
+ @triton.jit
179
+ def layer_norm_bwd_kernel(
180
+ X, # pointer to the input
181
+ W, # pointer to the weights
182
+ B, # pointer to the biases
183
+ Z, # pointer to the other branch
184
+ Y, # pointer to the output to be recomputed
185
+ DY, # pointer to the output gradient
186
+ DX, # pointer to the input gradient
187
+ DW, # pointer to the partial sum of weights gradient
188
+ DB, # pointer to the partial sum of biases gradient
189
+ DZ, # pointer to the other branch
190
+ Mean, # pointer to the mean
191
+ Rstd, # pointer to the 1/std
192
+ stride_x_row, # how much to increase the pointer when moving by 1 row
193
+ stride_z_row,
194
+ stride_y_row,
195
+ stride_dy_row,
196
+ stride_dx_row,
197
+ stride_dz_row,
198
+ stride_dw_row,
199
+ stride_db_row,
200
+ M, # number of rows in X
201
+ N, # number of columns in X
202
+ eps, # epsilon to avoid division by zero
203
+ rows_per_program,
204
+ NORM_BEFORE_GATE: tl.constexpr,
205
+ IS_RMS_NORM: tl.constexpr,
206
+ HAS_BIAS: tl.constexpr,
207
+ HAS_Z: tl.constexpr,
208
+ RECOMPUTE_OUTPUT: tl.constexpr,
209
+ BLOCK_N: tl.constexpr,
210
+ ):
211
+ # Map the program id to the elements of X, DX, and DY it should compute.
212
+ row_block_id = tl.program_id(0)
213
+ group = tl.program_id(1)
214
+ row_start = row_block_id * rows_per_program
215
+ cols = tl.arange(0, BLOCK_N)
216
+ mask = cols < N
217
+ X += row_start * stride_x_row + group * N
218
+ if HAS_Z:
219
+ Z += row_start * stride_z_row + group * N
220
+ DZ += row_start * stride_dz_row + group * N
221
+ DY += row_start * stride_dy_row + group * N
222
+ DX += row_start * stride_dx_row + group * N
223
+ if RECOMPUTE_OUTPUT:
224
+ Y += row_start * stride_y_row + group * N
225
+ if not IS_RMS_NORM:
226
+ Mean += group * M
227
+ Rstd += group * M
228
+ W += group * N
229
+ w = tl.load(W + cols, mask=mask).to(tl.float32)
230
+ if (RECOMPUTE_OUTPUT or HAS_Z) and HAS_BIAS:
231
+ B += group * N
232
+ b = tl.load(B + cols, mask=mask, other=0.).to(tl.float32)
233
+ dw = tl.zeros((BLOCK_N,), dtype=tl.float32)
234
+ if HAS_BIAS:
235
+ db = tl.zeros((BLOCK_N,), dtype=tl.float32)
236
+ row_end = min((row_block_id + 1) * rows_per_program, M)
237
+ for row in range(row_start, row_end):
238
+ # Load data to SRAM
239
+ x = tl.load(X + cols, mask=mask, other=0).to(tl.float32)
240
+ dy = tl.load(DY + cols, mask=mask, other=0).to(tl.float32)
241
+ if not IS_RMS_NORM:
242
+ mean = tl.load(Mean + row)
243
+ if HAS_Z and not NORM_BEFORE_GATE:
244
+ z = tl.load(Z + cols, mask=mask, other=0.).to(tl.float32)
245
+ x_og = x
246
+ x = x_og * z * tl.sigmoid(z)
247
+ rstd = tl.load(Rstd + row)
248
+ # Compute dx
249
+ xhat = (x - mean) * rstd if not IS_RMS_NORM else x * rstd
250
+ xhat = tl.where(mask, xhat, 0.)
251
+ if HAS_Z and NORM_BEFORE_GATE:
252
+ z = tl.load(Z + cols, mask=mask, other=0.).to(tl.float32)
253
+ z_sigmoid = tl.sigmoid(z)
254
+ y = xhat * w + b if HAS_BIAS else xhat * w
255
+ if RECOMPUTE_OUTPUT:
256
+ tl.store(Y + cols, y * z * z_sigmoid, mask=mask)
257
+ dz = dy * y * z_sigmoid * (1 + z * (1 - z_sigmoid))
258
+ tl.store(DZ + cols, dz, mask=mask)
259
+ dy *= z * z_sigmoid
260
+ else:
261
+ if RECOMPUTE_OUTPUT:
262
+ y = xhat * w + b if HAS_BIAS else xhat * w
263
+ tl.store(Y + cols, y, mask=mask)
264
+ wdy = w * dy
265
+ c1 = tl.sum(xhat * wdy, axis=0) / N
266
+ if not IS_RMS_NORM:
267
+ c2 = tl.sum(wdy, axis=0) / N
268
+ dx = (wdy - (xhat * c1 + c2)) * rstd
269
+ else:
270
+ dx = (wdy - xhat * c1) * rstd
271
+ dw += dy * xhat
272
+ if HAS_BIAS:
273
+ db += dy
274
+ if HAS_Z and not NORM_BEFORE_GATE:
275
+ z_sigmoid = tl.sigmoid(z)
276
+ dz = dx * x_og * z_sigmoid * (1 + z * (1 - z_sigmoid))
277
+ tl.store(DZ + cols, dz, mask=mask)
278
+ dx *= z * z_sigmoid
279
+ # Write dx
280
+ tl.store(DX + cols, dx, mask=mask)
281
+
282
+ X += stride_x_row
283
+ if HAS_Z:
284
+ Z += stride_z_row
285
+ DZ += stride_dz_row
286
+ if RECOMPUTE_OUTPUT:
287
+ Y += stride_y_row
288
+ DY += stride_dy_row
289
+ DX += stride_dx_row
290
+ tl.store(DW + row_block_id * stride_dw_row + group * N + cols, dw, mask=mask)
291
+ if HAS_BIAS:
292
+ tl.store(DB + row_block_id * stride_db_row + group * N + cols, db, mask=mask)
293
+
294
+
295
+ def layer_norm_bwd(
296
+ dy: torch.Tensor,
297
+ x: torch.Tensor,
298
+ weight: torch.Tensor,
299
+ bias: torch.Tensor,
300
+ eps: float,
301
+ mean: torch.Tensor,
302
+ rstd: torch.Tensor,
303
+ z: torch.Tensor = None,
304
+ group_size: int = None,
305
+ norm_before_gate: bool = True,
306
+ is_rms_norm: bool = False,
307
+ recompute_output: bool = False,
308
+ dz: torch.Tensor = None,
309
+ out: torch.Tensor = None,
310
+ ):
311
+ M, N = x.shape
312
+ if group_size is None:
313
+ group_size = N
314
+ assert N % group_size == 0
315
+ ngroups = N // group_size
316
+ assert x.stride(-1) == 1
317
+ assert dy.stride(-1) == 1
318
+ assert dy.shape == (M, N)
319
+ if z is not None:
320
+ assert z.stride(-1) == 1
321
+ assert z.shape == (M, N)
322
+ assert weight.shape == (N,)
323
+ assert weight.stride(-1) == 1
324
+ if bias is not None:
325
+ assert bias.stride(-1) == 1
326
+ assert bias.shape == (N,)
327
+ # allocate output
328
+ dx = torch.empty_like(x)
329
+ if dz is not None:
330
+ assert z is not None
331
+ assert dz.shape == z.shape
332
+ assert dz.stride(-1) == 1
333
+ else:
334
+ dz = torch.empty_like(z) if z is not None else None
335
+ if recompute_output:
336
+ if out is None:
337
+ out = torch.empty_like(x)
338
+ assert out.shape == x.shape
339
+
340
+ # Less than 64KB per feature: enqueue fused kernel
341
+ MAX_FUSED_SIZE = 65536 // x.element_size()
342
+ BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(group_size))
343
+ if group_size > BLOCK_N:
344
+ raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
345
+ # heuristics for number of warps
346
+ num_warps = min(max(BLOCK_N // 256, 1), 8)
347
+ sm_count = get_multiprocessor_count(x.device.index)
348
+ # If group size is small (e.g., 64), we're only using 1 warp. So having just 108 programs
349
+ # would limit the occupancy.
350
+ nrow_groups = math.ceil(sm_count * math.ceil(4 / num_warps) / ngroups)
351
+ _dw = torch.empty((nrow_groups, N), dtype=torch.float32, device=weight.device)
352
+ _db = torch.empty((nrow_groups, N), dtype=torch.float32, device=bias.device) if bias is not None else None
353
+ rows_per_program = math.ceil(M / nrow_groups)
354
+ grid = (nrow_groups, ngroups)
355
+ layer_norm_bwd_kernel[grid](
356
+ x,
357
+ weight,
358
+ bias,
359
+ z,
360
+ out if recompute_output else None,
361
+ dy,
362
+ dx,
363
+ _dw,
364
+ _db,
365
+ dz,
366
+ mean,
367
+ rstd,
368
+ x.stride(0),
369
+ z.stride(0) if z is not None else 0,
370
+ 0 if not recompute_output else out.stride(0),
371
+ dy.stride(0),
372
+ dx.stride(0),
373
+ dz.stride(0) if dz is not None else 0,
374
+ _dw.stride(0),
375
+ _db.stride(0) if _db is not None else 0,
376
+ M, group_size, eps,
377
+ rows_per_program,
378
+ BLOCK_N=BLOCK_N,
379
+ NORM_BEFORE_GATE=norm_before_gate,
380
+ IS_RMS_NORM=is_rms_norm,
381
+ num_warps=num_warps,
382
+ )
383
+ dw = _dw.sum(0).to(weight.dtype)
384
+ db = _db.sum(0).to(bias.dtype) if bias is not None else None
385
+ return (dx, dw, db, dz) if not recompute_output else (dx, dw, db, dz, out)
386
+
387
+
388
+ class LayerNormFn(torch.autograd.Function):
389
+
390
+ @input_guard
391
+ @staticmethod
392
+ def forward(ctx, x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True,
393
+ is_rms_norm=False):
394
+ """If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))
395
+ """
396
+
397
+ x_shape_og = x.shape
398
+ # reshape input data into 2D tensor
399
+ x = x.reshape(-1, x.shape[-1])
400
+ if x.stride(-1) != 1:
401
+ x = x.contiguous()
402
+ if z is not None:
403
+ assert z.shape == x_shape_og
404
+ z = z.reshape(-1, z.shape[-1])
405
+ if z.stride(-1) != 1:
406
+ z = z.contiguous()
407
+ weight = weight.contiguous()
408
+ if bias is not None:
409
+ bias = bias.contiguous()
410
+ y, mean, rstd = layer_norm_fwd(
411
+ x,
412
+ weight,
413
+ bias,
414
+ eps,
415
+ z=z,
416
+ group_size=group_size,
417
+ norm_before_gate=norm_before_gate,
418
+ is_rms_norm=is_rms_norm,
419
+ )
420
+ ctx.save_for_backward(x, weight, bias, mean, rstd, z)
421
+ ctx.x_shape_og = x_shape_og
422
+ ctx.eps = eps
423
+ ctx.group_size = group_size
424
+ ctx.norm_before_gate = norm_before_gate
425
+ ctx.is_rms_norm = is_rms_norm
426
+ return y.reshape(x_shape_og)
427
+
428
+ @input_guard
429
+ @staticmethod
430
+ def backward(ctx, dy):
431
+ x, weight, bias, mean, rstd, z = ctx.saved_tensors
432
+ dy = dy.reshape(-1, dy.shape[-1])
433
+ if dy.stride(-1) != 1:
434
+ dy = dy.contiguous()
435
+ assert dy.shape == x.shape
436
+ dx, dw, db, dz = layer_norm_bwd(
437
+ dy,
438
+ x,
439
+ weight,
440
+ bias,
441
+ ctx.eps,
442
+ mean,
443
+ rstd,
444
+ z,
445
+ ctx.group_size,
446
+ ctx.norm_before_gate,
447
+ ctx.is_rms_norm,
448
+ )
449
+ dx = dx.reshape(ctx.x_shape_og)
450
+ dz = dz.reshape(ctx.x_shape_og) if dz is not None else None
451
+ return dx, dw, db, dz, None, None, None, None
452
+
453
+
454
+ def layernorm_fn(x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True, is_rms_norm=False):
455
+ return LayerNormFn.apply(x, weight, bias, z, eps, group_size, norm_before_gate, is_rms_norm)
456
+
457
+
458
+ def rmsnorm_fn(x, weight, bias, z=None, eps=1e-6, group_size=None, norm_before_gate=True):
459
+ return LayerNormFn.apply(x, weight, bias, z, eps, group_size, norm_before_gate, True)
460
+
461
+
462
+ class LayerNormGated(nn.Module):
463
+
464
+ def __init__(
465
+ self,
466
+ hidden_size,
467
+ eps: float = 1e-5,
468
+ group_size: int | None = None,
469
+ norm_before_gate: bool = True,
470
+ device: torch.device | None = None,
471
+ dtype: torch.dtype | None = None,
472
+ ):
473
+ """If group_size is not None, we do GroupNorm with each group having group_size elements.
474
+ group_size=None is equivalent to group_size=hidden_size (i.e. there's only 1 group).
475
+ """
476
+
477
+ factory_kwargs = {"device": device, "dtype": dtype}
478
+ super().__init__()
479
+ self.eps = eps
480
+ self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
481
+ self.bias = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
482
+ self.group_size = group_size
483
+ self.norm_before_gate = norm_before_gate
484
+ self.reset_parameters()
485
+
486
+ def reset_parameters(self):
487
+ torch.nn.init.ones_(self.weight)
488
+ torch.nn.init.zeros_(self.bias)
489
+
490
+ def forward(self, x, z=None):
491
+ """If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))
492
+ """
493
+ return layernorm_fn(x, self.weight, self.bias, z=z, group_size=self.group_size, eps=self.eps,
494
+ norm_before_gate=self.norm_before_gate)
495
+
496
+
497
+ class RMSNormGated(nn.Module):
498
+
499
+ def __init__(
500
+ self,
501
+ hidden_size,
502
+ eps: float = 1e-5,
503
+ group_size: int | None = None,
504
+ norm_before_gate: bool = False,
505
+ device: torch.device | None = None,
506
+ dtype: torch.dtype | None = None,
507
+ ):
508
+ """If group_size is not None, we do GroupNorm with each group having group_size elements.
509
+ group_size=None is equivalent to group_size=hidden_size (i.e. there's only 1 group).
510
+ """
511
+ factory_kwargs = {"device": device, "dtype": dtype}
512
+ super().__init__()
513
+ self.eps = eps
514
+ self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
515
+ self.register_parameter("bias", None)
516
+ self.group_size = group_size
517
+ self.norm_before_gate = norm_before_gate
518
+ self.reset_parameters()
519
+
520
+ def reset_parameters(self):
521
+ torch.nn.init.ones_(self.weight)
522
+
523
+ def forward(self, x, z=None):
524
+ """If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))
525
+ """
526
+ return rmsnorm_fn(x, self.weight, self.bias, z=z, eps=self.eps, group_size=self.group_size,
527
+ norm_before_gate=self.norm_before_gate)
code/flash-linear-attention/fla/modules/mlp.py ADDED
@@ -0,0 +1,144 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+ from __future__ import annotations
4
+
5
+ from functools import partial
6
+ from typing import TYPE_CHECKING, Any
7
+
8
+ import torch
9
+ import torch.nn as nn
10
+ try:
11
+ from torch.distributed import DeviceMesh
12
+ except ImportError:
13
+ DeviceMesh = None
14
+ try:
15
+ from torch.distributed.tensor import Placement, Replicate, Shard, distribute_module
16
+ except ImportError:
17
+ Placement = None
18
+ Replicate = None
19
+ Shard = None
20
+ distribute_module = None
21
+ try:
22
+ from torch.distributed.tensor.parallel import ParallelStyle
23
+ except ImportError:
24
+ class ParallelStyle:
25
+ pass
26
+
27
+ from fla.modules.activations import swiglu, swiglu_linear
28
+
29
+ try:
30
+ from torch.distributed.tensor import DTensor
31
+ except (ImportError, AttributeError):
32
+ DTensor = None
33
+
34
+ if TYPE_CHECKING:
35
+ from transformers.processing_utils import Unpack
36
+
37
+
38
+ class GatedMLP(nn.Module):
39
+
40
+ def __init__(
41
+ self,
42
+ hidden_size: int,
43
+ hidden_ratio: int | None = None,
44
+ intermediate_size: int | None = None,
45
+ hidden_act: str = 'swish',
46
+ fuse_swiglu: bool = True,
47
+ ) -> GatedMLP:
48
+ super().__init__()
49
+
50
+ self.hidden_size = hidden_size
51
+ # the final number of params is `hidden_ratio * hidden_size^2`
52
+ # `intermediate_size` is chosen to be a multiple of 256 closest to `2/3 * hidden_size * hidden_ratio`
53
+ if hidden_ratio is None:
54
+ hidden_ratio = 4
55
+ if intermediate_size is None:
56
+ intermediate_size = int(hidden_size * hidden_ratio * 2 / 3)
57
+ intermediate_size = 256 * ((intermediate_size + 256 - 1) // 256)
58
+ self.hidden_ratio = hidden_ratio
59
+ self.intermediate_size = intermediate_size
60
+ self.hidden_act = hidden_act
61
+ self.fuse_swiglu = fuse_swiglu
62
+
63
+ if hidden_act != 'swish':
64
+ raise ValueError(f'Unsupported hidden_act: {hidden_act}')
65
+
66
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
67
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
68
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
69
+ if self.fuse_swiglu:
70
+ self.swiglu_linear = SwiGLULinear()
71
+
72
+ def forward(
73
+ self,
74
+ x: torch.Tensor,
75
+ **kwargs: Unpack[Any],
76
+ ) -> torch.Tensor:
77
+ gate, y = self.gate_proj(x), self.up_proj(x)
78
+ if self.fuse_swiglu:
79
+ return self.swiglu_linear(gate, y, self.down_proj.weight, self.down_proj.bias)
80
+ else:
81
+ return self.down_proj(swiglu(gate, y))
82
+
83
+
84
+ class SwiGLULinear(nn.Module):
85
+
86
+ def forward(self, x, y, weight, bias):
87
+ return swiglu_linear(x, y, weight, bias)
88
+
89
+
90
+ class SwiGLULinearParallel(ParallelStyle):
91
+ def __init__(
92
+ self,
93
+ *,
94
+ input_layouts: Placement | None = None,
95
+ output_layouts: Placement | None = None,
96
+ use_local_output: bool = True,
97
+ ):
98
+ super().__init__()
99
+ self.input_layouts = (input_layouts or Shard(-1),)
100
+ self.output_layouts = (output_layouts or Replicate(),)
101
+ self.desired_input_layouts = (Shard(-1),)
102
+ self.use_local_output = use_local_output
103
+
104
+ @staticmethod
105
+ def _prepare_input_fn(
106
+ input_layouts, desired_input_layouts, mod, inputs, device_mesh,
107
+ ):
108
+ x, y, weight, bias = inputs
109
+ if not isinstance(x, DTensor):
110
+ x = DTensor.from_local(x, device_mesh, input_layouts, run_check=False)
111
+ if x.placements != desired_input_layouts:
112
+ x = x.redistribute(placements=desired_input_layouts, async_op=True)
113
+
114
+ if not isinstance(y, DTensor):
115
+ y = DTensor.from_local(y, device_mesh, input_layouts, run_check=False)
116
+ if y.placements != desired_input_layouts:
117
+ y = y.redistribute(placements=desired_input_layouts, async_op=True)
118
+
119
+ if not isinstance(weight, DTensor):
120
+ weight = DTensor.from_local(weight, device_mesh, (Shard(1),))
121
+
122
+ if bias is not None and not isinstance(bias, DTensor):
123
+ bias = DTensor.from_local(bias, device_mesh, (Replicate(),))
124
+
125
+ return x, y, weight, bias
126
+
127
+ @staticmethod
128
+ def _prepare_output_fn(output_layouts, use_local_output, mod, outputs, device_mesh):
129
+ # Rowwise sharding produces partial output, depending on output layouts:
130
+ # 1. to replicate -> allreduce
131
+ # 2. to shard -> reduce_scatter
132
+ if outputs.placements != output_layouts:
133
+ outputs = outputs.redistribute(placements=output_layouts, async_op=True)
134
+ # back to local tensor if use_local_output is True
135
+ return outputs.to_local() if use_local_output else outputs
136
+
137
+ def _apply(self, module: nn.Module, device_mesh: DeviceMesh) -> nn.Module:
138
+ return distribute_module(
139
+ module,
140
+ device_mesh,
141
+ partition_fn=None,
142
+ input_fn=partial(self._prepare_input_fn, self.input_layouts, self.desired_input_layouts),
143
+ output_fn=partial(self._prepare_output_fn, self.output_layouts, self.use_local_output),
144
+ )
code/flash-linear-attention/fla/modules/parallel.py ADDED
@@ -0,0 +1,53 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch.nn as nn
5
+ try:
6
+ from torch.distributed import DeviceMesh
7
+ except ImportError:
8
+ DeviceMesh = None
9
+ try:
10
+ from torch.distributed.tensor import distribute_module
11
+ except ImportError:
12
+ distribute_module = None
13
+ try:
14
+ from torch.distributed.tensor.parallel import ParallelStyle
15
+ except ImportError:
16
+ class ParallelStyle:
17
+ pass
18
+ try:
19
+ from torch.distributed.tensor.placement_types import Placement
20
+ except ImportError:
21
+ Placement = None
22
+
23
+ try:
24
+ from torch.distributed.tensor import DTensor
25
+ except (ImportError, AttributeError):
26
+ DTensor = None
27
+
28
+
29
+ class PrepareModuleWeight(ParallelStyle):
30
+ def __init__(self, *, layouts: Placement | None = None):
31
+ super().__init__()
32
+ self.layouts = layouts
33
+
34
+ def _replicate_module_fn(
35
+ self,
36
+ name: str,
37
+ module: nn.Module,
38
+ device_mesh: DeviceMesh,
39
+ ):
40
+ for p_name, param in module.named_parameters():
41
+ replicated_param = nn.Parameter(
42
+ DTensor.from_local(param, device_mesh, [self.layouts], run_check=False),
43
+ )
44
+ module.register_parameter(p_name, replicated_param)
45
+
46
+ def _apply(self, module: nn.Module, device_mesh: DeviceMesh) -> nn.Module:
47
+ return distribute_module(
48
+ module,
49
+ device_mesh,
50
+ partition_fn=self._replicate_module_fn,
51
+ input_fn=None,
52
+ output_fn=None,
53
+ )
code/flash-linear-attention/fla/modules/rotary.py ADDED
@@ -0,0 +1,499 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import torch.nn as nn
6
+ import triton
7
+ import triton.language as tl
8
+ from einops import rearrange, repeat
9
+
10
+ from fla.ops.utils import prepare_chunk_indices
11
+ from fla.utils import autotune_cache_kwargs, get_multiprocessor_count, input_guard, is_amd
12
+
13
+ NUM_WARPS_AUTOTUNE = [2, 4, 8, 16] if is_amd else [2, 4, 8, 16, 32]
14
+
15
+
16
+ def rotate_half(x, interleaved=False):
17
+ if not interleaved:
18
+ x1, x2 = x.chunk(2, dim=-1)
19
+ return torch.cat((-x2, x1), dim=-1)
20
+ else:
21
+ x1, x2 = x[..., ::2], x[..., 1::2]
22
+ return rearrange(torch.stack((-x2, x1), dim=-1), '... d two -> ... (d two)', two=2)
23
+
24
+
25
+ def rotary_embedding_ref(x, cos, sin, interleaved=False):
26
+ ro_dim = cos.shape[-1] * 2
27
+ assert ro_dim <= x.shape[-1]
28
+ cos = repeat(cos, '... d -> ... 1 (2 d)' if not interleaved else '... d -> ... 1 (d 2)')
29
+ sin = repeat(sin, '... d -> ... 1 (2 d)' if not interleaved else '... d -> ... 1 (d 2)')
30
+ return torch.cat([x[..., :ro_dim] * cos + rotate_half(x[..., :ro_dim], interleaved) * sin, x[..., ro_dim:]], -1)
31
+
32
+
33
+ @triton.autotune(
34
+ configs=[
35
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
36
+ for num_warps in NUM_WARPS_AUTOTUNE
37
+ for num_stages in [2, 3, 4]
38
+ ],
39
+ key=['B', 'H', 'D', 'INTERLEAVED'],
40
+ **autotune_cache_kwargs,
41
+ )
42
+ @triton.jit(do_not_specialize=['T'])
43
+ def rotary_embedding_kernel(
44
+ x,
45
+ cos,
46
+ sin,
47
+ y,
48
+ cu_seqlens,
49
+ chunk_indices,
50
+ seq_offsets,
51
+ T,
52
+ B: tl.constexpr,
53
+ H: tl.constexpr,
54
+ D: tl.constexpr,
55
+ R: tl.constexpr,
56
+ TR: tl.constexpr,
57
+ BT: tl.constexpr,
58
+ BD: tl.constexpr,
59
+ IS_SEQLEN_OFFSETS_TENSOR: tl.constexpr,
60
+ IS_VARLEN: tl.constexpr,
61
+ INTERLEAVED: tl.constexpr,
62
+ CONJUGATE: tl.constexpr,
63
+ ):
64
+ i_t, i_b, i_h = tl.program_id(0), tl.program_id(1), tl.program_id(2)
65
+
66
+ if IS_VARLEN:
67
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
68
+ bos, eos = tl.load(cu_seqlens + i_n), tl.load(cu_seqlens + i_n + 1)
69
+ T = eos - bos
70
+ x = x + bos * H*D + i_h * D
71
+ y = y + bos * H*D + i_h * D
72
+ else:
73
+ i_n = i_b
74
+ x = x + i_n * T*H*D + i_h * D
75
+ y = y + i_n * T*H*D + i_h * D
76
+
77
+ if i_t * BT >= T:
78
+ return
79
+
80
+ o_t = i_t * BT + tl.arange(0, BT)
81
+ if not IS_SEQLEN_OFFSETS_TENSOR:
82
+ o_cs = o_t + seq_offsets
83
+ else:
84
+ o_cs = o_t + tl.load(seq_offsets + i_n)
85
+ m_t = (o_t >= 0) & (o_t < T) & (o_cs >= 0) & (o_cs < TR)
86
+
87
+ if not INTERLEAVED:
88
+ # Load the 1st and 2nd halves of x, do calculation, then store to 1st and 2nd halves of out
89
+ o_r = tl.arange(0, BD // 2)
90
+ p_x = x + o_t[:, None] * H*D + o_r[None, :]
91
+ p_cos = cos + (o_cs[:, None] * R + o_r[None, :])
92
+ p_sin = sin + (o_cs[:, None] * R + o_r[None, :])
93
+ mask = m_t[:, None] & (o_r < R)[None, :]
94
+
95
+ b_cos = tl.load(p_cos, mask=mask, other=1.0).to(tl.float32)
96
+ b_sin = tl.load(p_sin, mask=mask, other=0.0).to(tl.float32)
97
+ b_x0 = tl.load(p_x, mask=mask, other=0.0).to(tl.float32)
98
+ b_x1 = tl.load(p_x + R, mask=mask, other=0.0).to(tl.float32)
99
+ if CONJUGATE:
100
+ b_sin = -b_sin
101
+ b_o0 = b_x0 * b_cos - b_x1 * b_sin
102
+ b_o1 = b_x0 * b_sin + b_x1 * b_cos
103
+ # write back result
104
+ p_y = y + (o_t[:, None] * H*D + o_r[None, :])
105
+ tl.store(p_y, b_o0, mask=mask)
106
+ tl.store(p_y + R, b_o1, mask=mask)
107
+ else:
108
+ # We don't want to load x[0, 2, 4, ...] and x[1, 3, 5, ...] separately since both are slow.
109
+ # Instead, we load x0 = x[0, 1, 2, 3, ...] and x1 = x[1, 0, 3, 2, ...].
110
+ # Loading x0 will be fast but x1 will be slow.
111
+ # Then we load cos = cos[0, 0, 1, 1, ...] and sin = sin[0, 0, 1, 1, ...].
112
+ # Then we do the calculation and use tl.where to pick put the right outputs for the even
113
+ # and for the odd indices.
114
+ o_d = tl.arange(0, BD)
115
+ o_d_swap = o_d + ((o_d + 1) % 2) * 2 - 1 # 1, 0, 3, 2, 5, 4, ...
116
+ o_d_repeat = tl.arange(0, BD) // 2
117
+ p_x0 = x + o_t[:, None] * H*D + o_d[None, :]
118
+ p_x1 = x + o_t[:, None] * H*D + o_d_swap[None, :]
119
+ p_cos = cos + (o_cs[:, None] * R + o_d_repeat[None, :])
120
+ p_sin = sin + (o_cs[:, None] * R + o_d_repeat[None, :])
121
+ mask = m_t[:, None] & (o_d_repeat < R)[None, :]
122
+
123
+ b_cos = tl.load(p_cos, mask=mask, other=1.0).to(tl.float32)
124
+ b_sin = tl.load(p_sin, mask=mask, other=0.0).to(tl.float32)
125
+ b_x0 = tl.load(p_x0, mask=mask, other=0.0).to(tl.float32)
126
+ b_x1 = tl.load(p_x1, mask=mask, other=0.0).to(tl.float32)
127
+ if CONJUGATE:
128
+ b_sin = -b_sin
129
+ b_o0 = b_x0 * b_cos
130
+ b_o1 = b_x1 * b_sin
131
+ b_y = tl.where(o_d[None, :] % 2 == 0, b_o0 - b_o1, b_o0 + b_o1)
132
+ p_y = y + (o_t[:, None] * H*D + o_d[None, :])
133
+ tl.store(p_y, b_y, mask=mask)
134
+
135
+
136
+ def rotary_embedding_fwdbwd(
137
+ x: torch.Tensor,
138
+ cos: torch.Tensor,
139
+ sin: torch.Tensor,
140
+ seqlen_offsets: int | torch.Tensor = 0,
141
+ cu_seqlens: torch.Tensor | None = None,
142
+ interleaved: bool = False,
143
+ inplace: bool = False,
144
+ conjugate: bool = False,
145
+ ) -> torch.Tensor:
146
+ """
147
+ Args:
148
+ x: [B, T, H, D].
149
+ cos: [TR, R / 2]
150
+ sin: [TR, R / 2]
151
+ seqlen_offsets: integer or integer tensor of size [N]
152
+ cu_seqlens: [N + 1,] or None
153
+
154
+ Returns:
155
+ y: [B, T, H, D]
156
+ """
157
+ is_varlen = cu_seqlens is not None
158
+
159
+ B, T, H, D = x.shape
160
+ N = B if not is_varlen else cu_seqlens.shape[0] - 1
161
+ TR, R = cos.shape
162
+ R2 = R * 2
163
+
164
+ assert D <= 256, "Only support D <= 256"
165
+ assert TR >= T, f"TR must be >= T, got {TR} and {T}"
166
+
167
+ assert cos.dtype == sin.dtype, f"cos and sin must have the same dtype, got {cos.dtype} and {sin.dtype}"
168
+ assert x.dtype == cos.dtype, f"Input and cos/sin must have the same dtype, got {x.dtype} and {cos.dtype}"
169
+
170
+ if isinstance(seqlen_offsets, torch.Tensor):
171
+ assert seqlen_offsets.shape == (N,)
172
+ assert seqlen_offsets.dtype in [torch.int32, torch.int64]
173
+ else:
174
+ assert seqlen_offsets + T <= TR
175
+
176
+ y = torch.empty_like(x) if not inplace else x
177
+ if R2 < D and not inplace:
178
+ y[..., R2:].copy_(x[..., R2:])
179
+
180
+ BD = triton.next_power_of_2(R2)
181
+ BT = min(128, triton.next_power_of_2(triton.cdiv(T, get_multiprocessor_count(x.device.index))))
182
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if is_varlen else None
183
+ NT = len(chunk_indices) if is_varlen else triton.cdiv(T, BT)
184
+
185
+ grid = (NT, B, H)
186
+ rotary_embedding_kernel[grid](
187
+ x,
188
+ cos,
189
+ sin,
190
+ y,
191
+ cu_seqlens,
192
+ chunk_indices,
193
+ seqlen_offsets,
194
+ B=B,
195
+ T=T,
196
+ H=H,
197
+ D=D,
198
+ R=R,
199
+ TR=TR,
200
+ BT=BT,
201
+ BD=BD,
202
+ IS_SEQLEN_OFFSETS_TENSOR=isinstance(seqlen_offsets, torch.Tensor),
203
+ IS_VARLEN=is_varlen,
204
+ INTERLEAVED=interleaved,
205
+ CONJUGATE=conjugate,
206
+ )
207
+ return y
208
+
209
+
210
+ class RotaryEmbeddingFunction(torch.autograd.Function):
211
+
212
+ @staticmethod
213
+ @input_guard
214
+ def forward(
215
+ ctx,
216
+ x,
217
+ cos,
218
+ sin,
219
+ interleaved=False,
220
+ inplace=False,
221
+ seqlen_offsets: int | torch.Tensor = 0,
222
+ cu_seqlens: torch.Tensor | None = None,
223
+ ):
224
+ y = rotary_embedding_fwdbwd(
225
+ x,
226
+ cos,
227
+ sin,
228
+ seqlen_offsets=seqlen_offsets,
229
+ cu_seqlens=cu_seqlens,
230
+ interleaved=interleaved,
231
+ inplace=inplace,
232
+ )
233
+ if isinstance(seqlen_offsets, int):
234
+ # Can't save int with save_for_backward
235
+ ctx.save_for_backward(cos, sin, cu_seqlens)
236
+ ctx.seqlen_offsets = seqlen_offsets
237
+ else:
238
+ ctx.save_for_backward(cos, sin, cu_seqlens, seqlen_offsets)
239
+ ctx.seqlen_offsets = None
240
+ ctx.interleaved = interleaved
241
+ ctx.inplace = inplace
242
+ return y if not inplace else x
243
+
244
+ @staticmethod
245
+ @input_guard
246
+ def backward(ctx, do):
247
+ seqlen_offsets = ctx.seqlen_offsets
248
+ if seqlen_offsets is None:
249
+ cos, sin, cu_seqlens, seqlen_offsets = ctx.saved_tensors
250
+ else:
251
+ cos, sin, cu_seqlens = ctx.saved_tensors
252
+ # TD [2023-09-02]: For some reason Triton (2.0.0.post1) errors with
253
+ # "[CUDA]: invalid device context", and cloning makes it work. Idk why. Triton 2.1.0 works.
254
+ if not ctx.interleaved and not ctx.inplace:
255
+ do = do.clone()
256
+ dx = rotary_embedding_fwdbwd(
257
+ do,
258
+ cos,
259
+ sin,
260
+ seqlen_offsets=seqlen_offsets,
261
+ cu_seqlens=cu_seqlens,
262
+ interleaved=ctx.interleaved,
263
+ inplace=ctx.inplace,
264
+ conjugate=True,
265
+ )
266
+ return dx, None, None, None, None, None, None, None
267
+
268
+
269
+ def rotary_embedding(
270
+ x,
271
+ cos,
272
+ sin,
273
+ interleaved=False,
274
+ inplace=False,
275
+ seqlen_offsets: int | torch.Tensor = 0,
276
+ cu_seqlens: torch.Tensor | None = None,
277
+ ):
278
+ """
279
+ Args:
280
+ x: [B, T, H, D]
281
+ cos, sin: [TR, R//2]
282
+ interleaved:
283
+ If True, rotate pairs of even and odd dimensions (GPT-J style) instead of 1st half and 2nd half (GPT-NeoX style).
284
+ inplace:
285
+ If True, apply rotary embedding in-place.
286
+ seqlen_offsets: [N,] or int.
287
+ Each sequence in x is shifted by this amount.
288
+ Most commonly used in inference when we have KV cache.
289
+ cu_seqlens: [N + 1,] or None
290
+
291
+ Returns:
292
+ out: [B, T, H, D]
293
+ """
294
+ return RotaryEmbeddingFunction.apply(
295
+ x,
296
+ cos,
297
+ sin,
298
+ interleaved,
299
+ inplace,
300
+ seqlen_offsets,
301
+ cu_seqlens,
302
+ )
303
+
304
+
305
+ class RotaryEmbedding(nn.Module):
306
+ """
307
+ The rotary position embeddings from RoFormer_ (Su et. al).
308
+ A crucial insight from the method is that the query and keys are
309
+ transformed by rotation matrices which depend on the relative positions.
310
+
311
+ Other implementations are available in the Rotary Transformer repo_ and in
312
+ GPT-NeoX_, GPT-NeoX was an inspiration
313
+
314
+ .. _RoFormer: https://arxiv.org/abs/2104.09864
315
+ .. _repo: https://github.com/ZhuiyiTechnology/roformer
316
+ .. _GPT-NeoX: https://github.com/EleutherAI/gpt-neox
317
+
318
+ If scale_base is not None, this implements XPos (Sun et al., https://arxiv.org/abs/2212.10554).
319
+ A recommended value for scale_base is 512: https://github.com/HazyResearch/flash-attention/issues/96
320
+ Reference: https://github.com/sunyt32/torchscale/blob/main/torchscale/component/xpos_relative_position.py
321
+ """
322
+
323
+ def __init__(
324
+ self,
325
+ dim: int,
326
+ base: float = 10000.0,
327
+ scale_base: float | None = None,
328
+ interleaved: bool = False,
329
+ pos_idx_in_fp32: bool = True,
330
+ device: torch.device | None = None,
331
+ ):
332
+ """
333
+ interleaved:
334
+ If True, rotate pairs of even and odd dimensions (GPT-J style) instead of 1st half and 2nd half (GPT-NeoX style).
335
+ pos_idx_in_fp32:
336
+ If True, the position indices [0.0, ..., seqlen - 1] are in fp32, otherwise they might be in lower precision.
337
+ This option was added because previously (before 2023-07-02), when we construct
338
+ the position indices, we use the dtype of self.inv_freq.
339
+ In most cases this would be fp32, but if the model is trained in pure bf16 (not mixed precision), then
340
+ self.inv_freq would be bf16, and the position indices are also in bf16.
341
+ Because of the limited precision of bf16 (e.g. 1995.0 is rounded to 2000.0), the
342
+ embeddings for some positions will coincide.
343
+ To maintain compatibility with models previously trained in pure bf16, we add this option.
344
+ """
345
+ super().__init__()
346
+
347
+ self.dim = dim
348
+ self.base = float(base)
349
+ self.scale_base = scale_base
350
+ self.interleaved = interleaved
351
+ self.pos_idx_in_fp32 = pos_idx_in_fp32
352
+ self.device = device
353
+
354
+ # Generate and save the inverse frequency buffer (non trainable)
355
+ self.register_buffer("inv_freq", torch.empty(-(dim // -2), dtype=torch.float32, device=device), persistent=False)
356
+
357
+ scale = None
358
+ if scale_base is not None:
359
+ scale = torch.empty(-(dim // -2), dtype=torch.float32, device=device)
360
+ self.register_buffer("scale", scale, persistent=False)
361
+
362
+ self._seq_len_cached = 0
363
+ self._cos_cached = None
364
+ self._sin_cached = None
365
+ self._cos_k_cached = None
366
+ self._sin_k_cached = None
367
+
368
+ self.reset_parameters()
369
+
370
+ def reset_parameters(self):
371
+ with torch.no_grad():
372
+ self.inv_freq.copy_(self._compute_inv_freq(device=self.inv_freq.device))
373
+ if self.scale_base is not None:
374
+ self.scale.copy_(self._compute_scale(device=self.scale.device))
375
+
376
+ def __repr__(self):
377
+ s = f"{self.__class__.__name__}("
378
+ s += f"dim={self.dim}, "
379
+ s += f"base={self.base}, "
380
+ s += f"interleaved={self.interleaved}, "
381
+ if self.scale_base is not None:
382
+ s += f"scale_base={self.scale_base}, "
383
+ s += f"pos_idx_in_fp32={self.pos_idx_in_fp32})"
384
+ return s
385
+
386
+ def _compute_inv_freq(self, device=None):
387
+ return 1.0 / (
388
+ self.base
389
+ ** (torch.arange(0, self.dim, 2, device=device, dtype=torch.float32) / self.dim)
390
+ )
391
+
392
+ def _compute_scale(self, device=None):
393
+ return (torch.arange(0, self.dim, 2, device=device, dtype=torch.float32) + 0.4 * self.dim) / (1.4 * self.dim)
394
+
395
+ def _update_cos_sin_cache(self, seqlen, device=None, dtype=None):
396
+ # Reset the tables if the sequence length has changed,
397
+ # if we're on a new device (possibly due to tracing for instance),
398
+ # or if we're switching from inference mode to training
399
+ if (
400
+ seqlen > self._seq_len_cached
401
+ or self._cos_cached is None
402
+ or self._cos_cached.device != device
403
+ or self._cos_cached.dtype != dtype
404
+ or (self.training and self._cos_cached.is_inference())
405
+ ):
406
+ self._seq_len_cached = seqlen
407
+ # We want fp32 here, not self.inv_freq.dtype, since the model could be loaded in bf16
408
+ # And the output of arange can be quite large, so bf16 would lose a lot of precision.
409
+ # However, for compatibility reason, we add an option to use the dtype of self.inv_freq.
410
+ if self.pos_idx_in_fp32:
411
+ t = torch.arange(seqlen, device=device, dtype=torch.float32)
412
+ # We want fp32 here as well since inv_freq will be multiplied with t, and the output
413
+ # will be large. Having it in bf16 will lose a lot of precision and cause the
414
+ # cos & sin output to change significantly.
415
+ # We want to recompute self.inv_freq if it was not loaded in fp32
416
+ if self.inv_freq.dtype != torch.float32:
417
+ inv_freq = self._compute_inv_freq(device=device)
418
+ else:
419
+ inv_freq = self.inv_freq
420
+ else:
421
+ t = torch.arange(seqlen, device=device, dtype=self.inv_freq.dtype)
422
+ inv_freq = self.inv_freq
423
+ # Don't do einsum, it converts fp32 to fp16 under AMP
424
+ # freqs = torch.einsum("i,j->ij", t, self.inv_freq)
425
+ freqs = torch.outer(t, inv_freq)
426
+ if self.scale is None:
427
+ self._cos_cached = torch.cos(freqs).to(dtype)
428
+ self._sin_cached = torch.sin(freqs).to(dtype)
429
+ else:
430
+ power = (
431
+ torch.arange(seqlen, dtype=self.scale.dtype, device=self.scale.device)
432
+ - seqlen // 2
433
+ ) / self.scale_base
434
+ scale = self.scale.to(device=power.device) ** rearrange(power, "s -> s 1")
435
+ # We want the multiplication by scale to happen in fp32
436
+ self._cos_cached = (torch.cos(freqs) * scale).to(dtype)
437
+ self._sin_cached = (torch.sin(freqs) * scale).to(dtype)
438
+ self._cos_k_cached = (torch.cos(freqs) / scale).to(dtype)
439
+ self._sin_k_cached = (torch.sin(freqs) / scale).to(dtype)
440
+
441
+ def forward(
442
+ self,
443
+ q: torch.Tensor,
444
+ k: torch.Tensor,
445
+ seqlen_offset: int | torch.Tensor = 0,
446
+ cu_seqlens: torch.Tensor | None = None,
447
+ max_seqlen: int | None = None,
448
+ ) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
449
+ """
450
+ q: [B, T, H, D]
451
+ k: [B, T, H, D]
452
+ seqlen_offset:
453
+ [N] or int.
454
+ Each sequence in x is shifted by this amount.
455
+ Most commonly used in inference when we have KV cache.
456
+ cu_seqlens: [N + 1] or None
457
+ max_seqlen: int
458
+ """
459
+ if max_seqlen is not None:
460
+ self._update_cos_sin_cache(max_seqlen, device=q.device, dtype=q.dtype)
461
+ elif isinstance(seqlen_offset, int):
462
+ self._update_cos_sin_cache(q.shape[1] + seqlen_offset, device=q.device, dtype=q.dtype)
463
+ if self.scale is None:
464
+ q = rotary_embedding(
465
+ q,
466
+ self._cos_cached,
467
+ self._sin_cached,
468
+ interleaved=self.interleaved,
469
+ seqlen_offsets=seqlen_offset,
470
+ cu_seqlens=cu_seqlens,
471
+ )
472
+ k = rotary_embedding(
473
+ k,
474
+ self._cos_cached,
475
+ self._sin_cached,
476
+ interleaved=self.interleaved,
477
+ seqlen_offsets=seqlen_offset,
478
+ cu_seqlens=cu_seqlens,
479
+ )
480
+
481
+ else:
482
+ q = rotary_embedding(
483
+ q,
484
+ self._cos_cached,
485
+ self._sin_cached,
486
+ interleaved=self.interleaved,
487
+ seqlen_offsets=seqlen_offset,
488
+ cu_seqlens=cu_seqlens,
489
+ )
490
+ k = rotary_embedding(
491
+ k,
492
+ self._cos_k_cached,
493
+ self._sin_k_cached,
494
+ interleaved=self.interleaved,
495
+ seqlen_offsets=seqlen_offset,
496
+ cu_seqlens=cu_seqlens,
497
+ )
498
+
499
+ return q, k
code/flash-linear-attention/fla/modules/token_shift.py ADDED
@@ -0,0 +1,545 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+
3
+ import torch
4
+ import triton
5
+ import triton.language as tl
6
+
7
+ from fla.ops.utils import prepare_chunk_indices
8
+ from fla.utils import autotune_cache_kwargs, get_multiprocessor_count, input_guard, is_amd, tensor_cache
9
+
10
+ NUM_WARPS_AUTOTUNE = [2, 4, 8, 16] if is_amd else [2, 4, 8, 16, 32]
11
+
12
+
13
+ def token_shift_ref(
14
+ x: torch.Tensor,
15
+ cu_seqlens: torch.Tensor | None = None,
16
+ ) -> torch.Tensor:
17
+ if cu_seqlens is not None:
18
+ # Variable length mode with cu_seqlens
19
+ assert x.dim() == 3, "Input must be [B, T, D]"
20
+ B, T, D = x.shape
21
+ assert B == 1, "Batch size must be 1 when using cu_seqlens"
22
+
23
+ result = torch.zeros_like(x)
24
+ N = cu_seqlens.shape[0] - 1
25
+
26
+ for i in range(N):
27
+ start = cu_seqlens[i].item()
28
+ end = cu_seqlens[i+1].item()
29
+ seq_len = end - start
30
+
31
+ if seq_len <= 1:
32
+ # For sequences of length 1 or 0, delta is simply -x
33
+ result[0, start:end] = -x[0, start:end]
34
+ else:
35
+ # For longer sequences, handle padding manually
36
+ shifted = torch.zeros_like(x[0, start:end])
37
+ shifted[1:] = x[0, start:end-1]
38
+ delta = shifted - x[0, start:end]
39
+ result[0, start:end] = delta
40
+
41
+ return result
42
+ else:
43
+ time_shift = torch.nn.ZeroPad2d((0, 0, 1, -1))
44
+ shifted = time_shift(x)
45
+ delta = shifted - x
46
+ return delta
47
+
48
+
49
+ @triton.heuristics({
50
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
51
+ 'USE_INITIAL_STATE': lambda args: args['cache'] is not None,
52
+ })
53
+ @triton.autotune(
54
+ configs=[
55
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
56
+ for num_warps in NUM_WARPS_AUTOTUNE
57
+ for num_stages in [1, 2, 3]
58
+ ],
59
+ key=['BD'],
60
+ **autotune_cache_kwargs,
61
+ )
62
+ @triton.jit
63
+ def token_shift_fwd_kernel_short(
64
+ x,
65
+ y,
66
+ cu_seqlens,
67
+ cache,
68
+ cache_out,
69
+ T,
70
+ D: tl.constexpr,
71
+ BD: tl.constexpr,
72
+ IS_VARLEN: tl.constexpr,
73
+ USE_INITIAL_STATE: tl.constexpr,
74
+ STORE_FINAL_STATE: tl.constexpr,
75
+ IS_DECODE: tl.constexpr,
76
+ ):
77
+ i_b, i_t = tl.program_id(0), tl.program_id(1)
78
+
79
+ if IS_VARLEN:
80
+ i_n = i_b
81
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
82
+ g_t = i_t + bos
83
+
84
+ if g_t >= eos:
85
+ return
86
+
87
+ is_first_pos = (i_t == 0)
88
+ is_last_pos = (g_t == eos - 1)
89
+ else:
90
+ g_t = i_t
91
+ is_first_pos = (g_t == 0)
92
+ is_last_pos = (g_t == T - 1)
93
+
94
+ o_d = tl.arange(0, BD)
95
+ m_d = o_d < D
96
+
97
+ if IS_VARLEN:
98
+ base_offset = g_t * D + o_d
99
+ else:
100
+ base_offset = i_b * T*D + g_t * D + o_d
101
+
102
+ b_x = tl.load(x + base_offset, mask=m_d)
103
+ if IS_VARLEN:
104
+ cache_offset = i_n * D + o_d # i_n is seq index
105
+ else:
106
+ cache_offset = i_b * D + o_d # i_b is batch index
107
+
108
+ if IS_DECODE and USE_INITIAL_STATE:
109
+ b_cache = tl.load(cache + cache_offset, mask=m_d)
110
+ delta = b_cache - b_x
111
+ tl.store(y + base_offset, delta, mask=m_d)
112
+ if STORE_FINAL_STATE:
113
+ tl.store(cache_out + cache_offset, b_x, mask=m_d)
114
+ return
115
+
116
+ if is_first_pos:
117
+ # First position in sequence: delta = -hidden_states
118
+ if USE_INITIAL_STATE:
119
+ # cache shape: [N, D]
120
+ b_cache = tl.load(cache + cache_offset, mask=m_d)
121
+ delta = b_cache - b_x
122
+ tl.store(y + base_offset, delta, mask=m_d)
123
+ else:
124
+ tl.store(y + base_offset, -b_x, mask=m_d)
125
+ return
126
+
127
+ # Other positions: delta = prev - curr
128
+ if IS_VARLEN:
129
+ prev_offset = (g_t-1) * D + o_d
130
+ else:
131
+ prev_offset = i_b * T*D + (g_t-1) * D + o_d
132
+
133
+ prev_values = tl.load(x + prev_offset, mask=m_d)
134
+ delta = prev_values - b_x
135
+ tl.store(y + base_offset, delta, mask=m_d)
136
+ if STORE_FINAL_STATE:
137
+ if is_last_pos:
138
+ tl.store(cache_out + cache_offset, b_x, mask=m_d)
139
+
140
+
141
+ @triton.heuristics({
142
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
143
+ 'USE_INITIAL_STATE': lambda args: args['cache'] is not None,
144
+ })
145
+ @triton.autotune(
146
+ configs=[
147
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
148
+ for num_warps in NUM_WARPS_AUTOTUNE
149
+ for num_stages in [1, 2, 3]
150
+ ],
151
+ key=['BD', 'NB'],
152
+ **autotune_cache_kwargs,
153
+ )
154
+ @triton.jit
155
+ def token_shift_fwd_kernel_long(
156
+ x,
157
+ y,
158
+ cu_seqlens,
159
+ chunk_indices,
160
+ cache,
161
+ cache_out,
162
+ T,
163
+ D: tl.constexpr,
164
+ BD: tl.constexpr,
165
+ BT: tl.constexpr,
166
+ NB: tl.constexpr,
167
+ IS_VARLEN: tl.constexpr,
168
+ USE_INITIAL_STATE: tl.constexpr,
169
+ STORE_FINAL_STATE: tl.constexpr,
170
+ ):
171
+ i_d, i_t, i_b = tl.program_id(0), tl.program_id(1), tl.program_id(2)
172
+
173
+ if IS_VARLEN:
174
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), \
175
+ tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
176
+ bos, eos = tl.load(cu_seqlens + i_n), tl.load(cu_seqlens + i_n + 1)
177
+ t_start = i_t * BT
178
+ t_end = tl.minimum(t_start + BT, eos - bos)
179
+ else:
180
+ i_n = i_b
181
+ bos, eos = i_b * T, (i_b + 1) * T
182
+ t_start = i_t * BT
183
+ t_end = tl.minimum(t_start + BT, T)
184
+
185
+ o_d = i_d * BD + tl.arange(0, BD)
186
+ m_d = o_d < D
187
+
188
+ for t in range(t_start, t_end):
189
+ global_t = bos + t
190
+ offset = global_t * D + o_d
191
+ b_x = tl.load(x + offset, mask=m_d)
192
+ is_first = (global_t == bos)
193
+ if is_first:
194
+ if USE_INITIAL_STATE:
195
+ # cache shape: [N, D]
196
+ cache_off = i_n * D + o_d if IS_VARLEN else i_b * D + o_d
197
+ b_cache = tl.load(cache + cache_off, mask=m_d)
198
+ delta = b_cache - b_x
199
+ else:
200
+ delta = -b_x
201
+ else:
202
+ prev_off = offset - D
203
+ b_prev = tl.load(x + prev_off, mask=m_d)
204
+ delta = b_prev - b_x
205
+
206
+ tl.store(y + offset, delta, mask=m_d)
207
+
208
+ if STORE_FINAL_STATE:
209
+ if global_t == eos - 1:
210
+ cache_out_off = i_n * D + o_d if IS_VARLEN else i_b * D + o_d
211
+ tl.store(cache_out + cache_out_off, b_x, mask=m_d)
212
+
213
+
214
+ @triton.heuristics({
215
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
216
+ 'USE_INITIAL_STATE': lambda args: args['grad_cache_out'] is not None,
217
+ 'HAS_DCACHE': lambda args: args['grad_cache_in'] is not None,
218
+ })
219
+ @triton.autotune(
220
+ configs=[
221
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
222
+ for num_warps in NUM_WARPS_AUTOTUNE
223
+ for num_stages in [1, 2, 3]
224
+ ],
225
+ key=['BD'],
226
+ **autotune_cache_kwargs,
227
+ )
228
+ @triton.jit
229
+ def token_shift_bwd_kernel_short(
230
+ dx,
231
+ dy,
232
+ cu_seqlens,
233
+ grad_cache_in,
234
+ grad_cache_out,
235
+ T,
236
+ D: tl.constexpr,
237
+ BD: tl.constexpr,
238
+ IS_VARLEN: tl.constexpr,
239
+ USE_INITIAL_STATE: tl.constexpr,
240
+ HAS_DCACHE: tl.constexpr,
241
+ ):
242
+ i_b, i_t = tl.program_id(0), tl.program_id(1)
243
+
244
+ if IS_VARLEN:
245
+ i_n = i_b
246
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
247
+ g_t = i_t + bos
248
+ if g_t >= eos:
249
+ return
250
+ is_first_pos = (g_t == bos)
251
+ is_last_pos = (g_t == eos - 1)
252
+ else:
253
+ g_t = i_t
254
+ is_first_pos = (g_t == 0)
255
+ is_last_pos = (g_t == T - 1)
256
+
257
+ o_d = tl.arange(0, BD)
258
+ m_d = o_d < D
259
+
260
+ if IS_VARLEN:
261
+ base_offset = g_t * D + o_d
262
+ # This should not be used for varlen
263
+ cache_off = i_n * D + o_d
264
+ else:
265
+ base_offset = i_b * T * D + g_t * D + o_d
266
+ cache_off = i_b * D + o_d
267
+
268
+ b_dy = tl.load(dy + base_offset, mask=m_d)
269
+
270
+ if is_last_pos:
271
+ # grad = -grad_delta[t] + grad_cache_in(from next rank)
272
+ if HAS_DCACHE:
273
+ b_dy_cache = tl.load(grad_cache_in + cache_off, mask=m_d)
274
+ b_dx = -b_dy + b_dy_cache
275
+ else:
276
+ b_dx = -b_dy
277
+ else:
278
+ # grad = -grad_delta[t] + grad_delta[t+1]
279
+ if IS_VARLEN:
280
+ next_offset = (g_t + 1) * D + o_d
281
+ else:
282
+ next_offset = i_b * T * D + (g_t + 1) * D + o_d
283
+ b_dx = -b_dy + tl.load(dy + next_offset, mask=m_d)
284
+
285
+ tl.store(dx + base_offset, b_dx, mask=m_d)
286
+
287
+ if USE_INITIAL_STATE:
288
+ if is_first_pos:
289
+ tl.store(grad_cache_out + cache_off, b_dy, mask=m_d)
290
+
291
+
292
+ @triton.heuristics({
293
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
294
+ 'USE_INITIAL_STATE': lambda args: args['grad_cache_out'] is not None,
295
+ 'HAS_DCACHE': lambda args: args['grad_cache_in'] is not None,
296
+ })
297
+ @triton.autotune(
298
+ configs=[
299
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
300
+ for num_warps in NUM_WARPS_AUTOTUNE
301
+ for num_stages in [1, 2, 3]
302
+ ],
303
+ key=['BD', 'NB'],
304
+ **autotune_cache_kwargs,
305
+ )
306
+ @triton.jit
307
+ def token_shift_bwd_kernel_long(
308
+ dx,
309
+ dy,
310
+ cu_seqlens,
311
+ chunk_indices,
312
+ grad_cache_in,
313
+ grad_cache_out,
314
+ T,
315
+ D: tl.constexpr,
316
+ BD: tl.constexpr,
317
+ BT: tl.constexpr,
318
+ NB: tl.constexpr,
319
+ IS_VARLEN: tl.constexpr,
320
+ USE_INITIAL_STATE: tl.constexpr,
321
+ HAS_DCACHE: tl.constexpr,
322
+ ):
323
+ i_d, i_t_blk, i_b = tl.program_id(0), tl.program_id(1), tl.program_id(2)
324
+
325
+ if IS_VARLEN:
326
+ i_n, i_t_blk = tl.load(chunk_indices + i_t_blk * 2).to(tl.int32), \
327
+ tl.load(chunk_indices + i_t_blk * 2 + 1).to(tl.int32)
328
+ bos, eos = tl.load(cu_seqlens + i_n), tl.load(cu_seqlens + i_n + 1)
329
+ t_start = i_t_blk * BT
330
+ t_end = tl.minimum(t_start + BT, eos - bos)
331
+ else:
332
+ bos, eos = i_b * T, (i_b + 1) * T
333
+ t_start = i_t_blk * BT
334
+ t_end = tl.minimum(t_start + BT, T)
335
+
336
+ o_d = i_d * BD + tl.arange(0, BD)
337
+ m_d = o_d < D
338
+ cache_off = i_n * D + o_d if IS_VARLEN else i_b * D + o_d
339
+
340
+ for t in range(t_start, t_end):
341
+ global_t = bos + t
342
+ offset = global_t * D + o_d
343
+ b_dy = tl.load(dy + offset, mask=m_d)
344
+
345
+ if global_t == eos - 1:
346
+ if HAS_DCACHE:
347
+ b_dy_cache = tl.load(grad_cache_in + cache_off, mask=m_d)
348
+ b_dx = -b_dy + b_dy_cache
349
+ else:
350
+ b_dx = -b_dy
351
+ else:
352
+ next_off = offset + D
353
+ b_dx = -b_dy + tl.load(dy + next_off, mask=m_d)
354
+
355
+ tl.store(dx + offset, b_dx, mask=m_d)
356
+
357
+ if USE_INITIAL_STATE:
358
+ if global_t == bos:
359
+ tl.store(grad_cache_out + cache_off, b_dy, mask=m_d)
360
+
361
+
362
+ @tensor_cache
363
+ def prepare_maxlens(cu_seqlens: torch.LongTensor) -> int:
364
+ return torch.max(cu_seqlens[1:] - cu_seqlens[:-1]).item()
365
+
366
+
367
+ def token_shift_fwd(
368
+ x: torch.Tensor,
369
+ cu_seqlens: torch.Tensor | None = None,
370
+ cache: torch.Tensor | None = None,
371
+ output_cache: bool = False,
372
+ ) -> torch.Tensor:
373
+ B, T, D = x.shape
374
+ y = torch.empty_like(x)
375
+ use_short_kernel = T <= 4096
376
+
377
+ if cu_seqlens is not None:
378
+ T = prepare_maxlens(cu_seqlens)
379
+ N = len(cu_seqlens) - 1
380
+ else:
381
+ N = B
382
+
383
+ if output_cache:
384
+ cache_out = torch.empty((N, D), device=x.device, dtype=x.dtype)
385
+ else:
386
+ cache_out = None
387
+
388
+ if use_short_kernel:
389
+ if cu_seqlens is not None:
390
+ N = len(cu_seqlens) - 1
391
+ else:
392
+ N = B
393
+ BD = triton.next_power_of_2(D)
394
+ grid = (N, T)
395
+ IS_DECODE = T == 1 or (B == 1 and T == N)
396
+ token_shift_fwd_kernel_short[grid](
397
+ x=x,
398
+ y=y,
399
+ cu_seqlens=cu_seqlens,
400
+ cache=cache,
401
+ cache_out=cache_out,
402
+ T=T,
403
+ D=D,
404
+ BD=BD,
405
+ STORE_FINAL_STATE=output_cache,
406
+ IS_DECODE=IS_DECODE,
407
+ )
408
+ else:
409
+ BT = min(64, triton.next_power_of_2(triton.cdiv(max(16, B*T), get_multiprocessor_count(x.device.index))))
410
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
411
+ NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT)
412
+
413
+ BD = triton.next_power_of_2(D)
414
+ NB = triton.cdiv(B*T, 1024)
415
+
416
+ def grid(meta): return (triton.cdiv(D, meta['BD']), NT, N)
417
+ token_shift_fwd_kernel_long[grid](
418
+ x,
419
+ y,
420
+ cu_seqlens,
421
+ chunk_indices,
422
+ cache,
423
+ cache_out,
424
+ T,
425
+ D=D,
426
+ BD=BD,
427
+ BT=BT,
428
+ NB=NB,
429
+ STORE_FINAL_STATE=output_cache,
430
+ )
431
+
432
+ return y, N, T, use_short_kernel, cache_out
433
+
434
+
435
+ def token_shift_bwd(
436
+ dy: torch.Tensor,
437
+ N: int,
438
+ T: int,
439
+ dcache: torch.Tensor | None = None,
440
+ cu_seqlens: torch.Tensor | None = None,
441
+ use_short_kernel: bool = True,
442
+ has_init_cache: bool = False,
443
+ ) -> torch.Tensor:
444
+ D = dy.shape[2]
445
+ BD = triton.next_power_of_2(D)
446
+ dx = torch.empty_like(dy)
447
+ if has_init_cache:
448
+ grad_cache_out = torch.empty((N, D), device=dy.device, dtype=dy.dtype)
449
+ else:
450
+ grad_cache_out = None
451
+ if use_short_kernel:
452
+ grid = (N, T)
453
+ token_shift_bwd_kernel_short[grid](
454
+ dy=dy,
455
+ dx=dx,
456
+ cu_seqlens=cu_seqlens,
457
+ grad_cache_in=dcache,
458
+ grad_cache_out=grad_cache_out,
459
+ T=T,
460
+ D=D,
461
+ BD=BD,
462
+ )
463
+ else:
464
+ BT = min(64, triton.next_power_of_2(triton.cdiv(max(16, dy.numel() // D),
465
+ get_multiprocessor_count(dy.device.index))))
466
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
467
+ NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT)
468
+ NB = triton.cdiv(N * dy.shape[1], 1024)
469
+ BD = triton.next_power_of_2(D)
470
+
471
+ def grid(meta): return (triton.cdiv(D, meta['BD']), NT, N)
472
+ token_shift_bwd_kernel_long[grid](
473
+ dx,
474
+ dy,
475
+ cu_seqlens,
476
+ chunk_indices,
477
+ dcache,
478
+ grad_cache_out,
479
+ T,
480
+ D=D,
481
+ BD=BD,
482
+ BT=BT,
483
+ NB=NB,
484
+ )
485
+ return dx, grad_cache_out
486
+
487
+
488
+ class TokenShift(torch.autograd.Function):
489
+
490
+ @staticmethod
491
+ @input_guard
492
+ def forward(ctx, x: torch.Tensor, cu_seqlens: torch.Tensor | None = None,
493
+ cache: torch.Tensor | None = None, output_cache: bool = False):
494
+ output, N, T, use_short_kernel, cache_out = token_shift_fwd(x, cu_seqlens, cache, output_cache)
495
+ ctx.cu_seqlens = cu_seqlens
496
+ ctx.N = N
497
+ ctx.T = T
498
+ ctx.use_short_kernel = use_short_kernel
499
+ ctx.has_cache = cache is not None
500
+ return output, cache_out
501
+
502
+ @staticmethod
503
+ @input_guard
504
+ def backward(ctx, dy: torch.Tensor, dcache: torch.Tensor | None = None):
505
+ dx, grad_cache = token_shift_bwd(dy, ctx.N, ctx.T, dcache, ctx.cu_seqlens,
506
+ ctx.use_short_kernel, ctx.has_cache)
507
+ return dx, None, grad_cache, None
508
+
509
+
510
+ def token_shift(
511
+ x: torch.Tensor,
512
+ cu_seqlens: torch.LongTensor | None = None,
513
+ cache: torch.Tensor | None = None,
514
+ output_cache: bool = False,
515
+ ):
516
+ """
517
+ Token-shift operation implemented with Triton kernels.
518
+
519
+ Args:
520
+ x: Input tensor of shape [B, T, D] (or [1, T, D] when `cu_seqlens` is supplied).
521
+ cu_seqlens: Optional cumulative sequence lengths of shape [B + 1].
522
+ When supplied, `x.shape[0]` must be 1 and `x.dim()` must be 3.
523
+ cache: Optional cache tensor of shape [N, D] that holds the last token
524
+ from the previous call.
525
+ output_cache: Whether to return the updated cache alongside the output.
526
+ In previous versions this parameter did not exist and the
527
+ cache was always dropped; to preserve backward compatibility
528
+ the default is False.
529
+
530
+ Returns:
531
+ output: Tensor of shape [B, T, D] after applying the token-shift.
532
+
533
+ cache_out: Tensor of shape [B, 1, D] containing the last token that
534
+ should be fed as `cache` in the next call. Only returned
535
+ when `output_cache=True`.
536
+ """
537
+ if cu_seqlens is not None:
538
+ assert x.dim() == 3, "Input must be [B, T, D]"
539
+ assert x.shape[0] == 1, "Batch size must be 1 when using cu_seqlens"
540
+
541
+ output, cache_out = TokenShift.apply(x, cu_seqlens, cache, output_cache)
542
+ if output_cache:
543
+ return output, cache_out
544
+ else:
545
+ return output
code/flash-linear-attention/fla/ops/__init__.py ADDED
@@ -0,0 +1,54 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from .abc import chunk_abc
3
+ from .attn import parallel_attn
4
+ from .based import fused_chunk_based, parallel_based
5
+ from .comba import chunk_comba, fused_recurrent_comba
6
+ from .delta_rule import chunk_delta_rule, fused_chunk_delta_rule, fused_recurrent_delta_rule
7
+ from .forgetting_attn import parallel_forgetting_attn
8
+ from .gated_delta_rule import chunk_gated_delta_rule, fused_recurrent_gated_delta_rule
9
+ from .generalized_delta_rule import (
10
+ chunk_dplr_delta_rule,
11
+ chunk_iplr_delta_rule,
12
+ fused_recurrent_dplr_delta_rule,
13
+ fused_recurrent_iplr_delta_rule,
14
+ )
15
+ from .gla import chunk_gla, fused_chunk_gla, fused_recurrent_gla
16
+ from .gsa import chunk_gsa, fused_recurrent_gsa
17
+ from .hgrn import fused_recurrent_hgrn
18
+ from .kda import chunk_kda, fused_recurrent_kda
19
+ from .lightning_attn import chunk_lightning_attn, fused_recurrent_lightning_attn
20
+ from .linear_attn import chunk_linear_attn, fused_chunk_linear_attn, fused_recurrent_linear_attn
21
+ from .log_linear_attn import chunk_log_linear_attn
22
+ from .mesa_net import chunk_mesa_net
23
+ from .nsa import parallel_nsa
24
+ from .path_attn import parallel_path_attn
25
+ from .retention import chunk_retention, fused_chunk_retention, fused_recurrent_retention, parallel_retention
26
+ from .rwkv6 import chunk_rwkv6, fused_recurrent_rwkv6
27
+ from .rwkv7 import chunk_rwkv7, fused_recurrent_rwkv7
28
+ from .simple_gla import chunk_simple_gla, fused_chunk_simple_gla, fused_recurrent_simple_gla, parallel_simple_gla
29
+
30
+ __all__ = [
31
+ 'chunk_abc',
32
+ 'parallel_attn',
33
+ 'fused_chunk_based', 'parallel_based',
34
+ 'chunk_delta_rule', 'fused_chunk_delta_rule', 'fused_recurrent_delta_rule',
35
+ 'parallel_forgetting_attn',
36
+ 'chunk_gated_delta_rule', 'fused_recurrent_gated_delta_rule',
37
+ 'chunk_comba', 'fused_recurrent_comba',
38
+ 'chunk_dplr_delta_rule', 'chunk_iplr_delta_rule',
39
+ 'fused_recurrent_dplr_delta_rule', 'fused_recurrent_iplr_delta_rule',
40
+ 'chunk_kda', 'fused_recurrent_kda',
41
+ 'chunk_gla', 'fused_chunk_gla', 'fused_recurrent_gla',
42
+ 'chunk_gsa', 'fused_recurrent_gsa',
43
+ 'fused_recurrent_hgrn',
44
+ 'chunk_lightning_attn', 'fused_recurrent_lightning_attn',
45
+ 'chunk_linear_attn', 'fused_chunk_linear_attn', 'fused_recurrent_linear_attn',
46
+ 'chunk_log_linear_attn',
47
+ 'chunk_mesa_net',
48
+ 'parallel_nsa',
49
+ 'parallel_path_attn',
50
+ 'chunk_retention', 'fused_chunk_retention', 'fused_recurrent_retention', 'parallel_retention',
51
+ 'chunk_rwkv6', 'fused_recurrent_rwkv6',
52
+ 'chunk_rwkv7', 'fused_recurrent_rwkv7',
53
+ 'chunk_simple_gla', 'fused_chunk_simple_gla', 'fused_recurrent_simple_gla', 'parallel_simple_gla',
54
+ ]
code/flash-linear-attention/fla/ops/abc/__init__.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+
2
+ from .chunk import chunk_abc
3
+
4
+ __all__ = [
5
+ 'chunk_abc',
6
+ ]
code/flash-linear-attention/fla/ops/abc/chunk.py ADDED
@@ -0,0 +1,1115 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils import softmax_bwd, softmax_fwd
9
+ from fla.ops.utils.logcumsumexp import logcumsumexp_fwd_kernel
10
+ from fla.ops.utils.op import exp
11
+ from fla.utils import input_guard
12
+
13
+
14
+ @triton.jit(do_not_specialize=['T'])
15
+ def chunk_abc_fwd_kernel_h(
16
+ k,
17
+ v,
18
+ z,
19
+ h,
20
+ h0,
21
+ ht,
22
+ T,
23
+ K: tl.constexpr,
24
+ V: tl.constexpr,
25
+ BT: tl.constexpr,
26
+ BK: tl.constexpr,
27
+ BV: tl.constexpr,
28
+ NT: tl.constexpr,
29
+ NORMK: tl.constexpr,
30
+ USE_INITIAL_STATE: tl.constexpr,
31
+ STORE_FINAL_STATE: tl.constexpr,
32
+ ):
33
+ i_v, i_k, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
34
+
35
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
36
+ if USE_INITIAL_STATE:
37
+ p_h = tl.make_block_ptr(h0 + i_bh * K * V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
38
+ b_h += tl.load(p_h, boundary_check=(0, 1)).to(tl.float32)
39
+ if NORMK:
40
+ p_z0 = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), (i_k * BK,), (BK,), (0,))
41
+ else:
42
+ p_z0 = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), (i_v * BV,), (BV,), (0,))
43
+ b_zp = tl.load(p_z0).to(tl.float32)
44
+ for i_t in range(NT):
45
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (K, T), (1, K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
46
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
47
+ p_h = tl.make_block_ptr(h + i_bh * NT*K*V + i_t * K * V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
48
+
49
+ tl.store(p_h, b_h.to(p_h.dtype.element_ty), boundary_check=(0, 1))
50
+ # [BK, BT]
51
+ b_k = tl.load(p_k, boundary_check=(0, 1))
52
+ # [BT, BV]
53
+ b_v = tl.load(p_v, boundary_check=(0, 1))
54
+ if NORMK:
55
+ p_zc = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), ((i_t * BT + BT - 1) * K + i_k * BK,), (BK,), (0,))
56
+ # [BK,]
57
+ b_zc = tl.load(p_zc, boundary_check=(0,))
58
+ b_r, b_zp = exp(b_zp - b_zc), b_zc
59
+ # [BK, BV]
60
+ b_h = b_h * b_r[:, None]
61
+ b_k = exp(b_k - b_zc[:, None]).to(b_k.dtype)
62
+ else:
63
+ p_zc = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), ((i_t * BT + BT - 1) * V + i_v * BV,), (BV,), (0,))
64
+ # [BV,]
65
+ b_zc = tl.load(p_zc, boundary_check=(0,))
66
+ b_r, b_zp = exp(b_zp - b_zc), b_zc
67
+ # [BK, BV]
68
+ b_h = b_h * b_r[None, :]
69
+ b_v = exp(b_v - b_zc[None, :]).to(b_v.dtype)
70
+ # [BK, BV]
71
+ b_h += tl.dot(b_k, b_v, allow_tf32=False)
72
+
73
+ if STORE_FINAL_STATE:
74
+ p_h = tl.make_block_ptr(ht + i_bh * K * V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
75
+ tl.store(p_h, b_h.to(p_h.dtype.element_ty), boundary_check=(0, 1))
76
+
77
+
78
+ @triton.jit(do_not_specialize=['T'])
79
+ def chunk_abc_fwd_kernel_intra_K(
80
+ v,
81
+ z,
82
+ o,
83
+ A,
84
+ T,
85
+ V: tl.constexpr,
86
+ BT: tl.constexpr,
87
+ BC: tl.constexpr,
88
+ BV: tl.constexpr,
89
+ NC: tl.constexpr,
90
+ ):
91
+ i_v, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
92
+ i_t, i_i = i_c // NC, i_c % NC
93
+
94
+ p_z = tl.make_block_ptr(z + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_i * BC, i_v * BV), (BC, BV), (1, 0))
95
+ p_zn = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), ((i_t * BT + i_i * BC) * V + i_v * BV,), (BV,), (0,))
96
+ # [BV,]
97
+ b_zn = tl.load(p_zn, boundary_check=(0,))
98
+ # [BC, BV]
99
+ b_o = tl.zeros([BC, BV], dtype=tl.float32)
100
+ for i_j in range(0, i_i):
101
+ p_A = tl.make_block_ptr(A + i_bh * T * BT, (T, BT), (BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0))
102
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_j * BC, i_v * BV), (BC, BV), (1, 0))
103
+ # [BC, BV]
104
+ b_v = tl.load(p_v, boundary_check=(0, 1))
105
+ # [BC, BC]
106
+ b_A = tl.load(p_A, boundary_check=(0, 1))
107
+ b_o += tl.dot(b_A, exp(b_v - b_zn[None, :]).to(b_v.dtype), allow_tf32=False)
108
+ b_z = tl.load(p_z, boundary_check=(0, 1))
109
+ b_o *= exp(b_zn[None, :] - b_z)
110
+
111
+ o_i = tl.arange(0, BC)
112
+ o_A = i_bh * T * BT + (i_t * BT + i_i * BC + tl.arange(0, BC)) * BT + i_i * BC
113
+ m_A = (i_t * BT + i_i * BC + tl.arange(0, BC)) < T
114
+ for j in range(0, BC):
115
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T * V,), (1,), ((i_t * BT + i_i * BC + j) * V + i_v * BV,), (BV,), (0,))
116
+ # [BC,]
117
+ b_A = tl.load(A + o_A + j, mask=m_A, other=0)
118
+ # [BV,]
119
+ b_v = tl.load(p_v, boundary_check=(0,)).to(tl.float32)
120
+ # [BC, BV]
121
+ # avoid 0 * inf = inf
122
+ m_i = o_i[:, None] >= j
123
+ b_o += tl.where(m_i, b_A[:, None] * exp(b_v[None, :] - b_z), 0)
124
+ p_o = tl.make_block_ptr(o + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_i * BC, i_v * BV), (BC, BV), (1, 0))
125
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
126
+
127
+
128
+ @triton.jit(do_not_specialize=['T'])
129
+ def chunk_abc_fwd_kernel_K(
130
+ q,
131
+ k,
132
+ z,
133
+ h,
134
+ o,
135
+ A,
136
+ scale,
137
+ T,
138
+ K: tl.constexpr,
139
+ V: tl.constexpr,
140
+ BT: tl.constexpr,
141
+ BK: tl.constexpr,
142
+ BV: tl.constexpr,
143
+ NT: tl.constexpr,
144
+ ):
145
+ i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
146
+ i_p = tl.maximum(i_t * BT - 1, 0)
147
+
148
+ o_i = tl.arange(0, BT)
149
+ m_s = o_i[:, None] >= o_i[None, :]
150
+
151
+ b_o = tl.zeros([BT, BV], dtype=tl.float32)
152
+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
153
+ for i_k in range(tl.cdiv(K, BK)):
154
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
155
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (K, T), (1, K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
156
+ p_h = tl.make_block_ptr(h + i_bh * NT*K*V + i_t * K * V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
157
+
158
+ # [BT, BK]
159
+ b_q = tl.load(p_q, boundary_check=(0, 1))
160
+ b_q = (b_q * scale).to(b_q.dtype)
161
+ # [BK, BT]
162
+ b_k = tl.load(p_k, boundary_check=(0, 1))
163
+ # [BK, BV]
164
+ b_h = tl.load(p_h, boundary_check=(0, 1))
165
+ # [BT, BV]
166
+ b_o += tl.dot(b_q, b_h, allow_tf32=False)
167
+ # [BT, BT]
168
+ b_A += tl.dot(b_q, b_k, allow_tf32=False)
169
+ p_z = tl.make_block_ptr(z + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
170
+ p_o = tl.make_block_ptr(o + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
171
+ # [BT, BV]
172
+ b_z = tl.load(p_z, boundary_check=(0, 1))
173
+ # [BT, BV]
174
+ p_zp = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), (i_p * V + i_v * BV,), (BV,), (0,))
175
+ b_zp = tl.load(p_zp, boundary_check=(0,))
176
+ b_o = b_o * exp(b_zp[None, :] - b_z)
177
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
178
+
179
+ p_A = tl.make_block_ptr(A + i_bh * T * BT, (T, BT), (BT, 1), (i_t * BT, 0), (BT, BT), (1, 0))
180
+ # [BT, BT]
181
+ b_A = tl.where(m_s, b_A, 0.)
182
+ if i_v == 0:
183
+ tl.store(p_A, b_A.to(p_A.dtype.element_ty), boundary_check=(0, 1))
184
+
185
+
186
+ @triton.jit(do_not_specialize=['T'])
187
+ def chunk_abc_fwd_kernel_intra_V(
188
+ q,
189
+ k,
190
+ z,
191
+ A,
192
+ scale,
193
+ T,
194
+ K: tl.constexpr,
195
+ BT: tl.constexpr,
196
+ BC: tl.constexpr,
197
+ BK: tl.constexpr,
198
+ NC: tl.constexpr,
199
+ ):
200
+ i_k, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
201
+ i_t, i_i, i_j = i_c // (NC * NC), (i_c % (NC * NC)) // NC, (i_c % (NC * NC)) % NC
202
+ n_bh = tl.num_programs(2)
203
+
204
+ if i_i > i_j:
205
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0))
206
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (K, T), (1, K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC), (0, 1))
207
+ p_z = tl.make_block_ptr(z + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0))
208
+ p_A = tl.make_block_ptr(A + (i_k*n_bh+i_bh)*T*BT, (T, BT), (BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0))
209
+ p_zn = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), ((i_t * BT + i_i * BC) * K + i_k * BK,), (BK,), (0,))
210
+ # [BK,]
211
+ b_zn = tl.load(p_zn, boundary_check=(0,))
212
+ # [BC, BK]
213
+ b_q = tl.load(p_q, boundary_check=(0, 1))
214
+ b_z = tl.load(p_z, boundary_check=(0, 1))
215
+ b_q = (b_q * exp(b_zn[None, :] - b_z) * scale).to(b_q.dtype)
216
+ # [BK, BC]
217
+ b_k = tl.load(p_k, boundary_check=(0, 1))
218
+ b_k = exp(b_k - b_zn[:, None]).to(b_k.dtype)
219
+ # [BC, BC]
220
+ b_A = tl.dot(b_q, b_k, allow_tf32=False)
221
+ tl.store(p_A, b_A.to(A.dtype.element_ty), boundary_check=(0, 1))
222
+ elif i_i == i_j:
223
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0))
224
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T * K,), (1,), ((i_t * BT + i_j * BC) * K + i_k * BK,), (BK,), (0,))
225
+ p_z = tl.make_block_ptr(z + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0))
226
+ # [BC, BK]
227
+ b_q = tl.load(p_q, boundary_check=(0, 1))
228
+ b_z = tl.load(p_z, boundary_check=(0, 1))
229
+
230
+ o_i = tl.arange(0, BC)
231
+ o_A = (i_bh + i_k * n_bh) * T * BT + (i_t * BT + i_i * BC + tl.arange(0, BC)) * BT + i_j * BC
232
+ m_A = (i_t * BT + i_i * BC + tl.arange(0, BC)) < T
233
+ for j in range(0, BC):
234
+ # [BK,]
235
+ b_k = tl.load(p_k, boundary_check=(0,)).to(tl.float32)
236
+ # [BC,]
237
+ b_A = tl.sum(b_q * exp(b_k[None, :] - b_z) * scale, 1)
238
+ b_A = tl.where(o_i >= j, b_A, 0.)
239
+ tl.store(A + o_A + j, b_A.to(b_q.dtype), mask=m_A)
240
+
241
+ p_k = tl.advance(p_k, (K,))
242
+
243
+
244
+ @triton.jit(do_not_specialize=['T'])
245
+ def chunk_abc_fwd_kernel_V(
246
+ q,
247
+ v,
248
+ z,
249
+ h,
250
+ o,
251
+ A,
252
+ scale,
253
+ T,
254
+ K: tl.constexpr,
255
+ V: tl.constexpr,
256
+ BT: tl.constexpr,
257
+ BK: tl.constexpr,
258
+ BV: tl.constexpr,
259
+ NT: tl.constexpr,
260
+ ):
261
+ i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
262
+ i_p = tl.maximum(i_t * BT - 1, 0)
263
+
264
+ b_o = tl.zeros([BT, BV], dtype=tl.float32)
265
+ for i_k in range(tl.cdiv(K, BK)):
266
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
267
+ p_z = tl.make_block_ptr(z + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
268
+ p_h = tl.make_block_ptr(h + i_bh * NT*K*V + i_t * K * V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
269
+ p_zp = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), (i_p * K + i_k * BK,), (BK,), (0,))
270
+
271
+ # [BT, BK]
272
+ b_q = tl.load(p_q, boundary_check=(0, 1))
273
+ b_q = (b_q * scale).to(b_q.dtype)
274
+ # [BT, BK]
275
+ b_z = tl.load(p_z, boundary_check=(0, 1))
276
+ # [BT, BK]
277
+ b_zp = tl.load(p_zp, boundary_check=(0,))
278
+ b_q = (b_q * exp(b_zp[None, :] - b_z)).to(b_q.dtype)
279
+ # [BK, BV]
280
+ b_h = tl.load(p_h, boundary_check=(0, 1))
281
+ # works but dkw, owing to divine benevolence
282
+ # [BT, BV]
283
+ if i_k >= 0:
284
+ b_o += tl.dot(b_q, b_h, allow_tf32=False)
285
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
286
+ p_o = tl.make_block_ptr(o + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
287
+ p_A = tl.make_block_ptr(A + i_bh * T * BT, (T, BT), (BT, 1), (i_t * BT, 0), (BT, BT), (1, 0))
288
+ # [BT, BV]
289
+ b_v = tl.load(p_v, boundary_check=(0, 1))
290
+ # [BT, BT]
291
+ b_A = tl.load(p_A, boundary_check=(0, 1))
292
+ b_o += tl.dot(b_A.to(b_v.dtype), b_v, allow_tf32=False)
293
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
294
+
295
+
296
+ @triton.jit(do_not_specialize=['T'])
297
+ def chunk_abc_bwd_kernel_dh(
298
+ q,
299
+ z,
300
+ do,
301
+ dh,
302
+ scale,
303
+ T,
304
+ K: tl.constexpr,
305
+ V: tl.constexpr,
306
+ BT: tl.constexpr,
307
+ BK: tl.constexpr,
308
+ BV: tl.constexpr,
309
+ NT: tl.constexpr,
310
+ NORMK: tl.constexpr,
311
+ ):
312
+ i_k, i_v, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
313
+
314
+ b_dh = tl.zeros([BK, BV], dtype=tl.float32)
315
+ b_zp = tl.full([BK if NORMK else BV], float('inf'), dtype=tl.float32)
316
+ for i_t in range(NT - 1, -1, -1):
317
+ i_p = tl.maximum(i_t * BT - 1, 0)
318
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (K, T), (1, K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
319
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
320
+ p_dh = tl.make_block_ptr(dh + i_bh * NT*K*V + i_t * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
321
+
322
+ # [BK, BT]
323
+ b_q = tl.load(p_q, boundary_check=(0, 1))
324
+ b_q = (b_q * scale).to(b_q.dtype)
325
+ # [BT, BV]
326
+ b_do = tl.load(p_do, boundary_check=(0, 1))
327
+
328
+ tl.store(p_dh, b_dh.to(p_dh.dtype.element_ty), boundary_check=(0, 1))
329
+ if NORMK:
330
+ p_z = tl.make_block_ptr(z + i_bh * T*K, (K, T), (1, K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
331
+ p_zc = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), (i_p * K + i_k * BK,), (BK,), (0,))
332
+ # [BK,]
333
+ b_zc = tl.load(p_zc, boundary_check=(0,))
334
+ b_r, b_zp = exp(b_zc - b_zp), b_zc
335
+ # [BK, BT]
336
+ b_z = tl.load(p_z, boundary_check=(0, 1))
337
+ b_q = (b_q * exp(b_zc[:, None] - b_z)).to(b_q.dtype)
338
+ # [BK, BV]
339
+ b_dh = b_dh * b_r[:, None]
340
+ else:
341
+ p_z = tl.make_block_ptr(z + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
342
+ p_zc = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), (i_p * V + i_v * BV,), (BV,), (0,))
343
+ # [BV,]
344
+ b_zc = tl.load(p_zc, boundary_check=(0,))
345
+ b_r, b_zp = exp(b_zc - b_zp), b_zc
346
+ # [BT, BV]
347
+ b_z = tl.load(p_z, boundary_check=(0,))
348
+ b_do = (b_do * exp(b_zc[None, :] - b_z)).to(b_do.dtype)
349
+ # [BK, BV]
350
+ b_dh = b_dh * b_r[None, :]
351
+ # [BK, BV]
352
+ b_dh += tl.dot(b_q, b_do, allow_tf32=False)
353
+
354
+
355
+ @triton.jit(do_not_specialize=['T'])
356
+ def chunk_abc_bwd_kernel_V(
357
+ k,
358
+ v,
359
+ z,
360
+ h,
361
+ A,
362
+ do,
363
+ dh,
364
+ dq,
365
+ dk,
366
+ dv,
367
+ dA,
368
+ scale,
369
+ T,
370
+ K: tl.constexpr,
371
+ V: tl.constexpr,
372
+ BT: tl.constexpr,
373
+ BK: tl.constexpr,
374
+ BV: tl.constexpr,
375
+ NT: tl.constexpr,
376
+ ):
377
+ i_k, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
378
+ i_p = tl.maximum(i_t * BT - 1, 0)
379
+ n_bh = tl.num_programs(2)
380
+
381
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
382
+ p_zc = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), ((i_t * BT + BT - 1) * K + i_k * BK,), (BK,), (0,))
383
+ p_A = tl.make_block_ptr(A + i_bh * T * BT, (BT, T), (1, BT), (0, i_t * BT), (BT, BT), (0, 1))
384
+
385
+ # [BK,]
386
+ b_zc = tl.load(p_zc, boundary_check=(0,))
387
+ # [BT, BK]
388
+ b_k = tl.load(p_k, boundary_check=(0, 1))
389
+ b_k = exp(b_k - b_zc[None, :]).to(b_k.dtype)
390
+ # [BT, BT]
391
+ b_A = tl.load(p_A, boundary_check=(0, 1))
392
+
393
+ b_dq = tl.zeros([BT, BK], dtype=tl.float32)
394
+ b_dk = tl.zeros([BT, BK], dtype=tl.float32)
395
+ b_dA = tl.zeros([BT, BT], dtype=tl.float32)
396
+ for i_v in range(tl.cdiv(V, BV)):
397
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
398
+ p_h = tl.make_block_ptr(h + i_bh * NT*K*V + i_t * V * K, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))
399
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
400
+ p_dh = tl.make_block_ptr(dh + i_bh * NT*K*V + i_t * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
401
+ p_dv = tl.make_block_ptr(dv + (i_k*n_bh+i_bh) * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
402
+
403
+ # [BT, BV]
404
+ b_v = tl.load(p_v, boundary_check=(0, 1))
405
+ # [BV, BK]
406
+ b_h = tl.load(p_h, boundary_check=(0, 1))
407
+ # [BT, BV]
408
+ b_do = tl.load(p_do, boundary_check=(0, 1))
409
+ # [BK, BV]
410
+ b_dh = tl.load(p_dh, boundary_check=(0, 1))
411
+
412
+ # [BT, BV]
413
+ b_dv = tl.dot(b_k, b_dh, allow_tf32=False)
414
+ if i_k == 0:
415
+ b_dv += tl.dot(b_A.to(b_do.dtype), b_do, allow_tf32=False)
416
+ b_do = (b_do * scale).to(b_do.dtype)
417
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
418
+ # [BT, BT]
419
+ b_dA += tl.dot(b_do, tl.trans(b_v), allow_tf32=False)
420
+ # [BT, BK]
421
+ b_dq += tl.dot(b_do, b_h, allow_tf32=False)
422
+ # [BT, BK]
423
+ b_dk += tl.dot(b_v, tl.trans(b_dh), allow_tf32=False)
424
+ p_z = tl.make_block_ptr(z + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
425
+ p_zp = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), (i_p * K + i_k * BK,), (BK,), (0,))
426
+ # [BK,]
427
+ b_zp = tl.load(p_zp, boundary_check=(0,))
428
+ # [BT, BK]
429
+ b_z = tl.load(p_z, boundary_check=(0, 1))
430
+ b_z = exp(b_zp[None, :] - b_z)
431
+ # [BT, BK]
432
+ b_dq = b_dq * b_z
433
+ b_dk = b_dk * b_k
434
+
435
+ p_dq = tl.make_block_ptr(dq + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
436
+ p_dk = tl.make_block_ptr(dk + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
437
+ p_dA = tl.make_block_ptr(dA + i_bh * T * BT, (T, BT), (BT, 1), (i_t * BT, 0), (BT, BT), (1, 0))
438
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
439
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
440
+
441
+ o_i = tl.arange(0, BT)
442
+ m_s = o_i[:, None] >= o_i[None, :]
443
+ # [BT, BT]
444
+ b_dA = tl.where(m_s, b_dA, 0.).to(b_k.dtype)
445
+ if i_k == 0:
446
+ tl.store(p_dA, b_dA.to(p_dA.dtype.element_ty), boundary_check=(0, 1))
447
+
448
+
449
+ @triton.jit(do_not_specialize=['T'])
450
+ def chunk_abc_bwd_kernel_intra_V(
451
+ q,
452
+ k,
453
+ z,
454
+ dA,
455
+ dq,
456
+ dk,
457
+ T,
458
+ K: tl.constexpr,
459
+ BT: tl.constexpr,
460
+ BC: tl.constexpr,
461
+ BK: tl.constexpr,
462
+ NC: tl.constexpr,
463
+ ):
464
+ i_k, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
465
+ i_t, i_i = i_c // NC, i_c % NC
466
+
467
+ p_z = tl.make_block_ptr(z + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0))
468
+ p_zn = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), ((i_t * BT + i_i * BC) * K + i_k * BK,), (BK,), (0,))
469
+ # [BK,]
470
+ b_zn = tl.load(p_zn, boundary_check=(0,))
471
+ # [BC, BK]
472
+ b_z = tl.load(p_z, boundary_check=(0, 1))
473
+ b_zq = exp(b_zn[None, :] - b_z)
474
+ b_dq = tl.zeros([BC, BK], dtype=tl.float32)
475
+ for i_j in range(0, i_i):
476
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_j * BC, i_k * BK), (BC, BK), (1, 0))
477
+ p_dA = tl.make_block_ptr(dA + i_bh * T * BT, (T, BT), (BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0))
478
+ # [BC, BK]
479
+ b_k = tl.load(p_k, boundary_check=(0, 1))
480
+ b_kz = exp(b_k - b_zn[None, :]).to(b_k.dtype)
481
+ # [BC, BC]
482
+ b_dA = tl.load(p_dA, boundary_check=(0, 1))
483
+ # [BC, BK]
484
+ b_dq += tl.dot(b_dA, b_kz, allow_tf32=False)
485
+ b_dq *= b_zq
486
+
487
+ o_i = tl.arange(0, BC)
488
+ o_dA = i_bh * T * BT + (i_t * BT + i_i * BC + tl.arange(0, BC)) * BT + i_i * BC
489
+ m_dA = (i_t * BT + i_i * BC + tl.arange(0, BC)) < T
490
+ for j in range(0, BC):
491
+ p_kj = tl.make_block_ptr(k + i_bh * T*K, (T * K,), (1,), ((i_t * BT + i_i*BC+j) * K + i_k * BK,), (BK,), (0,))
492
+ # [BC,]
493
+ b_dA = tl.load(dA + o_dA + j, mask=m_dA, other=0)
494
+ # [BK,]
495
+ b_kj = tl.load(p_kj, boundary_check=(0,)).to(tl.float32)
496
+ # [BC, BK]
497
+ m_i = o_i[:, None] >= j
498
+ # [BC, BK]
499
+ b_dq += tl.where(m_i, b_dA[:, None] * exp(b_kj[None, :] - b_z), 0.)
500
+ p_dq = tl.make_block_ptr(dq + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0))
501
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
502
+
503
+ tl.debug_barrier()
504
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0))
505
+ p_zn = tl.make_block_ptr(z + i_bh * T*K, (T*K,), (1,), ((i_t * BT + i_i * BC + BC - 1) * K + i_k * BK,), (BK,), (0,))
506
+ # [BK,]
507
+ b_zn = tl.load(p_zn, boundary_check=(0,))
508
+ # [BC, BK]
509
+ b_k = tl.load(p_k, boundary_check=(0, 1))
510
+ b_kz = exp(b_k - b_zn[None, :])
511
+ b_dk = tl.zeros([BC, BK], dtype=tl.float32)
512
+ for i_j in range(i_i + 1, NC):
513
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_j * BC, i_k * BK), (BC, BK), (1, 0))
514
+ p_z = tl.make_block_ptr(z + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_j * BC, i_k * BK), (BC, BK), (1, 0))
515
+ p_dA = tl.make_block_ptr(dA + i_bh * T * BT, (T, BT), (BT, 1), (i_t * BT + i_j * BC, i_i * BC), (BC, BC), (1, 0))
516
+ # [BC, BK]
517
+ b_q = tl.load(p_q, boundary_check=(0, 1))
518
+ b_z = tl.load(p_z, boundary_check=(0, 1))
519
+ b_qz = (b_q * exp(b_zn[None, :] - b_z)).to(b_q.dtype)
520
+ # [BC, BC]
521
+ b_dA = tl.load(p_dA, boundary_check=(0, 1))
522
+ # [BC, BK]
523
+ b_dk += tl.dot(tl.trans(b_dA), b_qz, allow_tf32=False)
524
+ b_dk *= b_kz
525
+
526
+ o_dA = i_bh * T * BT + (i_t * BT + i_i * BC) * BT + i_i * BC + tl.arange(0, BC)
527
+ for j in range(0, BC):
528
+ p_qj = tl.make_block_ptr(q + i_bh * T*K, (T * K,), (1,), ((i_t * BT + i_i * BC + j) * K + i_k * BK,), (BK,), (0,))
529
+ p_zj = tl.make_block_ptr(z + i_bh * T*K, (T * K,), (1,), ((i_t * BT + i_i * BC + j) * K + i_k * BK,), (BK,), (0,))
530
+ # [BC,]
531
+ b_dA = tl.load(dA + o_dA + j * BT, mask=(i_t * BT + i_i * BC + j < T), other=0)
532
+ # [BK,]
533
+ b_qj = tl.load(p_qj, boundary_check=(0,)).to(tl.float32)
534
+ b_zj = tl.load(p_zj, boundary_check=(0,)).to(tl.float32)
535
+ # [BC, BK]
536
+ m_i = o_i[:, None] <= j
537
+ b_dk += tl.where(m_i, b_dA[:, None] * b_qj[None, :] * exp(b_k - b_zj[None, :]), 0.)
538
+ p_dk = tl.make_block_ptr(dk + i_bh * T*K, (T, K), (K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0))
539
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
540
+
541
+
542
+ @triton.jit(do_not_specialize=['T'])
543
+ def chunk_abc_bwd_kernel_intra_K(
544
+ v,
545
+ z,
546
+ do,
547
+ dA,
548
+ scale,
549
+ T,
550
+ V: tl.constexpr,
551
+ BT: tl.constexpr,
552
+ BC: tl.constexpr,
553
+ BV: tl.constexpr,
554
+ NC: tl.constexpr,
555
+ ):
556
+ i_v, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
557
+ i_t, i_i, i_j = i_c // (NC * NC), (i_c % (NC * NC)) // NC, (i_c % (NC * NC)) % NC
558
+ n_bh = tl.num_programs(2)
559
+
560
+ if i_i > i_j:
561
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (V, T), (1, V), (i_v * BV, i_t * BT + i_j * BC), (BV, BC), (0, 1))
562
+ p_z = tl.make_block_ptr(z + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_i * BC, i_v * BV), (BC, BV), (1, 0))
563
+ p_zn = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), ((i_t * BT + i_i * BC) * V + i_v * BV,), (BV,), (0,))
564
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_i * BC, i_v * BV), (BC, BV), (1, 0))
565
+ p_dA = tl.make_block_ptr(dA+(i_bh+i_v*n_bh)*T*BT, (T, BT), (BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0))
566
+ # [BV,]
567
+ b_zn = tl.load(p_zn, boundary_check=(0,))
568
+ # [BC, BV]
569
+ b_z = tl.load(p_z, boundary_check=(0, 1))
570
+ b_do = tl.load(p_do, boundary_check=(0, 1))
571
+ b_do = (b_do * exp(b_zn[None, :] - b_z) * scale).to(b_do.dtype)
572
+ # [BV, BC]
573
+ b_v = tl.load(p_v, boundary_check=(0, 1))
574
+ b_v = exp(b_v - b_zn[:, None]).to(b_v.dtype)
575
+ # [BC, BC]
576
+ b_dA = tl.dot(b_do, b_v, allow_tf32=False)
577
+ tl.store(p_dA, b_dA.to(dA.dtype.element_ty), boundary_check=(0, 1))
578
+ elif i_i == i_j:
579
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T * V,), (1,), ((i_t * BT + i_j * BC) * V + i_v * BV,), (BV,), (0,))
580
+ p_z = tl.make_block_ptr(z + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_i * BC, i_v * BV), (BC, BV), (1, 0))
581
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_i * BC, i_v * BV), (BC, BV), (1, 0))
582
+ # [BC, BV]
583
+ b_z = tl.load(p_z, boundary_check=(0, 1))
584
+ b_do = tl.load(p_do, boundary_check=(0, 1)) * scale
585
+
586
+ o_i = tl.arange(0, BC)
587
+ o_A = (i_bh + i_v * n_bh) * T * BT + (i_t * BT + i_i * BC + tl.arange(0, BC)) * BT + i_j * BC
588
+ m_A = (i_t * BT + i_i * BC + tl.arange(0, BC)) < T
589
+ for j in range(0, BC):
590
+ # [BV,]
591
+ b_v = tl.load(p_v, boundary_check=(0,)).to(tl.float32)
592
+ # [BC,]
593
+ b_dA = tl.sum(b_do * exp(b_v[None, :] - b_z), 1)
594
+ b_dA = tl.where(o_i >= j, b_dA, 0)
595
+ tl.store(dA + o_A + j, b_dA.to(b_do.dtype), mask=m_A)
596
+
597
+ p_v = tl.advance(p_v, (V,))
598
+
599
+
600
+ @triton.jit(do_not_specialize=['T'])
601
+ def chunk_abc_bwd_kernel_K(
602
+ q,
603
+ k,
604
+ v,
605
+ z,
606
+ h,
607
+ A,
608
+ do,
609
+ dh,
610
+ dq,
611
+ dk,
612
+ dv,
613
+ dA,
614
+ scale,
615
+ T,
616
+ K: tl.constexpr,
617
+ V: tl.constexpr,
618
+ BT: tl.constexpr,
619
+ BK: tl.constexpr,
620
+ BV: tl.constexpr,
621
+ NT: tl.constexpr,
622
+ ):
623
+ i_k, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
624
+ i_p = tl.maximum(i_t * BT - 1, 0)
625
+ n_bh = tl.num_programs(2)
626
+
627
+ o_i = tl.arange(0, BT)
628
+ m_s = o_i[:, None] >= o_i[None, :]
629
+
630
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
631
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
632
+ p_A = tl.make_block_ptr(A + (i_k*n_bh+i_bh) * T * BT, (T, BT ), (BT, 1), (i_t * BT, 0), (BT, BT), (1, 0))
633
+
634
+ # [BT, BK]
635
+ b_q = tl.load(p_q, boundary_check=(0, 1))
636
+ b_k = tl.load(p_k, boundary_check=(0, 1))
637
+ # [BT, BT]
638
+ b_A = tl.dot((b_q * scale).to(b_q.dtype), tl.trans(b_k), allow_tf32=False)
639
+ b_A = tl.where(m_s, b_A, 0.)
640
+ tl.store(p_A, b_A.to(p_A.dtype.element_ty), boundary_check=(0, 1))
641
+
642
+ b_dq = tl.zeros([BT, BK], dtype=tl.float32)
643
+ b_dk = tl.zeros([BT, BK], dtype=tl.float32)
644
+ for i_v in range(tl.cdiv(V, BV)):
645
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
646
+ p_z = tl.make_block_ptr(z + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
647
+ p_zp = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), (i_p * V + i_v * BV,), (BV,), (0,))
648
+ p_zc = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), ((i_t * BT + BT - 1) * V + i_v * BV,), (BV,), (0,))
649
+ p_h = tl.make_block_ptr(h + i_bh * NT*K*V + i_t * K*V, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))
650
+
651
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
652
+ p_dh = tl.make_block_ptr(dh + i_bh * NT*K*V + i_t * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
653
+ p_dv = tl.make_block_ptr(dv + (i_k*n_bh+i_bh) * T*V, (T, V), (V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
654
+
655
+ # [BV,]
656
+ b_zp = tl.load(p_zp, boundary_check=(0,))
657
+ b_zc = tl.load(p_zc, boundary_check=(0,))
658
+ # [BT, BV]
659
+ b_v = tl.load(p_v, boundary_check=(0, 1))
660
+ b_v = exp(b_v - b_zc[None, :]).to(b_v.dtype)
661
+ b_z = tl.load(p_z, boundary_check=(0, 1))
662
+ b_z = exp(b_zp[None, :] - b_z)
663
+ # [BV, BK]
664
+ b_h = tl.load(p_h, boundary_check=(0, 1))
665
+ # [BT, BV]
666
+ b_do = tl.load(p_do, boundary_check=(0, 1))
667
+ b_do = (b_do * b_z * scale).to(b_do.dtype)
668
+ # [BK, BV]
669
+ b_dh = tl.load(p_dh, boundary_check=(0, 1))
670
+
671
+ # [BT, BK]
672
+ b_dq += tl.dot(b_do, b_h, allow_tf32=False)
673
+ b_dk += tl.dot(b_v, tl.trans(b_dh), allow_tf32=False)
674
+ # [BT, BV]
675
+ b_dv = b_v * tl.dot(b_k, b_dh, allow_tf32=False)
676
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
677
+ p_dA = tl.make_block_ptr(dA + i_bh * T * BT, (T, BT ), (BT, 1), (i_t * BT, 0), (BT, BT), (1, 0))
678
+ # [BT, BT]
679
+ b_dA = tl.load(p_dA, boundary_check=(0, 1))
680
+ # [BT, BK]
681
+ b_dq += tl.dot(b_dA, b_k, allow_tf32=False)
682
+ b_dk += tl.dot(tl.trans(b_dA).to(b_k.dtype), b_q, allow_tf32=False)
683
+
684
+ p_dq = tl.make_block_ptr(dq + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
685
+ p_dk = tl.make_block_ptr(dk + i_bh * T*K, (T, K), (K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
686
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
687
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
688
+
689
+
690
+ @triton.jit(do_not_specialize=['T'])
691
+ def chunk_abc_bwd_kernel_intra_KV(
692
+ v,
693
+ z,
694
+ A,
695
+ do,
696
+ dv,
697
+ T,
698
+ V: tl.constexpr,
699
+ BT: tl.constexpr,
700
+ BC: tl.constexpr,
701
+ BV: tl.constexpr,
702
+ NC: tl.constexpr,
703
+ ):
704
+ i_v, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
705
+ i_t, i_i = i_c // NC, i_c % NC
706
+
707
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_i * BC, i_v * BV), (BC, BV), (1, 0))
708
+ p_zn = tl.make_block_ptr(z + i_bh * T*V, (T*V,), (1,), ((i_t * BT + i_i * BC + BC - 1) * V + i_v * BV,), (BV,), (0,))
709
+ # [BV,]
710
+ b_zn = tl.load(p_zn, boundary_check=(0,))
711
+ # [BC, BV]
712
+ b_v = tl.load(p_v, boundary_check=(0, 1))
713
+ b_dv = tl.zeros([BC, BV], dtype=tl.float32)
714
+ for i_j in range(i_i + 1, NC):
715
+ p_z = tl.make_block_ptr(z + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_j * BC, i_v * BV), (BC, BV), (1, 0))
716
+ p_A = tl.make_block_ptr(A + i_bh * T * BT, (BT, T), (1, BT), (i_i * BC, i_t * BT + i_j * BC), (BC, BC), (0, 1))
717
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_j * BC, i_v * BV), (BC, BV), (1, 0))
718
+ # [BC, BV]
719
+ b_z = tl.load(p_z, boundary_check=(0, 1))
720
+ b_do = tl.load(p_do, boundary_check=(0, 1))
721
+ b_do = (b_do * exp(b_zn[None, :] - b_z)).to(b_do.dtype)
722
+ # [BC, BC]
723
+ b_A = tl.load(p_A, boundary_check=(0, 1))
724
+ b_dv += tl.dot(b_A, b_do, allow_tf32=False)
725
+ b_dv *= exp(b_v - b_zn[None, :])
726
+
727
+ o_i = tl.arange(0, BC)
728
+ for j in range(0, BC):
729
+ p_z = tl.make_block_ptr(z + i_bh * T*V, (T * V,), (1,), ((i_t * BT + i_i * BC + j) * V + i_v * BV,), (BV,), (0,))
730
+ p_A = tl.make_block_ptr(A + i_bh * T * BT, (T * BT,), (1,), ((i_t * BT + i_i * BC + j) * BT + i_i * BC,), (BC,), (0,))
731
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T * V,), (1,), ((i_t * BT + i_i * BC + j) * V + i_v * BV,), (BV,), (0,))
732
+ # [BC,]
733
+ b_A = tl.load(p_A, boundary_check=(0,))
734
+ # [BV,]
735
+ b_z = tl.load(p_z, boundary_check=(0,))
736
+ b_do = tl.load(p_do, boundary_check=(0,))
737
+ # [BC, BV]
738
+ m_i = o_i[:, None] <= j
739
+ b_dv += tl.where(m_i, exp(b_v - b_z[None, :]) * b_A[:, None] * b_do[None, :], 0.)
740
+ p_dv = tl.make_block_ptr(dv + i_bh * T*V, (T, V), (V, 1), (i_t * BT + i_i * BC, i_v * BV), (BC, BV), (1, 0))
741
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
742
+
743
+
744
+ @triton.jit(do_not_specialize=['T'])
745
+ def chunk_abc_bwd_kernel_rcum_inter(
746
+ s,
747
+ z,
748
+ ss,
749
+ doo,
750
+ T,
751
+ S: tl.constexpr,
752
+ BT: tl.constexpr,
753
+ BS: tl.constexpr,
754
+ NT: tl.constexpr,
755
+ ):
756
+ i_m, i_bh = tl.program_id(0), tl.program_id(1)
757
+
758
+ b_sp = tl.zeros([BS], dtype=tl.float32)
759
+ b_zp = tl.full([BS], float('inf'), dtype=tl.float32)
760
+ for i_t in range(NT - 1, -1, -1):
761
+ p_s = tl.make_block_ptr(s + i_bh * T*S, (T, S), (S, 1), (i_t * BT, i_m * BS), (BT, BS), (1, 0))
762
+ p_z = tl.make_block_ptr(z + i_bh * T*S, (T, S), (S, 1), (i_t * BT, i_m * BS), (BT, BS), (1, 0))
763
+ p_zc = tl.make_block_ptr(z + i_bh * T*S, (T*S,), (1,), ((i_t * BT) * S + i_m * BS,), (BS,), (0,))
764
+ p_ss = tl.make_block_ptr(ss + i_bh * T*S, (T, S), (S, 1), (i_t * BT, i_m * BS), (BT, BS), (1, 0))
765
+ p_doo = tl.make_block_ptr(doo + i_bh * T*S, (T, S), (S, 1), (i_t * BT, i_m * BS), (BT, BS), (1, 0))
766
+ # [BS,]
767
+ b_zc = tl.load(p_zc, boundary_check=(0,))
768
+ # [BT, BS]
769
+ b_s = tl.load(p_s, boundary_check=(0, 1))
770
+ b_z = tl.load(p_z, boundary_check=(0, 1))
771
+ b_ss = tl.load(p_ss, boundary_check=(0, 1))
772
+
773
+ b_doo = exp(b_s - b_zp[None, :]) * b_sp[None, :]
774
+ tl.store(p_doo, b_doo.to(p_doo.dtype.element_ty), boundary_check=(0, 1))
775
+ # [BS,]
776
+ b_sp = b_sp * exp(b_zc - b_zp) + tl.sum(b_ss * exp(b_zc[None, :] - b_z), 0)
777
+ b_zp = b_zc
778
+
779
+
780
+ @triton.jit(do_not_specialize=['T'])
781
+ def chunk_abc_bwd_kernel_rcum_intra(
782
+ s,
783
+ z,
784
+ ss,
785
+ doo,
786
+ T,
787
+ S: tl.constexpr,
788
+ BT: tl.constexpr,
789
+ BC: tl.constexpr,
790
+ BS: tl.constexpr,
791
+ NC: tl.constexpr,
792
+ ):
793
+ i_s, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
794
+ i_t, i_i = i_c // NC, i_c % NC
795
+
796
+ o_i = tl.arange(0, BC)
797
+ m_o = tl.full([BC, BC], 1., dtype=tl.float32)
798
+
799
+ p_s = tl.make_block_ptr(s + i_bh * T*S, (T, S), (S, 1), (i_t * BT + i_i * BC, i_s * BS), (BC, BS), (1, 0))
800
+ p_zn = tl.make_block_ptr(z + i_bh * T*S, (T*S,), (1,), ((i_t * BT + i_i * BC + BC - 1) * S + i_s * BS,), (BS,), (0,))
801
+ p_doo = tl.make_block_ptr(doo + i_bh * T*S, (T, S), (S, 1), (i_t * BT + i_i * BC, i_s * BS), (BC, BS), (1, 0))
802
+ # [BC, BS]
803
+ b_s = tl.load(p_s, boundary_check=(0, 1))
804
+ # [BS,]
805
+ b_zn = tl.load(p_zn, boundary_check=(0,))
806
+
807
+ b_doo = tl.zeros([BC, BS], dtype=tl.float32)
808
+ for i_j in range(i_i + 1, NC):
809
+ p_z = tl.make_block_ptr(z + i_bh * T*S, (T, S), (S, 1), (i_t * BT + i_j * BC, i_s * BS), (BC, BS), (1, 0))
810
+ p_ss = tl.make_block_ptr(ss + i_bh * T*S, (T, S), (S, 1), (i_t * BT + i_j * BC, i_s * BS), (BC, BS), (1, 0))
811
+ # [BC, BS]
812
+ b_z = tl.load(p_z, boundary_check=(0, 1))
813
+ b_ss = tl.load(p_ss, boundary_check=(0, 1))
814
+ # [BC, BS]
815
+ b_doo += b_ss * exp(b_zn[None, :] - b_z)
816
+ b_doo = exp(b_s - b_zn[None, :]) * tl.dot(m_o.to(b_s.dtype), b_doo.to(b_s.dtype), allow_tf32=False)
817
+
818
+ for j in range(0, BC):
819
+ p_z = tl.make_block_ptr(z + i_bh * T*S, (T*S,), (1,), ((i_t * BT + i_i * BC + j) * S + i_s * BS,), (BS,), (0,))
820
+ p_ss = tl.make_block_ptr(ss + i_bh * T*S, (T*S,), (1,), ((i_t * BT + i_i * BC + j) * S + i_s * BS,), (BS,), (0,))
821
+ # [BS,]
822
+ b_z = tl.load(p_z, boundary_check=(0,))
823
+ b_ss = tl.load(p_ss, boundary_check=(0,))
824
+ # [BC, BS]
825
+ m_i = o_i[:, None] <= j
826
+ b_doo += tl.where(m_i, exp(b_s - b_z[None, :]) * b_ss[None, :], 0.)
827
+ b_doo += tl.load(p_doo, boundary_check=(0, 1))
828
+ tl.store(p_doo, b_doo.to(p_doo.dtype.element_ty), boundary_check=(0, 1))
829
+
830
+
831
+ class ChunkABCFunction(torch.autograd.Function):
832
+
833
+ @staticmethod
834
+ @input_guard
835
+ def forward(ctx, q, k, v, s, initial_state, output_final_state):
836
+ B, H, T, K, V, M = *q.shape, v.shape[-1], s.shape[-1]
837
+ BT, BC = 64, 16
838
+ BK = min(64, triton.next_power_of_2(K))
839
+ BV = min(64, triton.next_power_of_2(V))
840
+ BM = min(64, triton.next_power_of_2(M))
841
+ NT, NC = triton.cdiv(T, BT), triton.cdiv(BT, BC)
842
+ NV, NM = triton.cdiv(V, BV), triton.cdiv(M, BM)
843
+ num_warps = 4 if BK == 64 else 2
844
+ num_stages = 1
845
+
846
+ def fwd_pre(s, B, H, T, S):
847
+ # keep cummulative normalizer in fp32
848
+ z = torch.empty_like(s, dtype=torch.float)
849
+ grid = (B * H,)
850
+ logcumsumexp_fwd_kernel[grid](
851
+ s, z,
852
+ T=T, S=S,
853
+ )
854
+ return z
855
+
856
+ def fwd_inner(q, k, v, z, B, H, T, K, V, BT, BK, BV, NT, normk=False, h0=None, ht=None):
857
+ NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
858
+ h = q.new_empty(B, H, NT * K, V)
859
+ grid = (NV, NK, B * H)
860
+ chunk_abc_fwd_kernel_h[grid](
861
+ k, v, z, h, h0, ht,
862
+ T=T, K=K, V=V, BT=BT, BK=BK, BV=BV, NT=NT,
863
+ NORMK=normk,
864
+ USE_INITIAL_STATE=h0 is not None,
865
+ STORE_FINAL_STATE=ht is not None,
866
+ num_warps=num_warps,
867
+ num_stages=num_stages,
868
+ )
869
+ return h
870
+
871
+ final_state = None
872
+ if output_final_state:
873
+ final_state = (q.new_empty(B, H, K, M, dtype=torch.float),
874
+ q.new_empty(B, H, M, V, dtype=torch.float))
875
+
876
+ z = fwd_pre(s, B, H, T, M)
877
+ scale = K ** -0.5
878
+ hk = fwd_inner(
879
+ q=q, k=k, v=s, z=z,
880
+ B=B, H=H, T=T, K=K, V=M, BT=BT, BK=BK, BV=BM, NT=NT,
881
+ normk=False,
882
+ h0=initial_state[0] if initial_state is not None else None,
883
+ ht=final_state[0] if final_state is not None else None,
884
+ )
885
+ ok1 = torch.empty_like(s)
886
+ Ak = q.new_empty(B, H, T, BT)
887
+ grid = (NM, NT, B * H)
888
+ chunk_abc_fwd_kernel_K[grid](
889
+ q, k, z, hk, ok1, Ak,
890
+ scale=scale,
891
+ T=T, K=K, V=M, BT=BT, BK=BK, BV=BM, NT=NT,
892
+ num_warps=num_warps,
893
+ num_stages=num_stages,
894
+ )
895
+ ok0 = torch.empty_like(s)
896
+ grid = (NM, NT * NC, B * H)
897
+ chunk_abc_fwd_kernel_intra_K[grid](
898
+ s, z, ok0, Ak,
899
+ T=T, V=M, BT=BT, BC=BC, BV=BM, NC=NC,
900
+ num_warps=2,
901
+ num_stages=num_stages,
902
+ )
903
+ ok = ok0.add_(ok1)
904
+
905
+ scale = 1.
906
+ # p is kept in fp32 for safe softmax backward
907
+ p = softmax_fwd(ok, dtype=torch.float)
908
+ qv = p.to(q.dtype)
909
+
910
+ scale = 1.
911
+ hv = fwd_inner(
912
+ q=qv, k=s, v=v, z=z,
913
+ B=B, H=H, T=T, K=M, V=V, BT=BT, BK=BM, BV=BV, NT=NT,
914
+ normk=True,
915
+ h0=initial_state[1] if initial_state is not None else None,
916
+ ht=final_state[1] if final_state is not None else None,
917
+ )
918
+ Av = q.new_zeros(NM, B, H, T, BT)
919
+ grid = (NM, NT * NC * NC, B * H)
920
+ chunk_abc_fwd_kernel_intra_V[grid](
921
+ qv, s, z, Av,
922
+ scale=scale,
923
+ T=T, K=M, BT=BT, BC=BC, BK=BM, NC=NC,
924
+ num_warps=2,
925
+ num_stages=num_stages,
926
+ )
927
+ Av = Av.sum(0)
928
+ ov = torch.empty_like(v)
929
+ grid = (NV, NT, B * H)
930
+ chunk_abc_fwd_kernel_V[grid](
931
+ qv, v, z, hv, ov, Av,
932
+ scale=scale,
933
+ T=T,
934
+ K=M,
935
+ V=V,
936
+ BT=BT,
937
+ BK=BM,
938
+ BV=BV,
939
+ NT=NT,
940
+ num_warps=num_warps,
941
+ num_stages=num_stages,
942
+ )
943
+ ctx.save_for_backward(q, k, v, s, z, ok, p, hk, hv, Av)
944
+ ctx.BT = BT
945
+ return ov, final_state
946
+
947
+ @staticmethod
948
+ @input_guard
949
+ def backward(ctx, dov, dht=None):
950
+ q, k, v, s, z, ok, p, hk, hv, Av = ctx.saved_tensors
951
+ B, H, T, K, V, M = *q.shape, v.shape[-1], s.shape[-1]
952
+ BT, BC = ctx.BT, 16
953
+ BK = min(64, triton.next_power_of_2(K))
954
+ BV = min(64, triton.next_power_of_2(V))
955
+ BM = min(64, triton.next_power_of_2(M))
956
+ NT, NC = triton.cdiv(T, BT), triton.cdiv(BT, BC)
957
+ NK, NM = triton.cdiv(K, BK), triton.cdiv(M, BM)
958
+ num_warps = 4 if BK == 64 else 2
959
+ num_stages = 1
960
+
961
+ def bwd_inner(q, z, do, B, H, T, K, V, BT, BK, BV, NT, scale, normk=False):
962
+ NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
963
+ dh = q.new_empty(B, H, NT * K, V)
964
+ grid = (NK, NV, B * H)
965
+ chunk_abc_bwd_kernel_dh[grid](
966
+ q, z, do, dh,
967
+ scale=scale,
968
+ T=T, K=K, V=V, BT=BT, BK=BK, BV=BV, NT=NT,
969
+ NORMK=normk,
970
+ num_warps=num_warps,
971
+ num_stages=num_stages,
972
+ )
973
+ return dh
974
+
975
+ def bwd_post(s, z, ss, B, H, T, S, BT, BC, BS, NT, NC, NS):
976
+ doo = torch.empty_like(s)
977
+ grid = (NS, B * H)
978
+ chunk_abc_bwd_kernel_rcum_inter[grid](
979
+ s, z, ss, doo,
980
+ T=T, S=S, BT=BT, BS=BS, NT=NT,
981
+ num_warps=num_warps,
982
+ num_stages=num_stages,
983
+ )
984
+ grid = (NS, NT * NC, B * H)
985
+ chunk_abc_bwd_kernel_rcum_intra[grid](
986
+ s, z, ss, doo,
987
+ T=T, S=S, BT=BT, BC=BC, BS=BS, NC=NC,
988
+ num_warps=num_warps,
989
+ num_stages=num_stages,
990
+ )
991
+ return doo
992
+
993
+ scale = 1.
994
+ qv = p.to(q.dtype)
995
+ dhv = bwd_inner(
996
+ qv, z, dov,
997
+ B=B, H=H, T=T, K=M, V=V, BT=BT, BK=BM, BV=BV, NT=NT,
998
+ scale=scale,
999
+ normk=True,
1000
+ )
1001
+ dp1 = torch.empty_like(p)
1002
+ dsv1 = torch.empty_like(s, dtype=torch.float)
1003
+ dv = v.new_empty(NM, *v.shape)
1004
+ dAv = q.new_zeros(B, H, T, BT)
1005
+ grid = (NM, NT, B * H)
1006
+ chunk_abc_bwd_kernel_V[grid](
1007
+ s, v, z, hv, Av, dov, dhv, dp1, dsv1, dv, dAv,
1008
+ scale=scale,
1009
+ T=T, K=M, V=V, BT=BT, BK=BM, BV=BV, NT=NT,
1010
+ num_warps=num_warps,
1011
+ num_stages=num_stages,
1012
+ )
1013
+ dv = dv.sum(0)
1014
+ dp0 = torch.empty_like(p)
1015
+ dsv0 = s.new_zeros(s.shape, dtype=torch.float)
1016
+ grid = (NM, NT * NC, B * H)
1017
+ chunk_abc_bwd_kernel_intra_V[grid](
1018
+ qv, s, z, dAv, dp0, dsv0,
1019
+ T=T, K=M, BT=BT, BC=BC, BK=BM, NC=NC,
1020
+ num_warps=2,
1021
+ num_stages=num_stages,
1022
+ )
1023
+ dp = dp1.add_(dp0)
1024
+ dsv = dsv1.add_(dsv0)
1025
+
1026
+ # softmax gradient, equivalent to:
1027
+ # dok = p * (dp - (p * dp).sum(-1, True))
1028
+ dok = softmax_bwd(p, dp, dtype=ok.dtype)
1029
+
1030
+ scale = K ** -0.5
1031
+ dhk = bwd_inner(
1032
+ q, z, dok,
1033
+ B=B, H=H, T=T, K=K, V=M, BT=BT, BK=BK, BV=BM, NT=NT,
1034
+ scale=scale,
1035
+ normk=False,
1036
+ )
1037
+ dAk = q.new_zeros(NM, B, H, T, BT)
1038
+ grid = (NM, NT * NC * NC, B * H)
1039
+ chunk_abc_bwd_kernel_intra_K[grid](
1040
+ s, z, dok, dAk,
1041
+ scale=scale,
1042
+ T=T, V=M, BT=BT, BC=BC, BV=BM, NC=NC,
1043
+ num_warps=2,
1044
+ num_stages=num_stages,
1045
+ )
1046
+ dAk = dAk.sum(0)
1047
+
1048
+ Ak = q.new_zeros(NK, B, H, T, BT)
1049
+ dq = torch.empty_like(q)
1050
+ dk = torch.empty_like(k)
1051
+ dsk1 = s.new_empty(NK, *s.shape, dtype=torch.float)
1052
+ grid = (NK, NT, B * H)
1053
+ chunk_abc_bwd_kernel_K[grid](
1054
+ q, k, s, z, hk, Ak, dok, dhk, dq, dk, dsk1, dAk,
1055
+ scale=scale,
1056
+ T=T, K=K, V=M, BT=BT, BK=BK, BV=BM, NT=NT,
1057
+ num_warps=num_warps,
1058
+ num_stages=num_stages,
1059
+ )
1060
+ Ak = Ak.sum(0)
1061
+ dsk1 = dsk1.sum(0)
1062
+ dsk0 = torch.empty_like(s, dtype=torch.float)
1063
+ grid = (NM, NT * NC, B * H)
1064
+ chunk_abc_bwd_kernel_intra_KV[grid](
1065
+ s, z, Ak, dok, dsk0,
1066
+ T=T, V=M, BT=BT, BC=BC, BV=BM, NC=NC,
1067
+ num_warps=2,
1068
+ num_stages=num_stages,
1069
+ )
1070
+ ds = dsv.add_(dsk1.add_(dsk0))
1071
+ ds -= bwd_post(s, z, ok * dok + p * dp, B, H, T, M, BT, BC, BM, NT, NC, NM)
1072
+ ds = ds.to(s.dtype)
1073
+ return dq, dk, dv, ds, None, None
1074
+
1075
+
1076
+ @torch.compiler.disable
1077
+ def chunk_abc(
1078
+ q: torch.Tensor,
1079
+ k: torch.Tensor,
1080
+ v: torch.Tensor,
1081
+ s: torch.Tensor,
1082
+ initial_state: tuple[torch.Tensor] | None = None,
1083
+ output_final_state: bool = False,
1084
+ head_first: bool = False,
1085
+ ) -> tuple[torch.Tensor, torch.Tensor]:
1086
+ r"""
1087
+ Args:
1088
+ q (torch.Tensor):
1089
+ queries of shape `[B, T, H, K]`.
1090
+ k (torch.Tensor):
1091
+ keys of shape `[B, T, H, K]`.
1092
+ v (torch.Tensor):
1093
+ values of shape `[B, T, H, V]`.
1094
+ s (torch.Tensor):
1095
+ slot representations of shape `[B, T, H, M]`.
1096
+ initial_state (Optional[Tuple[torch.Tensor, torch.Tensor]]):
1097
+ Initial states of shape `[B, H, K, M]` and `[B, H, M, V]`. Default: `None`.
1098
+ output_final_state (Optional[bool]):
1099
+ Whether to output the final state of shape `[B, H, K, M]` and `[B, H, M, V]`. Default: `False`.
1100
+ head_first (Optional[bool]):
1101
+ Whether the inputs are in the head-first format. Default: `False`.
1102
+ This argument has been deprecated.
1103
+
1104
+ Returns:
1105
+ o (torch.Tensor):
1106
+ Outputs of shape `[B, T, H, V]`.
1107
+ final_state (torch.Tensor):
1108
+ Final state of shape `[B, H, K, M]` and `[B, H, M, V]` if `output_final_state=True` else `None`.
1109
+ """
1110
+ if not head_first:
1111
+ q, k, v, s = map(lambda x: x.transpose(1, 2), (q, k, v, s))
1112
+ o, final_state = ChunkABCFunction.apply(q, k, v, s, initial_state, output_final_state)
1113
+ if not head_first:
1114
+ o = o.transpose(1, 2)
1115
+ return o, final_state
code/flash-linear-attention/fla/ops/abc/naive.py ADDED
@@ -0,0 +1,94 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+
3
+ import torch
4
+ from einops import repeat
5
+
6
+
7
+ def naive_recurrent_abc(
8
+ q: torch.Tensor,
9
+ k: torch.Tensor,
10
+ v: torch.Tensor,
11
+ s: torch.Tensor,
12
+ g: torch.Tensor | None = None,
13
+ scale: int | None = None,
14
+ initial_state: torch.Tensor | None = None,
15
+ output_final_state: bool | None = False,
16
+ ) -> torch.Tensor:
17
+ dtype = q.dtype
18
+
19
+ NG = q.shape[1]//k.shape[1]
20
+ # [batch_size, n_heads, seq_len, n_slots]
21
+ if g is None:
22
+ z = s.float().logcumsumexp(2)
23
+ g = torch.cat((z[:, :, :1], z[:, :, :-1]), 2) - z
24
+ s = torch.exp(s - z)
25
+ q, k, v, s, g = map(lambda x: x.float(), (q, k, v, s, g))
26
+ k, v, s, g = map(lambda x: repeat(x, 'b h t d -> b (h g) t d', g=NG), (k, v, s, g))
27
+ if initial_state is not None:
28
+ initial_state = tuple(map(lambda x: repeat(x, 'b h k v -> b (h g) k v', g=NG), initial_state))
29
+
30
+ B, H, T, K, V, M = *q.shape, v.shape[-1], s.shape[-1]
31
+
32
+ hk = torch.zeros(B, H, K, M, dtype=torch.float, device=q.device)
33
+ ok = torch.zeros_like(s)
34
+
35
+ if scale is None:
36
+ scale = q.shape[-1] ** -0.5
37
+
38
+ final_state = None
39
+ if initial_state is not None:
40
+ hk += initial_state[0]
41
+
42
+ for i in range(T):
43
+ q_i = q[:, :, i] * scale
44
+ k_i = k[:, :, i]
45
+ v_i = s[:, :, i]
46
+ g_i = g[:, :, i].exp()
47
+ hk = hk * g_i[..., None, :] + k_i[..., None] * v_i[..., None, :]
48
+ ok[:, :, i] = (q_i[..., None] * hk).sum(-2)
49
+
50
+ qv = ok.softmax(-1)
51
+ hv = torch.zeros(B, H, M, V, dtype=torch.float, device=q.device)
52
+ ov = torch.zeros_like(v)
53
+ if initial_state is not None:
54
+ hv += initial_state[1]
55
+
56
+ for i in range(T):
57
+ q_i = qv[:, :, i]
58
+ k_i = s[:, :, i]
59
+ v_i = v[:, :, i]
60
+ g_i = g[:, :, i].exp()
61
+ hv = hv * g_i[..., :, None] + k_i[..., None] * v_i[..., None, :]
62
+ ov[:, :, i] = (q_i[..., None] * hv).sum(-2)
63
+
64
+ if output_final_state:
65
+ final_state = (hk.view(B, -1, NG, K, M)[:, :, 0], hv.view(B, -1, NG, M, V)[:, :, 0])
66
+ return ov.to(dtype), final_state
67
+
68
+
69
+ def naive_cumsum_abc(
70
+ q: torch.Tensor,
71
+ k: torch.Tensor,
72
+ v: torch.Tensor,
73
+ s: torch.Tensor,
74
+ ) -> torch.Tensor:
75
+ """
76
+ A simple implementation of vanilla ABC that is more aligned with the descriptions in the paper.
77
+ This is just for demonstration purposes, with no numerical stabilities guaranteed.
78
+ """
79
+
80
+ dtype = q.dtype
81
+ q, k, v, s = map(lambda x: x.float(), (q, k, v, s))
82
+
83
+ scale = q.shape[-1] ** -0.5
84
+ # [batch_size, n_heads, seq_len, n_slots]
85
+ s = (s - s.max(2, True)[0]).exp()
86
+ z = s.cumsum(2)
87
+ # [batch_size, n_heads, seq_len, n_slots, d_head]
88
+ K = (s.unsqueeze(-1) * k.unsqueeze(-2)).cumsum(2) / z.unsqueeze(-1)
89
+ V = (s.unsqueeze(-1) * v.unsqueeze(-2)).cumsum(2) / z.unsqueeze(-1)
90
+ # [batch_size, n_heads, seq_len, n_slots]
91
+ p = torch.einsum('...d,...md->...m', q * scale, K).softmax(-1)
92
+ # [batch_size, n_heads, seq_len, d_head]
93
+ o = torch.einsum('...m,...md->...d', p, V)
94
+ return o.to(dtype), None
code/flash-linear-attention/fla/ops/attn/__init__.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+
2
+ from .parallel import parallel_attn
3
+
4
+ __all__ = [
5
+ 'parallel_attn',
6
+ ]
code/flash-linear-attention/fla/ops/attn/decoding.py ADDED
@@ -0,0 +1,181 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils.cumsum import chunk_global_cumsum
9
+ from fla.ops.utils.op import exp
10
+ from fla.utils import autotune_cache_kwargs, check_shared_mem
11
+
12
+
13
+ @triton.heuristics({
14
+ 'USE_G': lambda args: args['g_cumsum'] is not None,
15
+ })
16
+ @triton.autotune(
17
+ configs=[
18
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
19
+ for num_warps in [1, 2, 4] + ([] if check_shared_mem('hopper') else [8])
20
+ for num_stages in [2, 3, 4, 5]
21
+ ],
22
+ key=['H', 'G', 'K', 'V', 'BK', 'BV', 'USE_G'],
23
+ **autotune_cache_kwargs,
24
+ )
25
+ @triton.jit
26
+ def naive_attn_decoding_kernel(
27
+ q,
28
+ k,
29
+ v,
30
+ o,
31
+ g_cumsum,
32
+ scale,
33
+ gate_scale,
34
+ cu_seqlens,
35
+ T,
36
+ B: tl.constexpr,
37
+ H: tl.constexpr,
38
+ HQ: tl.constexpr,
39
+ G: tl.constexpr,
40
+ K: tl.constexpr,
41
+ V: tl.constexpr,
42
+ BS: tl.constexpr,
43
+ BK: tl.constexpr,
44
+ BV: tl.constexpr,
45
+ USE_G: tl.constexpr,
46
+ ):
47
+ i_v, i_bh = tl.program_id(0), tl.program_id(1)
48
+ i_b, i_hq = i_bh // HQ, i_bh % HQ
49
+ i_h = i_hq // G
50
+
51
+ bos, eos = tl.load(cu_seqlens + i_b).to(tl.int32), tl.load(cu_seqlens + i_b + 1).to(tl.int32)
52
+ T = eos - bos
53
+
54
+ p_q = tl.make_block_ptr(q + i_bh * K, (K,), (1, ), (0, ), (BK,), (0,))
55
+ p_o = tl.make_block_ptr(o + i_bh * V, (V,), (1, ), (0, ), (BV,), (0,))
56
+
57
+ b_q = tl.load(p_q, boundary_check=(0,))
58
+ b_q = (b_q * scale).to(b_q.dtype)
59
+
60
+ b_o = tl.zeros([BV ], dtype=tl.float32)
61
+
62
+ b_m = tl.full([1], float('-inf'), dtype=tl.float32)
63
+ b_acc = tl.zeros([1], dtype=tl.float32)
64
+
65
+ if USE_G:
66
+ p_g = tl.make_block_ptr(g_cumsum + bos * HQ + i_hq, (T,), (HQ,), (T-1,), (1,), (0,))
67
+ b_gq = tl.load(p_g, boundary_check=(0,)).to(tl.float32)
68
+ else:
69
+ b_gq = None
70
+
71
+ for i_s in range(0, T, BS):
72
+ p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (T, K), (H*K, 1), (i_s, 0), (BS, BK), (1, 0))
73
+ p_v = tl.make_block_ptr(v + (bos * H + i_h) * V, (T, V), (H*V, 1), (i_s, i_v * BV), (BS, BV), (1, 0))
74
+ # [BK, BS]
75
+ b_k = tl.load(p_k, boundary_check=(0, 1))
76
+ # [BS, BV]
77
+ b_v = tl.load(p_v, boundary_check=(0, 1))
78
+ # [BT, BS]
79
+ b_s = tl.sum(b_q[None, :] * b_k, 1)
80
+
81
+ mask = i_s + tl.arange(0, BS) < T
82
+ b_s = tl.where(mask, b_s, float('-inf'))
83
+
84
+ if USE_G:
85
+ p_gk = tl.make_block_ptr(g_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_s,), (BS,), (0,))
86
+ b_gk = tl.load(p_gk, boundary_check=(0,)).to(tl.float32)
87
+ b_s += (b_gq - b_gk) * gate_scale
88
+ # [BT, BS]
89
+ b_m, b_mp = tl.maximum(b_m, tl.max(b_s)), b_m
90
+ b_r = exp(b_mp - b_m)
91
+ # [BT, BS]
92
+ b_p = exp(b_s - b_m)
93
+
94
+ # [BT]
95
+ b_acc = b_acc * b_r + tl.sum(b_p, 0)
96
+ # [BT, BV]
97
+ b_o = b_o * b_r + tl.sum(b_p[:, None] * b_v, 0)
98
+ b_mp = b_m
99
+ b_o = b_o / b_acc
100
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, ))
101
+
102
+
103
+ def attn_decoding_one_step(
104
+ q: torch.Tensor,
105
+ k: torch.Tensor,
106
+ v: torch.Tensor,
107
+ g: torch.Tensor | None = None,
108
+ scale: float | None = None,
109
+ cu_seqlens: torch.LongTensor = None,
110
+ do_gate_scale: bool = False,
111
+ ):
112
+ r"""
113
+ Args:
114
+ q (torch.Tensor):
115
+ query of shape `[1, B, HQ, K]`.
116
+ k (torch.Tensor):
117
+ keys of shape `[1, T, H, K]`.
118
+ GQA will be applied if HQ is divisible by H. T is the cumulative length for all batch.
119
+ v (torch.Tensor):
120
+ values of shape `[1, T, H, V]`.
121
+ g (Optional[torch.Tensor]):
122
+ log decay factors of shape `[1, T, H]`. Default: `None`.
123
+ scale (Optional[float]):
124
+ Scale factor for attention scores.
125
+ If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
126
+ cu_seqlens (torch.LongTensor):
127
+ Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
128
+ consistent with the FlashAttention API.
129
+ do_gate_scale (bool):
130
+ Whether to apply gate scale. Default: `False`. If `True`, the attention scale will also be applied
131
+ to the gating bias term in Forgetting Transformer or PaTH-FoX.
132
+
133
+ Returns:
134
+ o (torch.Tensor):
135
+ Outputs of shape `[B, 1, HQ, V]`.
136
+ """
137
+ assert cu_seqlens is not None, "The cu_seqlens must be provided for varlen decoding"
138
+ B, T, H, K, V = *k.shape, v.shape[-1]
139
+ N = len(cu_seqlens) - 1
140
+ HQ = q.shape[2]
141
+ G = HQ // H
142
+ if scale is None:
143
+ scale = K ** -0.5
144
+
145
+ BK = max(triton.next_power_of_2(K), 16)
146
+ if check_shared_mem('hopper', q.device.index):
147
+ BS = min(64, max(16, triton.next_power_of_2(T)))
148
+ BV = min(256, max(16, triton.next_power_of_2(V)))
149
+ elif check_shared_mem('ampere', q.device.index):
150
+ BS = min(32, max(16, triton.next_power_of_2(T)))
151
+ BV = min(128, max(16, triton.next_power_of_2(V)))
152
+ else:
153
+ BS = min(32, max(16, triton.next_power_of_2(T)))
154
+ BV = min(64, max(16, triton.next_power_of_2(V)))
155
+ g_cumsum = chunk_global_cumsum(g, cu_seqlens=cu_seqlens, output_dtype=torch.float32) if g is not None else None
156
+ NV = triton.cdiv(V, BV)
157
+ o = torch.empty(*q.shape[:-1], V, dtype=v.dtype, device=q.device)
158
+ gate_scale = 1.0 if not do_gate_scale else scale
159
+
160
+ grid = (NV, N * HQ)
161
+ naive_attn_decoding_kernel[grid](
162
+ q=q,
163
+ k=k,
164
+ v=v,
165
+ o=o,
166
+ g_cumsum=g_cumsum,
167
+ scale=scale,
168
+ gate_scale=gate_scale,
169
+ cu_seqlens=cu_seqlens,
170
+ B=B,
171
+ T=T,
172
+ H=H,
173
+ HQ=HQ,
174
+ G=G,
175
+ K=K,
176
+ V=V,
177
+ BS=BS,
178
+ BK=BK,
179
+ BV=BV,
180
+ )
181
+ return o
code/flash-linear-attention/fla/ops/attn/parallel.py ADDED
@@ -0,0 +1,728 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+ import warnings
4
+
5
+ import torch
6
+ import triton
7
+ import triton.language as tl
8
+ from einops import reduce
9
+
10
+ from fla.ops.utils import prepare_chunk_indices
11
+ from fla.ops.utils.cumsum import chunk_global_cumsum
12
+ from fla.ops.utils.op import exp2, log2
13
+ from fla.utils import autocast_custom_bwd, autocast_custom_fwd, check_shared_mem, contiguous
14
+
15
+
16
+ @triton.heuristics({
17
+ 'USE_G': lambda args: args['g_cumsum'] is not None,
18
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
19
+ })
20
+ @triton.jit
21
+ def parallel_attn_fwd_kernel(
22
+ q,
23
+ k,
24
+ v,
25
+ o,
26
+ g_cumsum,
27
+ lse,
28
+ scale,
29
+ cu_seqlens,
30
+ chunk_indices,
31
+ T,
32
+ B: tl.constexpr,
33
+ H: tl.constexpr,
34
+ HQ: tl.constexpr,
35
+ G: tl.constexpr,
36
+ K: tl.constexpr,
37
+ V: tl.constexpr,
38
+ BT: tl.constexpr,
39
+ BS: tl.constexpr,
40
+ BK: tl.constexpr,
41
+ BV: tl.constexpr,
42
+ USE_G: tl.constexpr,
43
+ IS_VARLEN: tl.constexpr,
44
+ ):
45
+ i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
46
+ i_b, i_hq = i_bh // HQ, i_bh % HQ
47
+ i_h = i_hq // G
48
+
49
+ if IS_VARLEN:
50
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
51
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
52
+ T = eos - bos
53
+ else:
54
+ i_n = i_b
55
+ bos, eos = i_n * T, i_n * T + T
56
+ RCP_LN2: tl.constexpr = 1.4426950216
57
+
58
+ p_q = tl.make_block_ptr(q + (bos * HQ + i_hq) * K, (T, K), (HQ*K, 1), (i_t * BT, 0), (BT, BK), (1, 0))
59
+ p_o = tl.make_block_ptr(o + (bos * HQ + i_hq) * V, (T, V), (HQ*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
60
+ p_lse = tl.make_block_ptr(lse + bos * HQ + i_hq, (T,), (HQ,), (i_t * BT,), (BT,), (0,))
61
+
62
+ # the Q block is kept in the shared memory throughout the whole kernel
63
+ # [BT, BK]
64
+ b_q = tl.load(p_q, boundary_check=(0, 1))
65
+ # [BT, BV]
66
+ b_o = tl.zeros([BT, BV], dtype=tl.float32)
67
+
68
+ b_m = tl.full([BT], float('-inf'), dtype=tl.float32)
69
+ b_acc = tl.zeros([BT], dtype=tl.float32)
70
+
71
+ if USE_G:
72
+ p_g = tl.make_block_ptr(g_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_t * BT,), (BT,), (0,))
73
+ b_gq = tl.load(p_g, boundary_check=(0,)).to(tl.float32)
74
+ else:
75
+ b_gq = None
76
+
77
+ for i_s in range(0, i_t * BT, BS):
78
+ p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (K, T), (1, H*K), (0, i_s), (BK, BS), (0, 1))
79
+ p_v = tl.make_block_ptr(v + (bos * H + i_h) * V, (T, V), (H*V, 1), (i_s, i_v * BV), (BS, BV), (1, 0))
80
+ # [BK, BS]
81
+ b_k = tl.load(p_k, boundary_check=(0, 1))
82
+ # [BS, BV]
83
+ b_v = tl.load(p_v, boundary_check=(0, 1))
84
+ # [BT, BS]
85
+ b_s = tl.dot(b_q, b_k) * scale * RCP_LN2
86
+
87
+ if USE_G:
88
+ o_k = i_s + tl.arange(0, BS)
89
+ m_k = o_k < T
90
+ b_gk = tl.load(g_cumsum + (bos + o_k) * HQ + i_hq, mask=m_k, other=0).to(tl.float32)
91
+ b_s += b_gq[:, None] - b_gk[None, :]
92
+
93
+ # [BT, BS]
94
+ b_m, b_mp = tl.maximum(b_m, tl.max(b_s, 1)), b_m
95
+ b_r = exp2(b_mp - b_m)
96
+ # [BT, BS]
97
+ b_p = exp2(b_s - b_m[:, None])
98
+ # [BT]
99
+ b_acc = b_acc * b_r + tl.sum(b_p, 1)
100
+ # [BT, BV]
101
+ b_o = b_o * b_r[:, None] + tl.dot(b_p.to(b_q.dtype), b_v)
102
+
103
+ b_mp = b_m
104
+
105
+ # [BT]
106
+ o_q = i_t * BT + tl.arange(0, BT)
107
+ for i_s in range(i_t * BT, min((i_t + 1) * BT, T), BS):
108
+ p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (K, T), (1, H*K), (0, i_s), (BK, BS), (0, 1))
109
+ p_v = tl.make_block_ptr(v + (bos * H + i_h) * V, (T, V), (H*V, 1), (i_s, i_v * BV), (BS, BV), (1, 0))
110
+
111
+ # [BS]
112
+ o_k = i_s + tl.arange(0, BS)
113
+ m_k = o_k < T
114
+ # [BK, BS]
115
+ b_k = tl.load(p_k, boundary_check=(0, 1))
116
+ # [BS, BV]
117
+ b_v = tl.load(p_v, boundary_check=(0, 1))
118
+ # [BT, BS]
119
+ b_s = tl.dot(b_q, b_k) * scale * RCP_LN2
120
+
121
+ if USE_G:
122
+ b_gk = tl.load(g_cumsum + (bos + o_k) * HQ + i_hq, mask=m_k, other=0).to(tl.float32)
123
+ b_s += b_gq[:, None] - b_gk[None, :]
124
+
125
+ b_s = tl.where((o_q[:, None] >= o_k[None, :]) & m_k[None, :], b_s, float('-inf'))
126
+
127
+ # [BT]
128
+ b_m, b_mp = tl.maximum(b_m, tl.max(b_s, 1)), b_m
129
+ b_r = exp2(b_mp - b_m)
130
+ # [BT, BS]
131
+ b_p = exp2(b_s - b_m[:, None])
132
+ # [BT]
133
+ b_acc = b_acc * b_r + tl.sum(b_p, 1)
134
+ # [BT, BV]
135
+ b_o = b_o * b_r[:, None] + tl.dot(b_p.to(b_q.dtype), b_v)
136
+ b_mp = b_m
137
+
138
+ b_o = b_o / b_acc[:, None]
139
+ b_m += log2(b_acc)
140
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
141
+ tl.store(p_lse, b_m.to(p_lse.dtype.element_ty), boundary_check=(0,))
142
+
143
+
144
+ @triton.jit
145
+ def parallel_attn_bwd_kernel_preprocess(
146
+ o,
147
+ do,
148
+ delta,
149
+ B: tl.constexpr,
150
+ V: tl.constexpr,
151
+ ):
152
+ i_n = tl.program_id(0)
153
+ o_d = tl.arange(0, B)
154
+ m_d = o_d < V
155
+
156
+ b_o = tl.load(o + i_n * V + o_d, mask=m_d, other=0)
157
+ b_do = tl.load(do + i_n * V + o_d, mask=m_d, other=0).to(tl.float32)
158
+ b_delta = tl.sum(b_o * b_do)
159
+
160
+ tl.store(delta + i_n, b_delta.to(delta.dtype.element_ty))
161
+
162
+
163
+ @triton.heuristics({
164
+ 'USE_G': lambda args: args['g_cumsum'] is not None,
165
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
166
+ })
167
+ @triton.jit(do_not_specialize=['T'])
168
+ def parallel_attn_bwd_kernel_dq(
169
+ q,
170
+ k,
171
+ v,
172
+ lse,
173
+ delta,
174
+ do,
175
+ dq,
176
+ dg_cumsum,
177
+ g_cumsum,
178
+ scale,
179
+ cu_seqlens,
180
+ chunk_indices,
181
+ T,
182
+ B: tl.constexpr,
183
+ H: tl.constexpr,
184
+ HQ: tl.constexpr,
185
+ G: tl.constexpr,
186
+ K: tl.constexpr,
187
+ V: tl.constexpr,
188
+ BT: tl.constexpr,
189
+ BS: tl.constexpr,
190
+ BK: tl.constexpr,
191
+ BV: tl.constexpr,
192
+ IS_VARLEN: tl.constexpr,
193
+ USE_G: tl.constexpr,
194
+ ):
195
+ i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
196
+ i_b, i_hq = i_bh // HQ, i_bh % HQ
197
+ i_h = i_hq // G
198
+
199
+ if IS_VARLEN:
200
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
201
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
202
+ T = eos - bos
203
+ else:
204
+ i_n = i_b
205
+ bos, eos = i_n * T, i_n * T + T
206
+ # NOTE: we must multiply RCP_LN2 after tl.dot for high precision
207
+ RCP_LN2: tl.constexpr = 1.4426950216
208
+
209
+ p_q = tl.make_block_ptr(q + (bos * HQ + i_hq) * K, (T, K), (HQ*K, 1), (i_t * BT, 0), (BT, BK), (1, 0))
210
+ p_dq = tl.make_block_ptr(dq + (bos * HQ + i_hq) * K, (T, K), (HQ*K, 1), (i_t * BT, 0), (BT, BK), (1, 0))
211
+ p_do = tl.make_block_ptr(do + (bos * HQ + i_hq) * V, (T, V), (HQ*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
212
+ p_lse = tl.make_block_ptr(lse + bos * HQ + i_hq, (T,), (HQ,), (i_t * BT,), (BT,), (0,))
213
+ p_delta = tl.make_block_ptr(delta + bos * HQ + i_hq, (T,), (HQ,), (i_t * BT,), (BT,), (0,))
214
+
215
+ # [BT, BK]
216
+ b_q = tl.load(p_q, boundary_check=(0, 1))
217
+ # [BT, BV]
218
+ b_do = tl.load(p_do, boundary_check=(0, 1))
219
+ # [BT]
220
+ b_lse = tl.load(p_lse, boundary_check=(0,))
221
+ b_delta = tl.load(p_delta, boundary_check=(0,))
222
+
223
+ # [BT, BK]
224
+ b_dq = tl.zeros([BT, BK], dtype=tl.float32)
225
+ if USE_G:
226
+ b_dg = tl.zeros([BT ], dtype=tl.float32)
227
+ p_gq = tl.make_block_ptr(g_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_t * BT,), (BT,), (0,))
228
+ b_gq = tl.load(p_gq, boundary_check=(0,)).to(tl.float32)
229
+ else:
230
+ b_gq = None
231
+ b_dg = None
232
+
233
+ o_q = i_t * BT + tl.arange(0, BT)
234
+ for i_s in range(0, i_t * BT, BS):
235
+ p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (K, T), (1, H*K), (0, i_s), (BK, BS), (0, 1))
236
+ p_v = tl.make_block_ptr(v + (bos * H + i_h) * V, (V, T), (1, H*V), (i_v * BV, i_s), (BV, BS), (0, 1))
237
+
238
+ o_k = i_s + tl.arange(0, BS)
239
+ m_k = o_k < T
240
+ # [BK, BS]
241
+ b_k = tl.load(p_k, boundary_check=(0, 1))
242
+ # [BV, BS]
243
+ b_v = tl.load(p_v, boundary_check=(0, 1))
244
+ # [BT, BS]
245
+ b_s = tl.dot(b_q, b_k) * scale * RCP_LN2
246
+ if USE_G:
247
+ b_gk = tl.load(g_cumsum + (bos + o_k) * HQ + i_hq, mask=m_k, other=0).to(tl.float32)
248
+ b_s += b_gq[:, None] - b_gk[None, :]
249
+
250
+ b_s = tl.where((o_q[:, None] >= o_k[None, :]) & m_k[None, :], b_s, float('-inf'))
251
+ b_p = exp2(b_s - b_lse[:, None])
252
+ # [BT, BV] @ [BV, BS] -> [BT, BS]
253
+ b_dp = tl.dot(b_do, b_v)
254
+ b_ds = b_p * (b_dp.to(tl.float32) - b_delta[:, None])
255
+ # [BT, BS] @ [BS, BK] -> [BT, BK]
256
+ b_dq += tl.dot(b_ds.to(b_k.dtype), tl.trans(b_k))
257
+ if USE_G:
258
+ b_dg += tl.sum(b_ds, 1)
259
+
260
+ # [BT]
261
+ o_q = i_t * BT + tl.arange(0, BT)
262
+ for i_s in range(i_t * BT, min((i_t + 1) * BT, T), BS):
263
+ p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (K, T), (1, H*K), (0, i_s), (BK, BS), (0, 1))
264
+ p_v = tl.make_block_ptr(v + (bos * H + i_h) * V, (V, T), (1, H*V), (i_v * BV, i_s), (BV, BS), (0, 1))
265
+
266
+ # [BS]
267
+ o_k = i_s + tl.arange(0, BS)
268
+ m_k = o_k < T
269
+ # [BK, BS]
270
+ b_k = tl.load(p_k, boundary_check=(0, 1))
271
+ # [BV, BS]
272
+ b_v = tl.load(p_v, boundary_check=(0, 1))
273
+ # [BT, BS]
274
+ b_s = tl.dot(b_q, b_k) * scale * RCP_LN2
275
+
276
+ if USE_G:
277
+ p_gk = tl.make_block_ptr(g_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_s,), (BS,), (0,))
278
+ b_gk = tl.load(p_gk, boundary_check=(0,)).to(tl.float32)
279
+ b_s += b_gq[:, None] - b_gk[None, :]
280
+ b_p = tl.where((o_q[:, None] >= o_k[None, :]) & m_k[None, :], exp2(b_s - b_lse[:, None]), 0)
281
+
282
+ # [BT, BV] @ [BV, BS] -> [BT, BS]
283
+ b_dp = tl.dot(b_do, b_v)
284
+ b_ds = b_p * (b_dp.to(tl.float32) - b_delta[:, None])
285
+ # [BT, BS] @ [BS, BK] -> [BT, BK]
286
+ b_dq += tl.dot(b_ds.to(b_k.dtype), tl.trans(b_k))
287
+ if USE_G:
288
+ b_dg += tl.sum(b_ds, 1)
289
+
290
+ b_dq *= scale
291
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
292
+ if USE_G:
293
+ p_dg = tl.make_block_ptr(dg_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_t * BT,), (BT,), (0,))
294
+ tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), boundary_check=(0,))
295
+
296
+
297
+ @triton.heuristics({
298
+ 'USE_G': lambda args: args['g_cumsum'] is not None,
299
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
300
+ })
301
+ @triton.jit(do_not_specialize=['T'])
302
+ def parallel_attn_bwd_kernel_dkv(
303
+ q,
304
+ k,
305
+ v,
306
+ g_cumsum,
307
+ lse,
308
+ delta,
309
+ do,
310
+ dk,
311
+ dv,
312
+ dg_cumsum,
313
+ cu_seqlens,
314
+ chunk_indices,
315
+ scale,
316
+ T,
317
+ B: tl.constexpr,
318
+ H: tl.constexpr,
319
+ HQ: tl.constexpr,
320
+ G: tl.constexpr,
321
+ K: tl.constexpr,
322
+ V: tl.constexpr,
323
+ BT: tl.constexpr,
324
+ BS: tl.constexpr,
325
+ BK: tl.constexpr,
326
+ BV: tl.constexpr,
327
+ USE_G: tl.constexpr,
328
+ IS_VARLEN: tl.constexpr,
329
+ ):
330
+ i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
331
+ i_b, i_hq = i_bh // HQ, i_bh % HQ
332
+ i_h = i_hq // G
333
+
334
+ if IS_VARLEN:
335
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
336
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
337
+ T = eos - bos
338
+ else:
339
+ i_n = i_b
340
+ bos, eos = i_n * T, i_n * T + T
341
+ RCP_LN2: tl.constexpr = 1.4426950216
342
+
343
+ p_k = tl.make_block_ptr(k + (bos * H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, 0), (BT, BK), (1, 0))
344
+ p_v = tl.make_block_ptr(v + (bos * H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
345
+ p_dk = tl.make_block_ptr(dk + (bos * HQ + i_hq) * K, (T, K), (HQ*K, 1), (i_t * BT, 0), (BT, BK), (1, 0))
346
+ p_dv = tl.make_block_ptr(dv + (bos * HQ + i_hq) * V, (T, V), (HQ*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
347
+
348
+ # [BT, BK]
349
+ b_k = tl.load(p_k, boundary_check=(0, 1))
350
+ b_dk = tl.zeros([BT, BK], dtype=tl.float32)
351
+ # [BT, BV]
352
+ b_v = tl.load(p_v, boundary_check=(0, 1))
353
+ b_dv = tl.zeros([BT, BV], dtype=tl.float32)
354
+
355
+ o_k = i_t * BT + tl.arange(0, BT)
356
+
357
+ if USE_G:
358
+ p_gk = tl.make_block_ptr(g_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_t * BT,), (BT,), (0,))
359
+ b_gk = tl.load(p_gk, boundary_check=(0,)).to(tl.float32)
360
+ b_dg = tl.zeros([BT], dtype=tl.float32)
361
+ else:
362
+ b_gk = None
363
+ b_dg = None
364
+
365
+ for i_s in range(i_t * BT, min((i_t + 1) * BT, T), BS):
366
+ p_q = tl.make_block_ptr(q + (bos * HQ + i_hq) * K, (T, K), (HQ*K, 1), (i_s, 0), (BS, BK), (1, 0))
367
+ p_do = tl.make_block_ptr(do + (bos * HQ + i_hq) * V, (T, V), (HQ*V, 1), (i_s, i_v * BV), (BS, BV), (1, 0))
368
+ p_lse = tl.make_block_ptr(lse + bos * HQ + i_hq, (T,), (HQ,), (i_s,), (BS,), (0,))
369
+ p_delta = tl.make_block_ptr(delta + bos * HQ + i_hq, (T,), (HQ,), (i_s,), (BS,), (0,))
370
+
371
+ # [BS]
372
+ o_q = i_s + tl.arange(0, BS)
373
+ m_q = o_q < T
374
+ # [BS, BK]
375
+ b_q = tl.load(p_q, boundary_check=(0, 1))
376
+ # [BS, BV]
377
+ b_do = tl.load(p_do, boundary_check=(0, 1))
378
+ # [BS]
379
+ b_lse = tl.load(p_lse, boundary_check=(0,))
380
+ b_delta = tl.load(p_delta, boundary_check=(0,))
381
+ # [BT, BS]
382
+ b_s = tl.dot(b_k, tl.trans(b_q)) * scale * RCP_LN2
383
+ if USE_G:
384
+ p_gq = tl.make_block_ptr(g_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_s,), (BS,), (0,))
385
+ b_gq = tl.load(p_gq, boundary_check=(0,)).to(tl.float32)
386
+ b_s += b_gq[None, :] - b_gk[:, None]
387
+ b_p = tl.where((o_k[:, None] <= o_q[None, :]) & m_q[None, :], exp2(b_s - b_lse[None, :]), 0)
388
+ # [BT, BS] @ [BS, BV] -> [BT, BV]
389
+ b_dv += tl.dot(b_p.to(b_do.dtype), b_do)
390
+ # [BT, BV] @ [BV, BS] -> [BT, BS]
391
+ b_dp = tl.dot(b_v, tl.trans(b_do))
392
+ # [BT, BS]
393
+ b_ds = b_p * (b_dp - b_delta[None, :])
394
+ # [BT, BS] @ [BS, BK] -> [BT, BK]
395
+ b_dk += tl.dot(b_ds.to(b_q.dtype), b_q)
396
+ if USE_G:
397
+ b_dg -= tl.sum(b_ds, 1)
398
+
399
+ for i_s in range((i_t + 1) * BT, tl.cdiv(T, BS) * BS, BS):
400
+ p_q = tl.make_block_ptr(q + (bos * HQ + i_hq) * K, (T, K), (HQ*K, 1), (i_s, 0), (BS, BK), (1, 0))
401
+ p_do = tl.make_block_ptr(do + (bos * HQ + i_hq) * V, (T, V), (HQ*V, 1), (i_s, i_v * BV), (BS, BV), (1, 0))
402
+ p_lse = tl.make_block_ptr(lse + bos * HQ + i_hq, (T,), (HQ,), (i_s,), (BS,), (0,))
403
+ p_delta = tl.make_block_ptr(delta + bos * HQ + i_hq, (T,), (HQ,), (i_s,), (BS,), (0,))
404
+
405
+ # [BS]
406
+ o_q = i_s + tl.arange(0, BS)
407
+ m_q = o_q < T
408
+ # [BS, BK]
409
+ b_q = tl.load(p_q, boundary_check=(0, 1))
410
+ # [BS, BV]
411
+ b_do = tl.load(p_do, boundary_check=(0, 1))
412
+ # [BS]
413
+ b_lse = tl.load(p_lse, boundary_check=(0,))
414
+ b_delta = tl.load(p_delta, boundary_check=(0,))
415
+ # [BT, BS]
416
+ b_s = tl.dot(b_k, tl.trans(b_q)) * scale * RCP_LN2
417
+ if USE_G:
418
+ p_gq = tl.make_block_ptr(g_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_s,), (BS,), (0,))
419
+ b_gq = tl.load(p_gq, boundary_check=(0,)).to(tl.float32)
420
+ b_s += b_gq[None, :] - b_gk[:, None]
421
+ b_p = tl.where(m_q[None, :], exp2(b_s - b_lse[None, :]), 0)
422
+ # [BT, BS] @ [BS, BV] -> [BT, BV]
423
+ b_dv += tl.dot(b_p.to(b_do.dtype), b_do)
424
+ # [BT, BV] @ [BV, BS] -> [BT, BS]
425
+ b_dp = tl.dot(b_v, tl.trans(b_do))
426
+ # [BT, BS]
427
+ b_ds = b_p * (b_dp - b_delta[None, :])
428
+ # [BT, BS] @ [BS, BK] -> [BT, BK]
429
+ b_dk += tl.dot(b_ds.to(b_q.dtype), b_q)
430
+ if USE_G:
431
+ b_dg -= tl.sum(b_ds, 1)
432
+
433
+ b_dk = b_dk * scale
434
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
435
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
436
+ if USE_G:
437
+ p_dg = tl.make_block_ptr(dg_cumsum + bos * HQ + i_hq, (T,), (HQ,), (i_t * BT,), (BT,), (0,))
438
+ tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), boundary_check=(0,))
439
+
440
+
441
+ def parallel_attn_fwd(
442
+ q: torch.Tensor,
443
+ k: torch.Tensor,
444
+ v: torch.Tensor,
445
+ g_cumsum: torch.Tensor,
446
+ scale: float,
447
+ cu_seqlens: torch.LongTensor | None = None,
448
+ ):
449
+ B, T, H, K, V = *k.shape, v.shape[-1]
450
+ HQ = q.shape[2]
451
+ G = HQ // H
452
+ BT = 128
453
+ if check_shared_mem('hopper', q.device.index):
454
+ BS = min(64, max(16, triton.next_power_of_2(T)))
455
+ BK = min(256, max(16, triton.next_power_of_2(K)))
456
+ BV = min(256, max(16, triton.next_power_of_2(V)))
457
+ num_warps = 8
458
+ elif check_shared_mem('ampere', q.device.index):
459
+ BS = min(32, max(16, triton.next_power_of_2(T)))
460
+ BK = min(256, max(16, triton.next_power_of_2(K)))
461
+ BV = min(128, max(16, triton.next_power_of_2(V)))
462
+ num_warps = 4
463
+ else:
464
+ BS = min(32, max(16, triton.next_power_of_2(T)))
465
+ BK = min(256, max(16, triton.next_power_of_2(K)))
466
+ BV = min(64, max(16, triton.next_power_of_2(V)))
467
+ num_warps = 2
468
+ NK = triton.cdiv(K, BK)
469
+ NV = triton.cdiv(V, BV)
470
+
471
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
472
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
473
+ assert NK == 1, "The key dimension can not be larger than 256"
474
+
475
+ o = torch.empty(B, T, HQ, V, dtype=v.dtype, device=q.device)
476
+ lse = torch.empty(B, T, HQ, dtype=torch.float, device=q.device)
477
+ grid = (NV, NT, B * HQ)
478
+ parallel_attn_fwd_kernel[grid](
479
+ q=q,
480
+ k=k,
481
+ v=v,
482
+ o=o,
483
+ g_cumsum=g_cumsum,
484
+ lse=lse,
485
+ scale=scale,
486
+ cu_seqlens=cu_seqlens,
487
+ chunk_indices=chunk_indices,
488
+ B=B,
489
+ T=T,
490
+ H=H,
491
+ HQ=HQ,
492
+ G=G,
493
+ K=K,
494
+ V=V,
495
+ BT=BT,
496
+ BS=BS,
497
+ BK=BK,
498
+ BV=BV,
499
+ num_warps=num_warps,
500
+ )
501
+ return o, lse
502
+
503
+
504
+ def parallel_attn_bwd_preprocess(
505
+ o: torch.Tensor,
506
+ do: torch.Tensor,
507
+ ):
508
+ V = o.shape[-1]
509
+ delta = torch.empty_like(o[..., 0], dtype=torch.float)
510
+ parallel_attn_bwd_kernel_preprocess[(delta.numel(),)](
511
+ o=o,
512
+ do=do,
513
+ delta=delta,
514
+ B=triton.next_power_of_2(V),
515
+ V=V,
516
+ )
517
+ return delta
518
+
519
+
520
+ def parallel_attn_bwd(
521
+ q: torch.Tensor,
522
+ k: torch.Tensor,
523
+ v: torch.Tensor,
524
+ o: torch.Tensor,
525
+ g_cumsum: torch.Tensor,
526
+ lse: torch.Tensor,
527
+ do: torch.Tensor,
528
+ scale: float = None,
529
+ chunk_size: int = 128,
530
+ cu_seqlens: torch.LongTensor | None = None,
531
+ ):
532
+ B, T, H, K, V = *k.shape, v.shape[-1]
533
+ HQ = q.shape[2]
534
+ G = HQ // H
535
+ if check_shared_mem('hopper'):
536
+ BT = 128
537
+ BS = 64
538
+ BK = max(triton.next_power_of_2(K), 16)
539
+ BV = max(triton.next_power_of_2(V), 16)
540
+ num_warps = 8
541
+ elif check_shared_mem('ampere'):
542
+ BS = 32
543
+ BK = max(triton.next_power_of_2(K), 16)
544
+ BV = max(triton.next_power_of_2(V), 16)
545
+ BT = 128 if K <= 64 else 64
546
+ num_warps = 4
547
+ else:
548
+ BT = 64
549
+ BS = 32
550
+ BK = max(triton.next_power_of_2(K), 16)
551
+ BV = min(max(triton.next_power_of_2(V), 16), 64)
552
+ num_warps = 2
553
+
554
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
555
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
556
+ NV = triton.cdiv(V, BV)
557
+
558
+ delta = parallel_attn_bwd_preprocess(o, do)
559
+
560
+ dq = torch.empty(B, T, HQ, K, dtype=k.dtype if H == HQ else torch.float, device=q.device)
561
+ dk = torch.empty(B, T, HQ, K, dtype=k.dtype if H == HQ else torch.float, device=q.device)
562
+ dv = torch.empty(B, T, HQ, V, dtype=v.dtype if H == HQ else torch.float, device=q.device)
563
+ grid = (NV, NT, B * HQ)
564
+
565
+ dg_cumsum, dg_cumsum_k = None, None
566
+ if g_cumsum is not None:
567
+ dg_cumsum = torch.empty(B, T, HQ, dtype=torch.float, device=q.device)
568
+ dg_cumsum_k = torch.empty(B, T, HQ, dtype=torch.float, device=q.device)
569
+
570
+ parallel_attn_bwd_kernel_dq[grid](
571
+ q=q,
572
+ k=k,
573
+ v=v,
574
+ g_cumsum=g_cumsum,
575
+ lse=lse,
576
+ delta=delta,
577
+ do=do,
578
+ dq=dq,
579
+ dg_cumsum=dg_cumsum,
580
+ cu_seqlens=cu_seqlens,
581
+ chunk_indices=chunk_indices,
582
+ scale=scale,
583
+ T=T,
584
+ B=B,
585
+ H=H,
586
+ HQ=HQ,
587
+ G=G,
588
+ K=K,
589
+ V=V,
590
+ BT=BT,
591
+ BS=BS,
592
+ BK=BK,
593
+ BV=BV,
594
+ num_warps=num_warps,
595
+ )
596
+ parallel_attn_bwd_kernel_dkv[grid](
597
+ q=q,
598
+ k=k,
599
+ v=v,
600
+ g_cumsum=g_cumsum,
601
+ lse=lse,
602
+ delta=delta,
603
+ do=do,
604
+ dk=dk,
605
+ dv=dv,
606
+ dg_cumsum=dg_cumsum_k,
607
+ cu_seqlens=cu_seqlens,
608
+ chunk_indices=chunk_indices,
609
+ scale=scale,
610
+ T=T,
611
+ B=B,
612
+ H=H,
613
+ HQ=HQ,
614
+ G=G,
615
+ K=K,
616
+ V=V,
617
+ BT=BT,
618
+ BS=BS,
619
+ BK=BK,
620
+ BV=BV,
621
+ num_warps=num_warps,
622
+ )
623
+ dk = reduce(dk, 'b t (h g) k -> b t h k', g=G, reduction='sum')
624
+ dv = reduce(dv, 'b t (h g) v -> b t h v', g=G, reduction='sum')
625
+ if g_cumsum is not None:
626
+ dg_cumsum.add_(dg_cumsum_k)
627
+ return dq, dk, dv, dg_cumsum
628
+
629
+
630
+ @torch.compile
631
+ class ParallelAttentionFunction(torch.autograd.Function):
632
+
633
+ @staticmethod
634
+ @contiguous
635
+ @autocast_custom_fwd
636
+ def forward(ctx, q, k, v, g, scale, cu_seqlens):
637
+ ctx.dtype = q.dtype
638
+
639
+ RCP_LN2: float = 1.4426950216
640
+ g_cumsum = chunk_global_cumsum(g, cu_seqlens=cu_seqlens, scale=RCP_LN2) if g is not None else None
641
+ o, lse = parallel_attn_fwd(
642
+ q=q,
643
+ k=k,
644
+ v=v,
645
+ g_cumsum=g_cumsum,
646
+ scale=scale,
647
+ cu_seqlens=cu_seqlens,
648
+ )
649
+ ctx.save_for_backward(q, k, v, o, g_cumsum, lse)
650
+ ctx.cu_seqlens = cu_seqlens
651
+ ctx.scale = scale
652
+ return o.to(q.dtype)
653
+
654
+ @staticmethod
655
+ @contiguous
656
+ @autocast_custom_bwd
657
+ def backward(ctx, do):
658
+ q, k, v, o, g_cumsum, lse = ctx.saved_tensors
659
+ dq, dk, dv, dg = parallel_attn_bwd(
660
+ q=q,
661
+ k=k,
662
+ v=v,
663
+ o=o,
664
+ g_cumsum=g_cumsum,
665
+ lse=lse,
666
+ do=do,
667
+ scale=ctx.scale,
668
+ cu_seqlens=ctx.cu_seqlens,
669
+ )
670
+ if dg is not None:
671
+ dg = chunk_global_cumsum(dg, cu_seqlens=ctx.cu_seqlens, reverse=True)
672
+
673
+ return dq.to(q), dk.to(k), dv.to(v), dg, None, None
674
+
675
+
676
+ def parallel_attn(
677
+ q: torch.Tensor,
678
+ k: torch.Tensor,
679
+ v: torch.Tensor,
680
+ g: torch.Tensor | None = None,
681
+ scale: float | None = None,
682
+ cu_seqlens: torch.LongTensor | None = None,
683
+ head_first: bool = False,
684
+ ) -> torch.Tensor:
685
+ r"""
686
+ Args:
687
+ q (torch.Tensor):
688
+ queries of shape `[B, T, HQ, K]`.
689
+ k (torch.Tensor):
690
+ keys of shape `[B, T, H, K]`.
691
+ GQA will be applied if HQ is divisible by H.
692
+ v (torch.Tensor):
693
+ values of shape `[B, T, H, V]`.
694
+ g (Optional[torch.Tensor]):
695
+ log decay factors of shape `[B, T, H]`.
696
+ scale (Optional[float]):
697
+ Scale factor for attention scores.
698
+ If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
699
+ cu_seqlens (torch.LongTensor):
700
+ Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
701
+ consistent with the FlashAttention API.
702
+ head_first (Optional[bool]):
703
+ Whether the inputs are in the head-first format. Default: `False`.
704
+ This argument has been deprecated.
705
+
706
+ Returns:
707
+ o (torch.Tensor):
708
+ Outputs of shape `[B, T, HQ, V]`.
709
+ """
710
+ if head_first:
711
+ raise DeprecationWarning(
712
+ "head_first is deprecated and will be removed in a future version. "
713
+ "Please use head_first=False for now instead.",
714
+ )
715
+ if not head_first and q.shape[1] < q.shape[2]:
716
+ warnings.warn(
717
+ f"Input tensor shape suggests potential format mismatch: seq_len ({q.shape[1]}) < num_heads ({q.shape[2]}). "
718
+ "This may indicate the inputs were passed in head-first format [B, H, T, ...] "
719
+ "when head_first=False was specified. "
720
+ "Please verify your input tensor format matches the expected shape [B, T, H, ...].",
721
+ )
722
+ if scale is None:
723
+ scale = k.shape[-1] ** -0.5
724
+ if cu_seqlens is not None:
725
+ assert q.shape[0] == 1, "batch size must be 1 when cu_seqlens are provided"
726
+
727
+ o = ParallelAttentionFunction.apply(q, k, v, g, scale, cu_seqlens)
728
+ return o
code/flash-linear-attention/fla/ops/based/__init__.py ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+
2
+ from .fused_chunk import fused_chunk_based
3
+ from .parallel import parallel_based
4
+
5
+ __all__ = [
6
+ 'fused_chunk_based',
7
+ 'parallel_based',
8
+ ]
code/flash-linear-attention/fla/ops/based/fused_chunk.py ADDED
@@ -0,0 +1,371 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
9
+
10
+
11
+ @triton.jit(do_not_specialize=['T'])
12
+ def fused_chunk_based_fwd_kernel(
13
+ q,
14
+ k,
15
+ v,
16
+ o,
17
+ z,
18
+ scale, # K ** -0.5
19
+ T,
20
+ B: tl.constexpr,
21
+ H: tl.constexpr,
22
+ K: tl.constexpr,
23
+ V: tl.constexpr,
24
+ BT: tl.constexpr,
25
+ BK: tl.constexpr,
26
+ BV: tl.constexpr,
27
+ ):
28
+ i_v, i_k, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
29
+
30
+ o_i = tl.arange(0, BT)
31
+
32
+ # [BT, BT]
33
+ m_s = o_i[:, None] >= o_i[None, :]
34
+
35
+ # [BV], zero-order taylor expansion
36
+ b_h_0o = tl.zeros([BV], dtype=tl.float32)
37
+ # [BK, BV], first-order taylor expansion
38
+ b_h_1o = tl.zeros([BK, BV], dtype=tl.float32)
39
+ # [BK, BK, BV] second-order taylor expansion
40
+ b_h_2o = tl.zeros([BK*BK, BV], dtype=tl.float32)
41
+
42
+ # make block pointers
43
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (0, i_k * BK), (BT, BK), (1, 0))
44
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (K, T), (1, K), (i_k * BK, 0), (BK, BT), (0, 1))
45
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (0, i_v * BV), (BT, BV), (1, 0))
46
+ p_o = tl.make_block_ptr(o + (i_bh + i_k*B*H) * T*V, (T, V), (V, 1), (0, i_v * BV), (BT, BV), (1, 0))
47
+
48
+ p_z = z + (i_bh + i_k * B * H) * T + tl.arange(0, BT)
49
+ k_2o = tl.zeros([1, BK * BK], dtype=tl.float32)
50
+ k_1o = tl.zeros([1, BK], dtype=tl.float32)
51
+ k_0o = 0
52
+
53
+ for i in range(0, tl.cdiv(T, BT)):
54
+ # [BK, BT]
55
+ b_k = tl.load(p_k, boundary_check=(0, 1))
56
+ # [BK*BK, BT]
57
+ b_k_2o = b_k[:, None, :] * b_k[None, :, :]
58
+ b_k_2o = tl.reshape(b_k_2o, [BK * BK, BT]).to(b_k.dtype)
59
+ # [BT, BV]
60
+ b_v = tl.load(p_v, boundary_check=(0, 1))
61
+ # [BT, BK]
62
+ b_q = (tl.load(p_q, boundary_check=(0, 1)) * scale).to(b_k.dtype)
63
+ b_o = tl.zeros([BT, BV], dtype=tl.float32)
64
+ b_z = tl.zeros([BT], dtype=tl.float32)
65
+
66
+ # interchunk
67
+ # zero-order
68
+ b_o += b_h_0o
69
+ b_z += k_0o
70
+ # first-order
71
+ b_o += tl.dot(b_q, b_h_1o.to(b_q.dtype), allow_tf32=False)
72
+ b_z += tl.sum(b_q * k_1o, axis=1)
73
+ # second-order
74
+ b_q_2o = b_q[:, :, None] * b_q[:, None, :]
75
+ b_q_2o = tl.reshape(b_q_2o, [BT, BK * BK]).to(b_k.dtype)
76
+ b_o += tl.dot(b_q_2o, b_h_2o.to(b_q_2o.dtype), allow_tf32=False) * 0.5
77
+ b_z += tl.sum(b_q_2o * k_2o, axis=1) * 0.5
78
+
79
+ # update running statistics
80
+ k_1o += tl.sum(b_k, axis=1)[None, :]
81
+ k_2o += tl.sum(b_k_2o, axis=1)[None, :]
82
+ k_0o += BT
83
+
84
+ # intrachunk
85
+ # [BT, BT]
86
+ b_s = tl.dot(b_q, b_k, allow_tf32=False)
87
+ b_s = 1 + b_s + 0.5 * b_s * b_s
88
+ b_s = tl.where(m_s, b_s, 0)
89
+ b_z += tl.sum(b_s, axis=1)
90
+ b_o += tl.dot(b_s.to(b_q.dtype), b_v, allow_tf32=False)
91
+ # [TB, BV]
92
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
93
+ tl.store(p_z, b_z.to(p_z.dtype.element_ty), mask=(i * BT + tl.arange(0, BT)) < T)
94
+
95
+ # update hidden state
96
+ # [BK, BV]
97
+ b_h_2o = b_h_2o + tl.dot(b_k_2o.to(b_v.dtype), b_v, allow_tf32=False)
98
+ b_h_1o = b_h_1o + tl.dot(b_k, b_v, allow_tf32=False)
99
+ b_h_0o = b_h_0o + tl.sum(b_v, axis=0)
100
+
101
+ p_q = tl.advance(p_q, (BT, 0))
102
+ p_k = tl.advance(p_k, (0, BT))
103
+ p_v = tl.advance(p_v, (BT, 0))
104
+ p_o = tl.advance(p_o, (BT, 0))
105
+ p_z += BT
106
+
107
+
108
+ # Similar to Algorithm1 of https://arxiv.org/abs/2006.16236
109
+ @triton.jit
110
+ def fused_chunk_based_bwd_kernel(
111
+ # NV: number of split in the V dimension. NK: number of split in the K dimension
112
+ q,
113
+ k,
114
+ v,
115
+ do,
116
+ dz,
117
+ dq,
118
+ dk,
119
+ dv,
120
+ scale, # K ** -0.5
121
+ T,
122
+ B: tl.constexpr,
123
+ H: tl.constexpr,
124
+ K: tl.constexpr,
125
+ V: tl.constexpr,
126
+ BT: tl.constexpr,
127
+ BK: tl.constexpr,
128
+ BV: tl.constexpr,
129
+ ):
130
+ i_v, i_k, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
131
+
132
+ o_i = tl.arange(0, BT)
133
+ m_s = o_i[:, None] >= o_i[None, :]
134
+
135
+ # [BV], zero-order taylor expansion
136
+ # b_h_0o = tl.zeros([BV], dtype=tl.float32)
137
+ # [BK, BV], first-order taylor expansion
138
+ b_h_1o = tl.zeros([BV, BK], dtype=tl.float32)
139
+ # [BK, BK, BV] second-order taylor expansion
140
+ b_h_2o = tl.zeros([BV, BK*BK], dtype=tl.float32)
141
+
142
+ k_1o = tl.zeros([1, BK], dtype=tl.float32)
143
+ k_2o = tl.zeros([1, BK * BK], dtype=tl.float32)
144
+
145
+ for i in range(0, tl.cdiv(T, BT)):
146
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (i * BT, i_k * BK), (BT, BK), (1, 0))
147
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (i * BT, i_k * BK), (BT, BK), (1, 0))
148
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (V, T), (1, V), (i_v * BV, i * BT), (BV, BT), (0, 1))
149
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i * BT, i_v * BV), (BT, BV), (1, 0))
150
+ p_dq = tl.make_block_ptr(dq + (i_bh + i_v*B*H) * T*K, (T, K), (K, 1), (i*BT, i_k*BK), (BT, BK), (1, 0))
151
+ p_dz = dz + (i_bh) * T + tl.arange(0, BT) + i * BT
152
+ b_dq = tl.zeros([BT, BK], dtype=tl.float32)
153
+
154
+ # load tensors
155
+ # [BT, BK]
156
+ b_q = tl.load(p_q, boundary_check=(0, 1))
157
+ b_q = (b_q * scale).to(b_q.dtype)
158
+ b_k = tl.load(p_k, boundary_check=(0, 1))
159
+ b_do = tl.load(p_do, boundary_check=(0, 1)).to(b_q.dtype)
160
+ b_dz = tl.load(p_dz, mask=(tl.arange(0, BT) + i * BT) < T)
161
+ # [BV, BT]
162
+ b_v = tl.load(p_v, boundary_check=(0, 1))
163
+
164
+ # inter-chunk
165
+ b_dq += tl.dot(b_do, (b_h_1o).to(b_do.dtype), allow_tf32=False)
166
+ if i_v == 0:
167
+ b_dq += b_dz[:, None] * k_1o
168
+ b_dq_2o = tl.dot(b_do, (b_h_2o).to(b_do.dtype), allow_tf32=False) * 0.5
169
+ if i_v == 0:
170
+ b_dq_2o += (b_dz[:, None] * k_2o) * 0.5
171
+ b_dq_2o = tl.reshape(b_dq_2o, [BT, BK, BK])
172
+ b_dq += tl.sum(b_dq_2o * b_q[:, :, None], axis=1)
173
+ b_dq += tl.sum(b_dq_2o * b_q[:, None, :], axis=2)
174
+ b_dq *= scale
175
+
176
+ # intra-chunk
177
+ # [BT, BT]
178
+ b_ds = tl.dot(b_do, b_v, allow_tf32=False)
179
+ if i_v == 0:
180
+ b_ds += b_dz[:, None]
181
+ b_ds = tl.where(m_s, b_ds, 0) * scale
182
+ b_s = tl.dot(b_q, tl.trans(b_k), allow_tf32=False)
183
+ b_s = tl.where(m_s, b_s, 0)
184
+ b_dq += tl.dot((b_ds * (1 + b_s)).to(b_q.dtype), b_k, allow_tf32=False)
185
+
186
+ # store
187
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
188
+
189
+ # update hidden state
190
+ # [BT, BK*BK]
191
+ b_k_2o = b_k[:, :, None] * b_k[:, None, :]
192
+ b_k_2o = tl.reshape(b_k_2o, [BT, BK * BK]).to(b_k.dtype)
193
+ # [BV, BK*BK]
194
+ b_h_2o = b_h_2o + tl.dot(b_v, b_k_2o.to(b_v.dtype), allow_tf32=False)
195
+ # [BV, BK]
196
+ b_h_1o = b_h_1o + tl.dot(b_v, b_k, allow_tf32=False)
197
+
198
+ if i_v == 0:
199
+ # update running statistics
200
+ k_1o += tl.sum(b_k, axis=0)[None, :]
201
+ k_2o += tl.sum(b_k_2o, axis=0)[None, :]
202
+
203
+ tl.debug_barrier()
204
+ b_h_1o = None
205
+ b_h_2o = None
206
+
207
+ # [BK, BV], first-order taylor expansion
208
+ b_dh_1o = tl.zeros([BK, BV], dtype=tl.float32)
209
+ # [BK, BK, BV] second-order taylor expansion
210
+ b_dh_2o = tl.zeros([BK*BK, BV], dtype=tl.float32)
211
+ b_dh_0o = tl.zeros([BV], dtype=tl.float32)
212
+ m_s = tl.arange(0, BT)[:, None] <= tl.arange(0, BT)[None, :]
213
+
214
+ dq_1o = tl.zeros([1, BK], dtype=tl.float32)
215
+ dq_2o = tl.zeros([BK * BK, 1], dtype=tl.float32)
216
+
217
+ for i in range(tl.cdiv(T, BT) * BT - BT, -BT, -BT):
218
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (K, T), (1, K), (i_k * BK, i), (BK, BT), (0, 1))
219
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (i, i_k * BK), (BT, BK), (1, 0))
220
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i, i_v * BV), (BT, BV), (1, 0))
221
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i, i_v * BV), (BT, BV), (1, 0))
222
+ p_dk = tl.make_block_ptr(dk + (i_bh+i_v*B*H) * T*K, (T, K), (K, 1), (i, i_k*BK), (BT, BK), (1, 0))
223
+ p_dv = tl.make_block_ptr(dv + (i_bh+i_k*B*H) * T*V, (T, V), (V, 1), (i, i_v*BV), (BT, BV), (1, 0))
224
+ p_dz = dz + (i_bh) * T + tl.arange(0, BT) + i
225
+
226
+ b_dk = tl.zeros([BT, BK], dtype=tl.float32)
227
+ b_dv = tl.zeros([BT, BV], dtype=tl.float32)
228
+
229
+ b_q = tl.load(p_q, boundary_check=(0, 1))
230
+ b_k = tl.load(p_k, boundary_check=(0, 1))
231
+ b_v = tl.load(p_v, boundary_check=(0, 1))
232
+ b_do = tl.load(p_do, boundary_check=(0, 1)).to(b_q.dtype)
233
+ b_dz = tl.load(p_dz, mask=(tl.arange(0, BT)+i) < T)
234
+ b_q = (b_q * scale).to(b_k.dtype)
235
+
236
+ # intra chunk
237
+ b_ds = tl.dot(b_v, tl.trans(b_do), allow_tf32=False)
238
+ if i_v == 0:
239
+ b_ds += b_dz[None, :]
240
+ b_ds = tl.where(m_s, b_ds, 0)
241
+ b_s = tl.dot(b_k, b_q, allow_tf32=False)
242
+ b_s2 = 1 + b_s + 0.5 * b_s * b_s
243
+ b_s = tl.where(m_s, b_s, 0)
244
+ b_s2 = tl.where(m_s, b_s2, 0)
245
+ b_ds *= (1+b_s)
246
+
247
+ b_dk += tl.dot(b_ds.to(b_k.dtype), tl.trans(b_q), allow_tf32=False)
248
+ b_dv += tl.dot(b_s2.to(b_do.dtype), b_do, allow_tf32=False)
249
+
250
+ # inter chunk
251
+ b_k_2o = b_k[:, :, None] * b_k[:, None, :]
252
+ b_k_2o = tl.reshape(b_k_2o, [BT, BK * BK]).to(b_k.dtype)
253
+
254
+ b_dv += tl.dot(b_k, b_dh_1o.to(b_k.dtype), allow_tf32=False)
255
+ b_dv += tl.dot(b_k_2o, b_dh_2o.to(b_k.dtype), allow_tf32=False)
256
+ b_dv += b_dh_0o
257
+
258
+ b_dk += tl.dot(b_v, tl.trans(b_dh_1o).to(b_k.dtype), allow_tf32=False)
259
+
260
+ if i_v == 0:
261
+ b_dk += dq_1o
262
+
263
+ b_dk_2o = tl.dot(b_dh_2o.to(b_k.dtype), tl.trans(b_v), allow_tf32=False)
264
+ if i_v == 0:
265
+ b_dk_2o += dq_2o
266
+ b_dk_2o = tl.reshape(b_dk_2o, [BK, BK, BT])
267
+ b_k_fp32 = tl.trans(b_k.to(tl.float32))
268
+ b_dk2 = tl.sum(b_dk_2o * b_k_fp32[:, None, :], axis=0)
269
+ b_dk2 += tl.sum(b_dk_2o * b_k_fp32[None, :, :], axis=1)
270
+ b_dk += tl.trans(b_dk2)
271
+
272
+ # hidden state update
273
+ b_dh_0o += tl.sum(b_do, axis=0)
274
+ b_dh_1o = b_dh_1o + tl.dot(b_q, b_do, allow_tf32=False)
275
+ b_q_2o = b_q[None, :, :] * b_q[:, None, :]
276
+ b_q_2o = tl.reshape(b_q_2o, [BK * BK, BT]).to(b_k.dtype)
277
+ b_dh_2o = b_dh_2o + tl.dot(b_q_2o, b_do, allow_tf32=False) * 0.5
278
+
279
+ if i_v == 0:
280
+ dq_1o += (tl.sum(b_dz[None, :] * b_q, axis=1))[None, :]
281
+ dq_2o += (tl.sum(b_dz[None, :] * b_q_2o, axis=1) * 0.5)[:, None]
282
+
283
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
284
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
285
+
286
+
287
+ class FusedChunkBasedFunction(torch.autograd.Function):
288
+
289
+ @staticmethod
290
+ @input_guard
291
+ @autocast_custom_fwd
292
+ def forward(ctx, q, k, v, scale=1):
293
+ B, H, T, K, V = *k.shape, v.shape[-1]
294
+
295
+ scale = scale
296
+ BT = 16
297
+ BK, BV = min(K, 16), min(V, 32)
298
+ BK, BV = max(BK, 16), max(BV, 16)
299
+ NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
300
+
301
+ num_warps = 4
302
+
303
+ # the norm of o might explode, so we need to use float32 here
304
+ o = q.new_empty(NK, B, H, T, V, dtype=torch.float32)
305
+ z = q.new_empty(NK, B, H, T, dtype=torch.float32)
306
+
307
+ grid = (NV, NK, B * H)
308
+ fused_chunk_based_fwd_kernel[grid](
309
+ q, k, v, o, z,
310
+ scale,
311
+ T=T, B=B, H=H, K=K, V=V, BT=BT, BK=BK, BV=BV,
312
+ num_warps=num_warps,
313
+ )
314
+ o = o.sum(0)
315
+ z = z.sum(0)
316
+ ctx.save_for_backward(q, k, v)
317
+ ctx.scale = scale
318
+ return o.to(q.dtype), z.to(z.dtype)
319
+
320
+ @staticmethod
321
+ @input_guard
322
+ @autocast_custom_bwd
323
+ def backward(ctx, do, dz):
324
+ q, k, v = ctx.saved_tensors
325
+ B, H, T, K, V = *k.shape, v.shape[-1]
326
+ scale = ctx.scale
327
+
328
+ BT = 16
329
+ BK, BV = min(K, 16), min(V, 32)
330
+ BK, BV = max(BK, 16), max(BV, 16)
331
+ NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
332
+ num_stages = 1
333
+ num_warps = 4
334
+
335
+ dq = q.new_empty(NV, B, H, T, K)
336
+ dk = q.new_empty(NV, B, H, T, K)
337
+ dv = q.new_empty(NK, B, H, T, V)
338
+ grid = (NV, NK, B * H)
339
+
340
+ fused_chunk_based_bwd_kernel[grid](
341
+ q, k, v, do, dz, dq, dk, dv,
342
+ scale,
343
+ T=T, B=B, H=H, K=K, V=V, BT=BT, BK=BK, BV=BV,
344
+ num_warps=num_warps,
345
+ num_stages=num_stages,
346
+ )
347
+ dq = dq.sum(0)
348
+ dk = dk.sum(0)
349
+ dv = dv.sum(0)
350
+ return dq.to(q.dtype), dk.to(k.dtype), dv.to(v.dtype), None
351
+
352
+
353
+ def fused_chunk_based(
354
+ q: torch.Tensor,
355
+ k: torch.Tensor,
356
+ v: torch.Tensor,
357
+ scale: float | None = None,
358
+ use_norm: bool = True,
359
+ head_first: bool = False,
360
+ ):
361
+ assert q.shape[-1] <= 16, 'only support feature dimension up to 16.'
362
+ if scale is None:
363
+ scale = q.shape[-1] ** -0.5
364
+ if not head_first:
365
+ q, k, v = map(lambda x: x.transpose(1, 2), (q, k, v))
366
+ o, z = FusedChunkBasedFunction.apply(q, k, v, scale)
367
+ if use_norm:
368
+ o = o / (z[..., None] + 1e-6)
369
+ if not head_first:
370
+ o = o.transpose(1, 2)
371
+ return o.to(q.dtype)
code/flash-linear-attention/fla/ops/based/naive.py ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+
3
+ import torch
4
+ from einops import rearrange
5
+
6
+
7
+ def naive_parallel_based(
8
+ q: torch.Tensor,
9
+ k: torch.Tensor,
10
+ v: torch.Tensor,
11
+ scale: float | None = None,
12
+ use_norm: bool = True,
13
+ ):
14
+ if scale is None:
15
+ scale = q.shape[-1] ** -0.5
16
+ q = q * scale
17
+ attn = q @ k.transpose(-2, -1)
18
+ attn = 1 + attn + 1/2 * (attn ** 2)
19
+ attn.masked_fill_(~torch.tril(torch.ones(
20
+ q.shape[-2], q.shape[-2], dtype=torch.bool, device=q.device)), 0)
21
+ o = attn @ v
22
+ if use_norm:
23
+ z = attn.sum(-1)
24
+ return o / (z[..., None] + 1e-6)
25
+ else:
26
+ return o
27
+
28
+
29
+ def naive_chunk_based(q, k, v, chunk_size=256):
30
+ q = q * (q.shape[-1] ** -0.5)
31
+ # compute normalizer.
32
+ k_cumsum = torch.cumsum(k, dim=-2)
33
+ kk_cumsum = torch.cumsum(k.unsqueeze(-1) * k.unsqueeze(-2), dim=-3)
34
+ # first
35
+ z = (q * k_cumsum).sum(-1)
36
+ # second order
37
+ z += (q.unsqueeze(-1) * q.unsqueeze(-2) * kk_cumsum).sum((-1, -2)) * 0.5
38
+ # zero-th order
39
+ z += (torch.arange(0, q.shape[-2]).to(z.device) * 1.0 + 1.0)[None, None, :]
40
+
41
+ # compute o
42
+ # constant term
43
+ _o = v.cumsum(-2)
44
+
45
+ q = rearrange(q, 'b h (n c) d -> b h n c d', c=chunk_size)
46
+
47
+ k = rearrange(k, 'b h (n c) d -> b h n c d', c=chunk_size)
48
+ v = rearrange(v, 'b h (n c) d -> b h n c d', c=chunk_size)
49
+
50
+ intra_chunk_attn = q @ k.transpose(-2, -1)
51
+ intra_chunk_attn = intra_chunk_attn + 1/2 * (intra_chunk_attn ** 2)
52
+ intra_chunk_attn.masked_fill_(~torch.tril(torch.ones(chunk_size, chunk_size, dtype=torch.bool, device=q.device)), 0)
53
+ o = intra_chunk_attn @ v
54
+
55
+ # quadractic term
56
+ kv = torch.einsum('b h n c x, b h n c y, b h n c z -> b h n x y z', k, k, v)
57
+ kv = kv.cumsum(2)
58
+ kv = torch.cat([torch.zeros_like(kv[:, :, :1]), kv[:, :, :-1]], dim=2)
59
+
60
+ o += 0.5 * torch.einsum('b h n x y z, b h n c x, b h n c y -> b h n c z', kv, q, q)
61
+
62
+ # linear term
63
+ kv = torch.einsum('b h n c x, b h n c y -> b h n x y', k, v)
64
+ kv = kv.cumsum(2)
65
+ kv = torch.cat([torch.zeros_like(kv[:, :, :1]), kv[:, :, :-1]], dim=2)
66
+ o += torch.einsum('b h n x y, b h n c x -> b h n c y', kv, q)
67
+
68
+ o = rearrange(o, 'b h n c d -> b h (n c) d')
69
+ o = o + _o
70
+ return o / (z[..., None] + 1e-6)
code/flash-linear-attention/fla/ops/based/parallel.py ADDED
@@ -0,0 +1,406 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
9
+
10
+ # Based: An Educational and Effective Sequence Mixer
11
+ # https://hazyresearch.stanford.edu/blog/2023-12-11-zoology2-based
12
+
13
+
14
+ @triton.jit(do_not_specialize=['T'])
15
+ def parallel_based_fwd_kernel(
16
+ q,
17
+ k,
18
+ v,
19
+ o,
20
+ z,
21
+ scale,
22
+ T,
23
+ B: tl.constexpr,
24
+ H: tl.constexpr,
25
+ K: tl.constexpr,
26
+ V: tl.constexpr,
27
+ BTL: tl.constexpr,
28
+ BTS: tl.constexpr,
29
+ BK: tl.constexpr,
30
+ BV: tl.constexpr,
31
+ ):
32
+ # i_c: chunk index. used for sequence parallelism
33
+ i_kv, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
34
+ NV = tl.cdiv(V, BV)
35
+ i_k = i_kv // (NV)
36
+ i_v = i_kv % (NV)
37
+
38
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (T, K), (K, 1), (i_c * BTL, i_k * BK), (BTL, BK), (1, 0))
39
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (K, T), (1, K), (i_k * BK, 0), (BK, BTS), (0, 1))
40
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (0, i_v * BV), (BTS, BV), (1, 0))
41
+
42
+ # [BQ, BD] block Q, in the shared memory throughout the whole kernel
43
+ b_q = tl.load(p_q, boundary_check=(0, 1))
44
+ b_q = (b_q * scale).to(b_q.dtype)
45
+ b_o = tl.zeros([BTL, BV], dtype=tl.float32)
46
+ b_z = tl.zeros([BTL], dtype=tl.float32)
47
+
48
+ # Q block and K block have no overlap
49
+ # no need for mask, thereby saving flops
50
+ for _ in range(0, i_c * BTL, BTS):
51
+ # [BK, BTS]
52
+ b_k = tl.load(p_k, boundary_check=(0, 1))
53
+
54
+ # [BTS, BV]
55
+ b_v = tl.load(p_v, boundary_check=(0, 1))
56
+ # [BTL, BTS]
57
+ b_s = tl.dot(b_q, (b_k), allow_tf32=False)
58
+ b_s = 1 + b_s + 0.5 * b_s * b_s
59
+ b_z += tl.sum(b_s, axis=1)
60
+
61
+ # [BQ, BD]
62
+ b_o = b_o + tl.dot(b_s.to(b_v.dtype), b_v, allow_tf32=False)
63
+ p_k = tl.advance(p_k, (0, BTS))
64
+ p_v = tl.advance(p_v, (BTS, 0))
65
+
66
+ # # rescale interchunk output
67
+ tl.debug_barrier()
68
+ o_q = tl.arange(0, BTL)
69
+ # # sync threads, easy for compiler to optimize
70
+ # tl.debug_barrier()
71
+
72
+ o_k = tl.arange(0, BTS)
73
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (K, T), (1, K), (i_k * BK, i_c * BTL), (BK, BTS), (0, 1))
74
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i_c * BTL, i_v * BV), (BTS, BV), (1, 0))
75
+ # Q block and K block have overlap. masks required
76
+ for _ in range(i_c * BTL, (i_c + 1) * BTL, BTS):
77
+ # [BK, BTS]
78
+ b_k = tl.load(p_k, boundary_check=(0, 1))
79
+ # [BTS, BV]
80
+ b_v = tl.load(p_v, boundary_check=(0, 1))
81
+ # [BTL, BTS]
82
+ m_s = o_q[:, None] >= o_k[None, :]
83
+ b_s = tl.dot(b_q, b_k, allow_tf32=False)
84
+ b_s = 1 + b_s + 0.5 * b_s * b_s
85
+ b_s = tl.where(m_s, b_s, 0)
86
+ b_z += tl.sum(b_s, axis=1)
87
+ # [BTL, BV]
88
+ b_o += tl.dot(b_s.to(b_q.dtype), b_v, allow_tf32=False)
89
+
90
+ p_k = tl.advance(p_k, (0, BTS))
91
+ p_v = tl.advance(p_v, (BTS, 0))
92
+ o_k += BTS
93
+
94
+ p_o = tl.make_block_ptr(o + (i_bh + B * H * i_k) * T*V, (T, V), (V, 1), (i_c*BTL, i_v*BV), (BTL, BV), (1, 0))
95
+ p_z = z + (i_bh + B * H * i_k) * T + i_c * BTL + tl.arange(0, BTL)
96
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
97
+ tl.store(p_z, b_z.to(p_z.dtype.element_ty), mask=((i_c * BTL + tl.arange(0, BTL)) < T))
98
+
99
+
100
+ @triton.jit
101
+ def _parallel_based_bwd_dq(
102
+ i_bh,
103
+ i_c,
104
+ i_k,
105
+ i_v,
106
+ q,
107
+ k,
108
+ v,
109
+ do,
110
+ dz,
111
+ dq,
112
+ scale,
113
+ T,
114
+ B: tl.constexpr,
115
+ H: tl.constexpr,
116
+ BTL: tl.constexpr,
117
+ BTS: tl.constexpr,
118
+ BK: tl.constexpr,
119
+ BV: tl.constexpr,
120
+ K: tl.constexpr,
121
+ V: tl.constexpr,
122
+ ):
123
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (T, V), (V, 1), (i_c * BTL, i_v * BV), (BTL, BV), (1, 0))
124
+ p_q = tl.make_block_ptr(q + (i_bh) * T*K, (T, K), (K, 1), (i_c*BTL, i_k*BK), (BTL, BK), (1, 0))
125
+ b_q = tl.load(p_q, boundary_check=(0, 1))
126
+ b_q = (b_q * scale).to(b_q.dtype)
127
+
128
+ b_do = tl.load(p_do, boundary_check=(0, 1)).to(b_q.dtype)
129
+ b_dq = tl.zeros([BTL, BK], dtype=tl.float32)
130
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (0, i_k * BK), (BTS, BK), (1, 0))
131
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (V, T), (1, V), (i_v * BV, 0), (BV, BTS), (0, 1))
132
+ p_dz = dz + i_bh * T + i_c * BTL + tl.arange(0, BTL)
133
+ b_dz = tl.load(p_dz, mask=(i_c * BTL + tl.arange(0, BTL)) < T)
134
+
135
+ for _ in range(0, i_c * BTL, BTS):
136
+ # [BTS, BK]
137
+ b_k = tl.load(p_k, boundary_check=(0, 1))
138
+ # [BV, BTS]
139
+ b_v = tl.load(p_v, boundary_check=(0, 1))
140
+ # [BTL, BTS]
141
+ b_ds = tl.dot(b_do, b_v, allow_tf32=False)
142
+ if i_v == 0:
143
+ b_ds += b_dz[:, None]
144
+ else:
145
+ b_ds = b_ds
146
+ b_s = tl.dot(b_q, tl.trans(b_k), allow_tf32=False)
147
+ # [BQ, BD]
148
+ b_dq += tl.dot((b_ds * (1 + b_s)).to(b_v.dtype), b_k, allow_tf32=False)
149
+ p_k = tl.advance(p_k, (BTS, 0))
150
+ p_v = tl.advance(p_v, (0, BTS))
151
+
152
+ b_dq *= scale
153
+ o_q = tl.arange(0, BTL)
154
+ o_k = tl.arange(0, BTS)
155
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (i_c * BTL, i_k * BK), (BTS, BK), (1, 0))
156
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (V, T), (1, V), (i_v * BV, i_c * BTL), (BV, BTS), (0, 1))
157
+ # Q block and K block have overlap. masks required
158
+ for _ in range(i_c * BTL, (i_c + 1) * BTL, BTS):
159
+ # [BTS, BK]
160
+ b_k = tl.load(p_k, boundary_check=(0, 1))
161
+ # [BV, BTS]
162
+ b_v = tl.load(p_v, boundary_check=(0, 1))
163
+ # [BTL, BTS]
164
+ m_s = o_q[:, None] >= o_k[None, :]
165
+ b_ds = tl.dot(b_do, b_v, allow_tf32=False)
166
+ if i_v == 0:
167
+ b_ds += b_dz[:, None]
168
+ else:
169
+ b_ds = b_ds
170
+ b_ds = tl.where(m_s, b_ds, 0) * scale
171
+ b_s = tl.dot(b_q, tl.trans(b_k), allow_tf32=False)
172
+ b_s = tl.where(m_s, b_s, 0)
173
+ # [BTL, BK]
174
+ b_dq += tl.dot((b_ds + b_ds * b_s).to(b_k.dtype), b_k, allow_tf32=False)
175
+ p_k = tl.advance(p_k, (BTS, 0))
176
+ p_v = tl.advance(p_v, (0, BTS))
177
+ o_k += BTS
178
+ p_dq = tl.make_block_ptr(dq + (i_bh + B * H * i_v) * T*K, (T, K), (K, 1), (i_c*BTL, i_k*BK), (BTL, BK), (1, 0))
179
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
180
+ return
181
+
182
+
183
+ @triton.jit
184
+ def _parallel_based_bwd_dkv(
185
+ i_bh,
186
+ i_c,
187
+ i_k,
188
+ i_v,
189
+ q,
190
+ k,
191
+ v,
192
+ do,
193
+ dz,
194
+ dk,
195
+ dv,
196
+ scale,
197
+ T,
198
+ B: tl.constexpr,
199
+ H: tl.constexpr,
200
+ BTL: tl.constexpr,
201
+ BTS: tl.constexpr,
202
+ BK: tl.constexpr,
203
+ BV: tl.constexpr,
204
+ K: tl.constexpr,
205
+ V: tl.constexpr,
206
+ ):
207
+ # compute dk dv
208
+ p_k = tl.make_block_ptr(k + i_bh * T*K, (T, K), (K, 1), (i_c * BTL, i_k * BK), (BTL, BK), (1, 0))
209
+ p_v = tl.make_block_ptr(v + i_bh * T*V, (T, V), (V, 1), (i_c * BTL, i_v * BV), (BTL, BV), (1, 0))
210
+ b_k, b_v = tl.load(p_k, boundary_check=(0, 1)), tl.load(p_v, boundary_check=(0, 1))
211
+ b_dk, b_dv = tl.zeros([BTL, BK], dtype=tl.float32), tl.zeros([BTL, BV], dtype=tl.float32)
212
+
213
+ for i in range((tl.cdiv(T, BTS) * BTS)-BTS, (i_c + 1) * BTL - BTS, -BTS):
214
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (K, T), (1, K), (i_k * BK, i), (BK, BTS), (0, 1))
215
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (V, T), (1, V), (i_v * BV, i), (BV, BTS), (0, 1))
216
+ p_dz = dz + i_bh * T + i + tl.arange(0, BTS)
217
+ b_q = tl.load(p_q, boundary_check=(0, 1)) # [BK, BTS]
218
+ b_do = tl.load(p_do, boundary_check=(0, 1)).to(b_q.dtype) # [BV, BTS]
219
+ b_dz = tl.load(p_dz, mask=(i + tl.arange(0, BTS)) < T)
220
+ b_s = tl.dot(b_k.to(b_q.dtype), b_q, allow_tf32=False) * scale # [BTL, BTS]
221
+ b_s2 = 1 + b_s + 0.5 * b_s * b_s
222
+ b_dv += tl.dot(b_s2.to(b_q.dtype), tl.trans(b_do), allow_tf32=False)
223
+ b_ds = tl.dot(b_v, b_do, allow_tf32=False) * scale
224
+ if i_v == 0:
225
+ b_ds += b_dz[None, :] * scale
226
+ else:
227
+ b_ds = b_ds
228
+ b_dk += tl.dot((b_ds + b_ds * b_s).to(b_q.dtype), tl.trans(b_q), allow_tf32=False)
229
+
230
+ tl.debug_barrier()
231
+ o_q, o_k = tl.arange(0, BTS), tl.arange(0, BTL)
232
+ for i in range(i_c*BTL, (i_c+1)*BTL, BTS):
233
+ p_q = tl.make_block_ptr(q + i_bh * T*K, (K, T), (1, K), (i_k * BK, i), (BK, BTS), (0, 1))
234
+ p_do = tl.make_block_ptr(do + i_bh * T*V, (V, T), (1, V), (i_v * BV, i), (BV, BTS), (0, 1))
235
+ p_dz = dz + i_bh * T + i + tl.arange(0, BTS)
236
+ b_q = tl.load(p_q, boundary_check=(0, 1)) # [BD, BQ]
237
+ b_do = tl.load(p_do, boundary_check=(0, 1)).to(b_q.dtype)
238
+ b_dz = tl.load(p_dz, mask=(i + tl.arange(0, BTS)) < T)
239
+ # [BK, BQ]
240
+ m_s = o_k[:, None] <= o_q[None, :]
241
+ b_s = tl.dot(b_k, b_q, allow_tf32=False) * scale
242
+ b_s2 = 1 + b_s + 0.5 * b_s * b_s
243
+ b_s = tl.where(m_s, b_s, 0)
244
+ b_s2 = tl.where(m_s, b_s2, 0)
245
+
246
+ b_ds = tl.dot(b_v, b_do, allow_tf32=False)
247
+ if i_v == 0:
248
+ b_ds += b_dz[None, :]
249
+ else:
250
+ b_ds = b_ds
251
+ b_ds = tl.where(m_s, b_ds, 0) * scale
252
+ # [BK, BD]
253
+ b_dv += tl.dot(b_s2.to(b_q.dtype), tl.trans(b_do), allow_tf32=False)
254
+ b_dk += tl.dot((b_ds + b_ds * b_s).to(b_q.dtype), tl.trans(b_q), allow_tf32=False)
255
+ o_q += BTS
256
+
257
+ p_dk = tl.make_block_ptr(dk + (i_bh + B * H * i_v) * T*K, (T, K), (K, 1), (i_c*BTL, i_k*BK), (BTL, BK), (1, 0))
258
+ p_dv = tl.make_block_ptr(dv + (i_bh + B * H * i_k) * T*V, (T, V), (V, 1), (i_c*BTL, i_v*BV), (BTL, BV), (1, 0))
259
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
260
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
261
+ return
262
+
263
+
264
+ @triton.jit(do_not_specialize=['T'])
265
+ def parallel_based_bwd_kernel(
266
+ q,
267
+ k,
268
+ v,
269
+ do,
270
+ dz,
271
+ dq,
272
+ dk,
273
+ dv,
274
+ scale,
275
+ T,
276
+ B: tl.constexpr,
277
+ H: tl.constexpr,
278
+ K: tl.constexpr,
279
+ V: tl.constexpr,
280
+ BTL: tl.constexpr,
281
+ BTS: tl.constexpr,
282
+ BK: tl.constexpr,
283
+ BV: tl.constexpr,
284
+ ):
285
+ i_kv, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
286
+ NV = tl.cdiv(V, BV)
287
+ i_k = i_kv // (NV)
288
+ i_v = i_kv % NV
289
+ _parallel_based_bwd_dq(
290
+ i_bh, i_c, i_k, i_v,
291
+ q, k, v, do, dz, dq,
292
+ scale, T, B, H, BTL, BTS, BK, BV, K, V,
293
+ )
294
+ tl.debug_barrier()
295
+ _parallel_based_bwd_dkv(
296
+ i_bh, i_c, i_k, i_v,
297
+ q, k, v, do, dz, dk, dv,
298
+ scale, T, B, H, BTL, BTS, BK, BV, K, V,
299
+ )
300
+
301
+
302
+ class ParallelBasedFunction(torch.autograd.Function):
303
+
304
+ @staticmethod
305
+ @input_guard
306
+ @autocast_custom_fwd
307
+ def forward(ctx, q, k, v, scale):
308
+ BTL, BTS = 128, 32
309
+ assert BTL % BTS == 0
310
+ # assert q.shape[-1] % 16 == 0
311
+ BK = min(128, max(triton.next_power_of_2(k.shape[-1]), 16))
312
+ BV = min(128, max(triton.next_power_of_2(v.shape[-1]), 16))
313
+ B, H, T, K, V = *k.shape, v.shape[-1]
314
+ num_stages = 2
315
+ num_warps = 4
316
+ NK = triton.cdiv(K, BK)
317
+ NV = triton.cdiv(V, BV)
318
+ grid = (NK * NV, triton.cdiv(T, BTL), B * H)
319
+
320
+ assert NK == 1, "will encounter some synchronization issue if not."
321
+
322
+ o = torch.empty(NK, B, H, T, V, device=q.device)
323
+ z = torch.empty(NK, B, H, T, device=q.device)
324
+ parallel_based_fwd_kernel[grid](
325
+ q, k, v, o, z,
326
+ scale,
327
+ B=B,
328
+ H=H,
329
+ T=T,
330
+ K=K,
331
+ V=V,
332
+ BTL=BTL,
333
+ BTS=BTS,
334
+ BK=BK,
335
+ BV=BV,
336
+ num_warps=num_warps,
337
+ num_stages=num_stages,
338
+ )
339
+ ctx.save_for_backward(q, k, v)
340
+ ctx.scale = scale
341
+ return o.sum(0).to(q.dtype), z.sum(0).to(q.dtype)
342
+
343
+ @staticmethod
344
+ @input_guard
345
+ @autocast_custom_bwd
346
+ def backward(ctx, do, dz):
347
+ q, k, v = ctx.saved_tensors
348
+ scale = ctx.scale
349
+ BTL, BTS = 64, 32
350
+ assert BTL % BTS == 0
351
+ BK = min(128, max(triton.next_power_of_2(k.shape[-1]), 16))
352
+ BV = min(128, max(triton.next_power_of_2(v.shape[-1]), 16))
353
+ B, H, T, K, V = *k.shape, v.shape[-1]
354
+ num_stages = 2
355
+ num_warps = 4
356
+ NK = triton.cdiv(K, BK)
357
+ NV = triton.cdiv(V, BV)
358
+ grid = (NK * NV, triton.cdiv(T, BTL), B * H)
359
+
360
+ assert NK == 1, "will encounter some synchronization issue if not"
361
+
362
+ dq = torch.empty(NV, B, H, T, K, dtype=q.dtype, device=q.device)
363
+ dk = torch.empty(NV, B, H, T, K, dtype=q.dtype, device=q.device)
364
+ dv = torch.empty(NK, B, H, T, V, dtype=q.dtype, device=q.device)
365
+
366
+ parallel_based_bwd_kernel[grid](
367
+ q, k, v, do, dz, dq, dk, dv,
368
+ scale,
369
+ B=B,
370
+ H=H,
371
+ T=T,
372
+ K=K,
373
+ V=V,
374
+ BTL=BTL,
375
+ BTS=BTS,
376
+ BK=BK,
377
+ BV=BV,
378
+ num_warps=num_warps,
379
+ num_stages=num_stages,
380
+ )
381
+
382
+ return dq.sum(0).to(q.dtype), dk.sum(0).to(k.dtype), dv.sum(0).to(v.dtype), None
383
+
384
+
385
+ triton_parallel_based = ParallelBasedFunction.apply
386
+
387
+
388
+ def parallel_based(
389
+ q: torch.Tensor,
390
+ k: torch.Tensor,
391
+ v: torch.Tensor,
392
+ scale: float | None = None,
393
+ use_norm: bool = True,
394
+ head_first: bool = False,
395
+ ):
396
+ assert q.shape[-1] <= 128, "only support feature dim up to 128"
397
+ if scale is None:
398
+ scale = q.shape[-1] ** -0.5
399
+ if not head_first:
400
+ q, k, v = map(lambda x: x.transpose(1, 2), (q, k, v))
401
+ o, z = triton_parallel_based(q, k, v, scale)
402
+ if use_norm:
403
+ o = o / (z[..., None] + 1e-6)
404
+ if not head_first:
405
+ o = o.transpose(1, 2)
406
+ return o.to(q.dtype)
code/flash-linear-attention/fla/ops/comba/__init__.py ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ from .chunk import chunk_comba
2
+ from .fused_recurrent import fused_recurrent_comba
3
+
4
+ __all__ = [
5
+ "chunk_comba",
6
+ "fused_recurrent_comba",
7
+ ]
code/flash-linear-attention/fla/ops/comba/chunk.py ADDED
@@ -0,0 +1,340 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+
6
+ from fla.modules.l2norm import l2norm_bwd, l2norm_fwd
7
+ from fla.ops.comba.utils import chunk_comba_cumsum_scalar_bwd, chunk_comba_cumsum_scalar_fwd
8
+ from fla.ops.comba.wy_fast import chunk_scaled_dot_comba_pkt_fwd, prepare_wy_repr_bwd, recompute_w_u_fwd
9
+ from fla.ops.common.chunk_delta_h import chunk_gated_delta_rule_bwd_dhu, chunk_gated_delta_rule_fwd_h
10
+ from fla.ops.common.chunk_o import chunk_bwd_dqkwg, chunk_bwd_dv_local, chunk_fwd_o
11
+ from fla.ops.utils import chunk_local_cumsum, solve_tril
12
+ from fla.utils import autocast_custom_bwd, autocast_custom_fwd, input_guard
13
+
14
+
15
+ def chunk_comba_fwd(
16
+ q: torch.Tensor,
17
+ k: torch.Tensor,
18
+ v: torch.Tensor,
19
+ p: torch.Tensor,
20
+ g: torch.Tensor,
21
+ beta: torch.Tensor,
22
+ scale: float,
23
+ initial_state: torch.Tensor,
24
+ output_final_state: bool,
25
+ cu_seqlens: torch.LongTensor | None = None,
26
+ ):
27
+ g0, g = chunk_comba_cumsum_scalar_fwd(g, chunk_size=64, cu_seqlens=cu_seqlens)
28
+ # obtain WY representation. u is actually the new v.
29
+ A = chunk_scaled_dot_comba_pkt_fwd(
30
+ k=k,
31
+ p=p,
32
+ beta=beta,
33
+ g0=g0,
34
+ g=g,
35
+ cu_seqlens=cu_seqlens,
36
+ output_dtype=torch.float32,
37
+ )
38
+ A = solve_tril(
39
+ A=A,
40
+ cu_seqlens=cu_seqlens,
41
+ output_dtype=k.dtype,
42
+ )
43
+ w, u = recompute_w_u_fwd(
44
+ k=p,
45
+ v=v,
46
+ beta=beta,
47
+ A=A,
48
+ g_cumsum=g0,
49
+ cu_seqlens=cu_seqlens,
50
+ )
51
+ h, v_new, final_state = chunk_gated_delta_rule_fwd_h(
52
+ k=k,
53
+ w=w,
54
+ u=u,
55
+ g=g,
56
+ initial_state=initial_state,
57
+ output_final_state=output_final_state,
58
+ cu_seqlens=cu_seqlens,
59
+ )
60
+ o = chunk_fwd_o(
61
+ q=q,
62
+ k=k,
63
+ v=v_new,
64
+ h=h,
65
+ g=g,
66
+ scale=scale,
67
+ cu_seqlens=cu_seqlens,
68
+ )
69
+ return g0, g, o, A, final_state
70
+
71
+
72
+ def chunk_comba_bwd(
73
+ q: torch.Tensor,
74
+ k: torch.Tensor,
75
+ v: torch.Tensor,
76
+ p: torch.Tensor,
77
+ g0: torch.Tensor,
78
+ g: torch.Tensor,
79
+ beta: torch.Tensor,
80
+ A: torch.Tensor,
81
+ scale: float,
82
+ initial_state: torch.Tensor,
83
+ do: torch.Tensor,
84
+ dht: torch.Tensor,
85
+ cu_seqlens: torch.LongTensor | None = None,
86
+ ):
87
+ w, u = recompute_w_u_fwd(
88
+ k=p,
89
+ v=v,
90
+ beta=beta,
91
+ A=A,
92
+ g_cumsum=g0,
93
+ cu_seqlens=cu_seqlens,
94
+ )
95
+ h, v_new, _ = chunk_gated_delta_rule_fwd_h(
96
+ k=k,
97
+ w=w,
98
+ u=u,
99
+ g=g,
100
+ initial_state=initial_state,
101
+ output_final_state=False,
102
+ cu_seqlens=cu_seqlens,
103
+ )
104
+ dv = chunk_bwd_dv_local(
105
+ q=q,
106
+ k=k,
107
+ g=g,
108
+ do=do,
109
+ scale=scale,
110
+ cu_seqlens=cu_seqlens,
111
+ )
112
+ dh, dh0, dv = chunk_gated_delta_rule_bwd_dhu(
113
+ q=q,
114
+ k=k,
115
+ w=w,
116
+ g=g,
117
+ h0=initial_state,
118
+ dht=dht,
119
+ do=do,
120
+ dv=dv,
121
+ scale=scale,
122
+ cu_seqlens=cu_seqlens,
123
+ )
124
+ dq, dk, dw, dg = chunk_bwd_dqkwg(
125
+ q=q,
126
+ k=k,
127
+ v=v_new,
128
+ w=w,
129
+ g=g,
130
+ h=h,
131
+ dv=dv,
132
+ do=do,
133
+ dh=dh,
134
+ scale=scale,
135
+ cu_seqlens=cu_seqlens,
136
+ )
137
+ dk2, dv, dp, db, dg0, dg2 = prepare_wy_repr_bwd(
138
+ k=k,
139
+ v=v,
140
+ p=p,
141
+ beta=beta,
142
+ g0=g0,
143
+ g=g,
144
+ A=A,
145
+ dw=dw,
146
+ du=dv,
147
+ cu_seqlens=cu_seqlens,
148
+ )
149
+ dk.add_(dk2)
150
+ dg.add_(dg2)
151
+ assert dg.dtype == torch.float32, "dg should be fp32"
152
+ dg = chunk_local_cumsum(dg, chunk_size=64, reverse=True, cu_seqlens=cu_seqlens)
153
+ # dg0 = d(g_cumsum - g)
154
+ dg += chunk_comba_cumsum_scalar_bwd(dg0, chunk_size=64, cu_seqlens=cu_seqlens)
155
+ return dq, dk, dv, dp, db, dg, dh0
156
+
157
+
158
+ class ChunkCombaFunction(torch.autograd.Function):
159
+
160
+ @staticmethod
161
+ @input_guard
162
+ @autocast_custom_fwd
163
+ def forward(
164
+ ctx,
165
+ q: torch.Tensor,
166
+ k: torch.Tensor,
167
+ v: torch.Tensor,
168
+ p: torch.Tensor,
169
+ g: torch.Tensor,
170
+ beta: torch.Tensor,
171
+ scale: float,
172
+ initial_state: torch.Tensor,
173
+ output_final_state: bool,
174
+ use_qk_l2norm_in_kernel: bool = False,
175
+ cu_seqlens: torch.LongTensor | None = None,
176
+ ):
177
+ if use_qk_l2norm_in_kernel:
178
+ q, q_rstd = l2norm_fwd(q)
179
+ k, k_rstd = l2norm_fwd(k)
180
+ p, p_rstd = l2norm_fwd(p)
181
+ else:
182
+ q_rstd, k_rstd, p_rstd = None, None, None
183
+
184
+ g0, g, o, A, final_state = chunk_comba_fwd(
185
+ q=q,
186
+ k=k,
187
+ v=v,
188
+ p=p,
189
+ g=g,
190
+ beta=beta,
191
+ scale=scale,
192
+ initial_state=initial_state,
193
+ output_final_state=output_final_state,
194
+ cu_seqlens=cu_seqlens,
195
+ )
196
+ ctx.save_for_backward(q, q_rstd, k, k_rstd, p, p_rstd, v, g0, g, beta, A, initial_state, cu_seqlens)
197
+ ctx.scale = scale
198
+ ctx.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel
199
+ return o.to(q.dtype), final_state
200
+
201
+ @staticmethod
202
+ @input_guard
203
+ @autocast_custom_bwd
204
+ def backward(
205
+ ctx,
206
+ do: torch.Tensor,
207
+ dht: torch.Tensor,
208
+ ):
209
+ q, q_rstd, k, k_rstd, p, p_rstd, v, g0, g, beta, A, initial_state, cu_seqlens = ctx.saved_tensors
210
+ dq, dk, dv, dp, db, dg, dh0 = chunk_comba_bwd(
211
+ q=q,
212
+ k=k,
213
+ v=v,
214
+ p=p,
215
+ g0=g0,
216
+ g=g,
217
+ beta=beta,
218
+ A=A,
219
+ scale=ctx.scale,
220
+ initial_state=initial_state,
221
+ do=do,
222
+ dht=dht,
223
+ cu_seqlens=cu_seqlens,
224
+ )
225
+ if ctx.use_qk_l2norm_in_kernel:
226
+ dq = l2norm_bwd(q, q_rstd, dq)
227
+ dk = l2norm_bwd(k, k_rstd, dk)
228
+ dp = l2norm_bwd(p, p_rstd, dp)
229
+ return dq.to(q), dk.to(k), dv.to(v), dp.to(p), dg.to(g), db.to(beta), None, dh0, None, None, None
230
+
231
+
232
+ @torch.compiler.disable
233
+ def chunk_comba(
234
+ q: torch.Tensor,
235
+ k: torch.Tensor,
236
+ v: torch.Tensor,
237
+ p: torch.Tensor,
238
+ g: torch.Tensor,
239
+ beta: torch.Tensor = None,
240
+ scale: float = None,
241
+ initial_state: torch.Tensor = None,
242
+ output_final_state: bool = False,
243
+ use_qk_l2norm_in_kernel: bool = False,
244
+ cu_seqlens: torch.LongTensor | None = None,
245
+ ):
246
+ r"""
247
+ Args:
248
+ q (torch.Tensor):
249
+ queries of shape `[B, T, H, K]`.
250
+ k (torch.Tensor):
251
+ keys of shape `[B, T, H, K]`.
252
+ v (torch.Tensor):
253
+ values of shape `[B, T, H, V]`.
254
+ p (torch.Tensor):
255
+ auxiliary keys of shape `[B, T, H, K]`.
256
+ g (torch.Tensor):
257
+ (forget) gating tensor (in log space!) of shape `[B, T, H]`.
258
+ beta (torch.Tensor):
259
+ betas of shape `[B, T, H]`.
260
+ scale (Optional[int]):
261
+ Scale factor for the RetNet attention scores.
262
+ If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
263
+ initial_state (Optional[torch.Tensor]):
264
+ Initial state of shape `[N, H, K, V]` for `N` input sequences.
265
+ For equal-length input sequences, `N` equals the batch size `B`.
266
+ Default: `None`.
267
+ output_final_state (Optional[bool]):
268
+ Whether to output the final state of shape `[N, H, K, V]`. Default: `False`.
269
+ use_qk_l2norm_in_kernel (bool):
270
+ Whether to apply L2norm to the q/k tensor internally. Default: `False`.
271
+ cu_seqlens (torch.LongTensor):
272
+ Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
273
+ consistent with the FlashAttention API.
274
+
275
+ Returns:
276
+ o (torch.Tensor):
277
+ Outputs of shape `[B, T, H, V]`.
278
+ final_state (torch.Tensor):
279
+ Final state of shape `[N, H, K, V]` if `output_final_state=True` else `None`.
280
+
281
+ Examples::
282
+ >>> import torch
283
+ >>> import torch.nn.functional as F
284
+ >>> from einops import rearrange
285
+ >>> from fla.ops.comba import chunk_comba
286
+ # inputs with equal lengths
287
+ >>> B, T, H, K, V = 4, 2048, 4, 512, 512
288
+ >>> q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
289
+ >>> k = F.normalize(torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda'), p=2, dim=-1)
290
+ >>> v = torch.randn(B, T, H, V, dtype=torch.bfloat16, device='cuda')
291
+ >>> b = torch.rand(H, dtype=torch.bfloat16, device='cuda').sigmoid()
292
+ >>> p = k * b[:, None]
293
+ >>> beta = torch.rand(B, T, H, dtype=torch.bfloat16, device='cuda').sigmoid()
294
+ >>> g = F.logsigmoid(torch.rand(B, T, H, dtype=torch.bfloat16, device='cuda'))
295
+ >>> h0 = torch.randn(B, H, K, V, dtype=torch.bfloat16, device='cuda')
296
+ >>> o, ht = chunk_comba(
297
+ q, k, v, p, g, beta,
298
+ initial_state=h0,
299
+ output_final_state=True
300
+ )
301
+ # for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
302
+ >>> q, k, v, beta, g = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, beta, g))
303
+ # for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
304
+ >>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
305
+ >>> o_var, ht_var = chunk_comba(
306
+ q, k, v, p, g, beta,
307
+ initial_state=h0,
308
+ output_final_state=True,
309
+ cu_seqlens=cu_seqlens
310
+ )
311
+ """
312
+ if p is None:
313
+ p = k
314
+ if cu_seqlens is not None:
315
+ if q.shape[0] != 1:
316
+ raise ValueError(
317
+ f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
318
+ f"Please flatten variable-length inputs before processing.",
319
+ )
320
+ if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
321
+ raise ValueError(
322
+ f"The number of initial states is expected to be equal to the number of input sequences, "
323
+ f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.",
324
+ )
325
+ if scale is None:
326
+ scale = k.shape[-1] ** -0.5
327
+ o, final_state = ChunkCombaFunction.apply(
328
+ q,
329
+ k,
330
+ v,
331
+ p,
332
+ g,
333
+ beta,
334
+ scale,
335
+ initial_state,
336
+ output_final_state,
337
+ use_qk_l2norm_in_kernel,
338
+ cu_seqlens,
339
+ )
340
+ return o, final_state
code/flash-linear-attention/fla/ops/comba/fused_recurrent.py ADDED
@@ -0,0 +1,330 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils.op import exp
9
+ from fla.utils import input_guard
10
+
11
+
12
+ @triton.heuristics({
13
+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
14
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
15
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
16
+ })
17
+ @triton.jit(do_not_specialize=['T'])
18
+ def fused_recurrent_comba_fwd_kernel(
19
+ q,
20
+ k,
21
+ p,
22
+ v,
23
+ g,
24
+ beta,
25
+ o,
26
+ h0,
27
+ ht,
28
+ cu_seqlens,
29
+ scale,
30
+ T,
31
+ B: tl.constexpr,
32
+ H: tl.constexpr,
33
+ HV: tl.constexpr,
34
+ K: tl.constexpr,
35
+ V: tl.constexpr,
36
+ BK: tl.constexpr,
37
+ BV: tl.constexpr,
38
+ USE_INITIAL_STATE: tl.constexpr, # whether to use initial state
39
+ STORE_FINAL_STATE: tl.constexpr, # whether to store final state
40
+ IS_BETA_HEADWISE: tl.constexpr, # whether beta is headwise vector or scalar,
41
+ USE_QK_L2NORM_IN_KERNEL: tl.constexpr,
42
+ IS_VARLEN: tl.constexpr,
43
+ ):
44
+ i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
45
+ i_n, i_hv = i_nh // HV, i_nh % HV
46
+ i_h = i_hv // (HV // H)
47
+ if IS_VARLEN:
48
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
49
+ all = T
50
+ T = eos - bos
51
+ else:
52
+ bos, eos = i_n * T, i_n * T + T
53
+ all = B * T
54
+ o_k = i_k * BK + tl.arange(0, BK)
55
+ o_v = i_v * BV + tl.arange(0, BV)
56
+
57
+ p_q = q + (bos * H + i_h) * K + o_k
58
+ p_k = k + (bos * H + i_h) * K + o_k
59
+ p_v = v + (bos * HV + i_hv) * V + o_v
60
+ p_p = p + (bos * H + i_h) * K + o_k
61
+ if IS_BETA_HEADWISE:
62
+ p_beta = beta + (bos * HV + i_hv) * V + o_v
63
+ else:
64
+ p_beta = beta + bos * HV + i_hv
65
+ p_g = g + bos * HV + i_hv
66
+ p_o = o + ((i_k * all + bos) * HV + i_hv) * V + o_v
67
+
68
+ mask_k = o_k < K
69
+ mask_v = o_v < V
70
+ mask_h = mask_k[:, None] & mask_v[None, :]
71
+
72
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
73
+ if USE_INITIAL_STATE:
74
+ p_h0 = h0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
75
+ b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32)
76
+
77
+ for _ in range(0, T):
78
+ b_q = tl.load(p_q, mask=mask_k, other=0).to(tl.float32)
79
+ b_k = tl.load(p_k, mask=mask_k, other=0).to(tl.float32)
80
+ b_v = tl.load(p_v, mask=mask_v, other=0).to(tl.float32)
81
+ b_p = tl.load(p_p, mask=mask_k, other=0).to(tl.float32)
82
+ b_g = tl.load(p_g).to(tl.float32)
83
+
84
+ if USE_QK_L2NORM_IN_KERNEL:
85
+ b_q = b_q / tl.sqrt(tl.sum(b_q * b_q) + 1e-6)
86
+ b_k = b_k / tl.sqrt(tl.sum(b_k * b_k) + 1e-6)
87
+ b_p = b_p / tl.sqrt(tl.sum(b_p * b_p) + 1e-6)
88
+ b_q = b_q * scale
89
+ # [BV]
90
+ b_v -= tl.sum(b_h * b_p[:, None], 0)
91
+ # [BK, BV]
92
+ b_h *= exp(b_g)
93
+ if IS_BETA_HEADWISE:
94
+ b_beta = tl.load(p_beta, mask=mask_v, other=0).to(tl.float32)
95
+ else:
96
+ b_beta = tl.load(p_beta).to(tl.float32)
97
+ b_v *= b_beta
98
+ # [BK, BV]
99
+ b_h += b_k[:, None] * b_v[None, :]
100
+ # [BV]
101
+ b_o = tl.sum(b_h * b_q[:, None], 0)
102
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v)
103
+
104
+ p_q += H*K
105
+ p_k += H*K
106
+ p_o += HV*V
107
+ p_v += HV*V
108
+ p_p += H*K
109
+ p_g += HV
110
+ p_beta += HV * (V if IS_BETA_HEADWISE else 1)
111
+
112
+ if STORE_FINAL_STATE:
113
+ p_ht = ht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
114
+ tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h)
115
+
116
+
117
+ def fused_recurrent_comba_fwd(
118
+ q: torch.Tensor,
119
+ k: torch.Tensor,
120
+ v: torch.Tensor,
121
+ p: torch.Tensor,
122
+ g: torch.Tensor,
123
+ beta: torch.Tensor,
124
+ scale: float,
125
+ initial_state: torch.Tensor,
126
+ output_final_state: bool,
127
+ use_qk_l2norm_in_kernel: bool = False,
128
+ cu_seqlens: torch.LongTensor | None = None,
129
+ ) -> tuple[torch.Tensor, torch.Tensor]:
130
+ B, T, H, K, V = *k.shape, v.shape[-1]
131
+ HV = v.shape[2]
132
+ N = B if cu_seqlens is None else len(cu_seqlens) - 1
133
+ BK, BV = triton.next_power_of_2(K), min(triton.next_power_of_2(V), 8)
134
+ NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
135
+ assert NK == 1, "NK > 1 is not supported yet"
136
+ num_stages = 3
137
+ num_warps = 1
138
+
139
+ o = q.new_empty(NK, *v.shape)
140
+ if output_final_state:
141
+ final_state = q.new_empty(N, HV, K, V, dtype=torch.float32)
142
+ else:
143
+ final_state = None
144
+
145
+ grid = (NK, NV, N * HV)
146
+ fused_recurrent_comba_fwd_kernel[grid](
147
+ q=q,
148
+ k=k,
149
+ p=p,
150
+ v=v,
151
+ g=g,
152
+ beta=beta,
153
+ o=o,
154
+ h0=initial_state,
155
+ ht=final_state,
156
+ cu_seqlens=cu_seqlens,
157
+ scale=scale,
158
+ T=T,
159
+ B=B,
160
+ H=H,
161
+ HV=HV,
162
+ K=K,
163
+ V=V,
164
+ BK=BK,
165
+ BV=BV,
166
+ IS_BETA_HEADWISE=beta.ndim == v.ndim,
167
+ USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
168
+ num_warps=num_warps,
169
+ num_stages=num_stages,
170
+ )
171
+ o = o.squeeze(0)
172
+ return o, final_state
173
+
174
+
175
+ class FusedRecurrentCombaFunction(torch.autograd.Function):
176
+
177
+ @staticmethod
178
+ @input_guard
179
+ def forward(
180
+ ctx,
181
+ q: torch.Tensor,
182
+ k: torch.Tensor,
183
+ p: torch.Tensor,
184
+ v: torch.Tensor,
185
+ g: torch.Tensor,
186
+ beta: torch.Tensor,
187
+ scale: float,
188
+ initial_state: torch.Tensor,
189
+ output_final_state: bool,
190
+ use_qk_l2norm_in_kernel: bool = False,
191
+ cu_seqlens: torch.LongTensor | None = None,
192
+ ):
193
+ o, final_state = fused_recurrent_comba_fwd(
194
+ q=q,
195
+ k=k,
196
+ p=p,
197
+ v=v,
198
+ g=g,
199
+ beta=beta,
200
+ scale=scale,
201
+ initial_state=initial_state,
202
+ output_final_state=output_final_state,
203
+ use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
204
+ cu_seqlens=cu_seqlens,
205
+ )
206
+
207
+ return o, final_state
208
+
209
+ @staticmethod
210
+ @input_guard
211
+ def backward(ctx, do, dht):
212
+ raise NotImplementedError(
213
+ "Backward pass is not implemented yet and we do not have plans to implement it "
214
+ "because we haven't figured out how to compute dg without materializing the full "
215
+ "hidden states for all time steps.",
216
+ )
217
+
218
+
219
+ def fused_recurrent_comba(
220
+ q: torch.Tensor,
221
+ k: torch.Tensor,
222
+ p: torch.Tensor,
223
+ v: torch.Tensor,
224
+ g: torch.Tensor,
225
+ beta: torch.Tensor = None,
226
+ scale: float = None,
227
+ initial_state: torch.Tensor = None,
228
+ output_final_state: bool = False,
229
+ use_qk_l2norm_in_kernel: bool = False,
230
+ cu_seqlens: torch.LongTensor | None = None,
231
+ ) -> tuple[torch.Tensor, torch.Tensor]:
232
+ r"""
233
+ Args:
234
+ q (torch.Tensor):
235
+ queries of shape `[B, T, H, K]`.
236
+ k (torch.Tensor):
237
+ keys of shape `[B, T, H, K]`.
238
+ p (torch.Tensor):
239
+ auxiliary keys of shape `[B, T, H, K]`.
240
+ v (torch.Tensor):
241
+ values of shape `[B, T, HV, V]`.
242
+ GVA is applied if `HV > H`.
243
+ g (torch.Tensor):
244
+ g (decays) of shape `[B, T, HV]`.
245
+ beta (torch.Tensor):
246
+ betas of shape `[B, T, HV]`.
247
+ scale (Optional[int]):
248
+ Scale factor for the RetNet attention scores.
249
+ If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
250
+ initial_state (Optional[torch.Tensor]):
251
+ Initial state of shape `[N, HV, K, V]` for `N` input sequences.
252
+ For equal-length input sequences, `N` equals the batch size `B`.
253
+ Default: `None`.
254
+ output_final_state (Optional[bool]):
255
+ Whether to output the final state of shape `[N, HV, K, V]`. Default: `False`.
256
+ use_qk_l2norm_in_kernel (Optional[bool]):
257
+ Whether to use qk l2norm within the kernel for saving GPU memory.
258
+ Default: `False`.
259
+ cu_seqlens (torch.LongTensor):
260
+ Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
261
+ consistent with the FlashAttention API.
262
+
263
+ Returns:
264
+ o (torch.Tensor):
265
+ Outputs of shape `[B, T, HV, V]`.
266
+ final_state (torch.Tensor):
267
+ Final state of shape `[N, HV, K, V]` if `output_final_state=True` else `None`.
268
+
269
+ Examples::
270
+ >>> import torch
271
+ >>> import torch.nn.functional as F
272
+ >>> from einops import rearrange
273
+ >>> from fla.ops.comba import fused_recurrent_comba
274
+ # inputs with equal lengths
275
+ >>> B, T, H, HV, K, V = 4, 2048, 4, 8, 512, 512
276
+ >>> q = torch.randn(B, T, H, K, device='cuda')
277
+ >>> k = F.normalize(torch.randn(B, T, H, K, device='cuda'), p=2, dim=-1)
278
+ >>> v = torch.randn(B, T, HV, V, device='cuda')
279
+ >>> b = torch.rand(H, dtype=torch.bfloat16, device='cuda').sigmoid()
280
+ >>> p = k * b[:, None]
281
+ >>> g = F.logsigmoid(torch.rand(B, T, HV, device='cuda'))
282
+ >>> beta = torch.rand(B, T, HV, device='cuda').sigmoid()
283
+ >>> h0 = torch.randn(B, HV, K, V, device='cuda')
284
+ >>> o, ht = fused_recurrent_comba(
285
+ q, k, v, p, g, beta,
286
+ initial_state=h0,
287
+ output_final_state=True
288
+ )
289
+ # for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
290
+ >>> q, k, v, p, g, beta = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, p, g, beta))
291
+ # for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
292
+ >>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
293
+ >>> o_var, ht_var = fused_recurrent_comba(
294
+ q, k, p, v, g, beta,
295
+ initial_state=h0,
296
+ output_final_state=True,
297
+ cu_seqlens=cu_seqlens
298
+ )
299
+ """
300
+ if cu_seqlens is not None:
301
+ if q.shape[0] != 1:
302
+ raise ValueError(
303
+ f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
304
+ f"Please flatten variable-length inputs before processing.",
305
+ )
306
+ if initial_state is not None and initial_state.shape[0] != len(cu_seqlens) - 1:
307
+ raise ValueError(
308
+ f"The number of initial states is expected to be equal to the number of input sequences, "
309
+ f"i.e., {len(cu_seqlens) - 1} rather than {initial_state.shape[0]}.",
310
+ )
311
+ if scale is None:
312
+ scale = k.shape[-1] ** -0.5
313
+ if beta is None:
314
+ beta = torch.ones_like(q[..., 0])
315
+ if p is None:
316
+ p = k
317
+ o, final_state = FusedRecurrentCombaFunction.apply(
318
+ q,
319
+ k,
320
+ p,
321
+ v,
322
+ g,
323
+ beta,
324
+ scale,
325
+ initial_state,
326
+ output_final_state,
327
+ use_qk_l2norm_in_kernel,
328
+ cu_seqlens,
329
+ )
330
+ return o, final_state
code/flash-linear-attention/fla/ops/comba/utils.py ADDED
@@ -0,0 +1,174 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ import torch
3
+ import triton
4
+ import triton.language as tl
5
+
6
+ from fla.ops.utils.index import prepare_chunk_indices
7
+ from fla.utils import autotune_cache_kwargs
8
+
9
+
10
+ @triton.heuristics({
11
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
12
+ })
13
+ @triton.autotune(
14
+ configs=[
15
+ triton.Config({}, num_warps=num_warps)
16
+ for num_warps in [1, 2, 4, 8]
17
+ ],
18
+ key=['B', 'H', 'BT', 'IS_VARLEN'],
19
+ **autotune_cache_kwargs,
20
+ )
21
+ @triton.jit(do_not_specialize=['T'])
22
+ def chunk_comba_cumsum_scalar_fwd_kernel(
23
+ g,
24
+ g0,
25
+ g1,
26
+ cu_seqlens,
27
+ chunk_indices,
28
+ T,
29
+ B: tl.constexpr,
30
+ H: tl.constexpr,
31
+ BT: tl.constexpr,
32
+ IS_VARLEN: tl.constexpr,
33
+ HEAD_FIRST: tl.constexpr,
34
+ ):
35
+ i_t, i_bh = tl.program_id(0), tl.program_id(1)
36
+ i_b, i_h = i_bh // H, i_bh % H
37
+ if IS_VARLEN:
38
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
39
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
40
+ T = eos - bos
41
+ else:
42
+ bos, eos = i_b * T, i_b * T + T
43
+
44
+ if HEAD_FIRST:
45
+ p_g = tl.make_block_ptr(g + bos*H + i_h*T, (T,), (1,), (i_t * BT,), (BT,), (0,))
46
+ p_g0 = tl.make_block_ptr(g0 + bos*H + i_h*T, (T,), (1,), (i_t * BT,), (BT,), (0,))
47
+ p_g1 = tl.make_block_ptr(g1 + bos*H + i_h*T, (T,), (1,), (i_t * BT,), (BT,), (0,))
48
+ else:
49
+ p_g = tl.make_block_ptr(g + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
50
+ p_g0 = tl.make_block_ptr(g0 + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
51
+ p_g1 = tl.make_block_ptr(g1 + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
52
+ # [BT]
53
+ b_g = tl.load(p_g, boundary_check=(0,)).to(tl.float32)
54
+ b_g1 = tl.cumsum(b_g, axis=0)
55
+ b_g0 = b_g1 - b_g
56
+ tl.store(p_g0, b_g0.to(p_g0.dtype.element_ty), boundary_check=(0,))
57
+ tl.store(p_g1, b_g1.to(p_g1.dtype.element_ty), boundary_check=(0,))
58
+
59
+
60
+ def chunk_comba_cumsum_scalar_fwd(
61
+ g: torch.Tensor,
62
+ chunk_size: int,
63
+ cu_seqlens: torch.Tensor | None = None,
64
+ head_first: bool = False,
65
+ output_dtype: torch.dtype | None = torch.float,
66
+ ) -> torch.Tensor:
67
+ if head_first:
68
+ B, H, T = g.shape
69
+ else:
70
+ B, T, H = g.shape
71
+ assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
72
+ BT = chunk_size
73
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
74
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
75
+ g0, g1 = torch.empty_like(g, dtype=output_dtype or g.dtype), torch.empty_like(g, dtype=output_dtype or g.dtype)
76
+ grid = (NT, B * H)
77
+ chunk_comba_cumsum_scalar_fwd_kernel[grid](
78
+ g,
79
+ g0,
80
+ g1,
81
+ cu_seqlens,
82
+ chunk_indices,
83
+ T=T,
84
+ B=B,
85
+ H=H,
86
+ BT=BT,
87
+ HEAD_FIRST=head_first,
88
+ )
89
+ return g0, g1
90
+
91
+
92
+ @triton.heuristics({
93
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
94
+ })
95
+ @triton.autotune(
96
+ configs=[
97
+ triton.Config({}, num_warps=num_warps)
98
+ for num_warps in [1, 2, 4, 8]
99
+ ],
100
+ key=['B', 'H', 'BT', 'IS_VARLEN'],
101
+ **autotune_cache_kwargs,
102
+ )
103
+ @triton.jit(do_not_specialize=['T'])
104
+ def chunk_comba_cumsum_scalar_bwd_kernel(
105
+ dg0,
106
+ dgr,
107
+ cu_seqlens,
108
+ chunk_indices,
109
+ T,
110
+ B: tl.constexpr,
111
+ H: tl.constexpr,
112
+ BT: tl.constexpr,
113
+ IS_VARLEN: tl.constexpr,
114
+ HEAD_FIRST: tl.constexpr,
115
+ ):
116
+ i_t, i_bh = tl.program_id(0), tl.program_id(1)
117
+ i_b, i_h = i_bh // H, i_bh % H
118
+ if IS_VARLEN:
119
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
120
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
121
+ T = eos - bos
122
+ else:
123
+ bos, eos = i_b * T, i_b * T + T
124
+
125
+ if HEAD_FIRST:
126
+ p_dg0 = tl.make_block_ptr(dg0 + bos*H + i_h*T, (T,), (1,), (i_t * BT,), (BT,), (0,))
127
+ p_dgr = tl.make_block_ptr(dgr + bos*H + i_h*T, (T,), (1,), (i_t * BT,), (BT,), (0,))
128
+ else:
129
+ p_dg0 = tl.make_block_ptr(dg0 + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
130
+ p_dgr = tl.make_block_ptr(dgr + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
131
+ # [BT]
132
+ """
133
+ b_dg: 1,2,3,4
134
+ b_dg0: 0,1,2,3
135
+ b_temp: 0,1,3,6
136
+ b_dz: 6
137
+ b_dgr: 6,5,3,0
138
+ """
139
+ b_dg0 = tl.load(p_dg0, boundary_check=(0,)).to(tl.float32)
140
+ b_temp = tl.cumsum(b_dg0, axis=0)
141
+ b_dz = tl.sum(b_dg0, axis=0)
142
+ b_dgr = -b_temp + b_dz[None]
143
+ tl.store(p_dgr, b_dgr.to(p_dgr.dtype.element_ty), boundary_check=(0,))
144
+
145
+
146
+ def chunk_comba_cumsum_scalar_bwd(
147
+ dg0: torch.Tensor,
148
+ chunk_size: int,
149
+ cu_seqlens: torch.Tensor | None = None,
150
+ head_first: bool = False,
151
+ output_dtype: torch.dtype | None = torch.float,
152
+ ) -> torch.Tensor:
153
+ if head_first:
154
+ B, H, T = dg0.shape
155
+ else:
156
+ B, T, H = dg0.shape
157
+ assert chunk_size == 2**(chunk_size.bit_length()-1), "chunk_size must be a power of 2"
158
+ BT = chunk_size
159
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
160
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
161
+ dg = torch.empty_like(dg0, dtype=output_dtype or dg0.dtype)
162
+ grid = (NT, B * H)
163
+ chunk_comba_cumsum_scalar_bwd_kernel[grid](
164
+ dg0,
165
+ dg,
166
+ cu_seqlens,
167
+ chunk_indices,
168
+ T=T,
169
+ B=B,
170
+ H=H,
171
+ BT=BT,
172
+ HEAD_FIRST=head_first,
173
+ )
174
+ return dg
code/flash-linear-attention/fla/ops/comba/wy_fast.py ADDED
@@ -0,0 +1,424 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils import prepare_chunk_indices
9
+ from fla.ops.utils.op import exp
10
+ from fla.utils import autotune_cache_kwargs, check_shared_mem
11
+
12
+
13
+ @triton.heuristics({
14
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
15
+ 'USE_G': lambda args: args['g'] is not None,
16
+ })
17
+ @triton.autotune(
18
+ configs=[
19
+ triton.Config({'BK': BK}, num_warps=num_warps, num_stages=num_stages)
20
+ for BK in [32, 64, 128]
21
+ for num_warps in [2, 4, 8]
22
+ for num_stages in [2, 3, 4]
23
+ ],
24
+ key=['H', 'K', 'BT', 'IS_VARLEN', 'USE_G'],
25
+ **autotune_cache_kwargs,
26
+ )
27
+ @triton.jit(do_not_specialize=['T'])
28
+ def chunk_scaled_dot_comba_pkt_fwd_kernel(
29
+ k,
30
+ p,
31
+ beta,
32
+ g0,
33
+ g,
34
+ A,
35
+ cu_seqlens,
36
+ chunk_indices,
37
+ T,
38
+ H: tl.constexpr,
39
+ K: tl.constexpr,
40
+ BT: tl.constexpr,
41
+ BK: tl.constexpr,
42
+ IS_VARLEN: tl.constexpr,
43
+ USE_G: tl.constexpr,
44
+ ):
45
+ i_t, i_bh = tl.program_id(0), tl.program_id(1)
46
+ i_b, i_h = i_bh // H, i_bh % H
47
+ if IS_VARLEN:
48
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
49
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
50
+ T = eos - bos
51
+ else:
52
+ bos, eos = i_b * T, i_b * T + T
53
+ o_t = i_t * BT + tl.arange(0, BT)
54
+ m_t = o_t < T
55
+
56
+ p_beta = tl.make_block_ptr(beta + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
57
+ b_beta = tl.load(p_beta, boundary_check=(0,))
58
+
59
+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
60
+ for i_k in range(tl.cdiv(K, BK)):
61
+ p_k = tl.make_block_ptr(k + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
62
+ p_p = tl.make_block_ptr(p + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
63
+ b_k = tl.load(p_k, boundary_check=(0, 1))
64
+ b_p = tl.load(p_p, boundary_check=(0, 1))
65
+ b_pb = b_p * b_beta[:, None]
66
+ b_A += tl.dot(b_pb.to(b_k.dtype), tl.trans(b_k))
67
+
68
+ if USE_G:
69
+ p_g0 = tl.make_block_ptr(g0 + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
70
+ p_g = tl.make_block_ptr(g + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
71
+ b_g0 = tl.load(p_g0, boundary_check=(0,))
72
+ b_g = tl.load(p_g, boundary_check=(0,))
73
+ b_A = b_A * exp(b_g0[:, None] - b_g[None, :])
74
+
75
+ m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
76
+ b_A = tl.where(m_A, b_A, 0)
77
+ p_A = tl.make_block_ptr(A + (bos*H + i_h) * BT, (T, BT), (BT*H, 1), (i_t * BT, 0), (BT, BT), (1, 0))
78
+ tl.store(p_A, b_A.to(p_A.dtype.element_ty), boundary_check=(0, 1))
79
+
80
+
81
+ def chunk_scaled_dot_comba_pkt_fwd(
82
+ k: torch.Tensor,
83
+ p: torch.Tensor,
84
+ beta: torch.Tensor,
85
+ g0: torch.Tensor | None = None,
86
+ g: torch.Tensor | None = None,
87
+ cu_seqlens: torch.LongTensor | None = None,
88
+ chunk_size: int = 64,
89
+ output_dtype: torch.dtype = torch.float32,
90
+ ) -> torch.Tensor:
91
+ r"""
92
+ Compute beta \mathcal{A}(i-1/j) * P * K^T.
93
+
94
+ Args:
95
+ k (torch.Tensor):
96
+ The key tensor of shape `[B, T, H, K]`.
97
+ p (torch.Tensor):
98
+ The auxiliary key tensor of shape `[B, T, H, K]`.
99
+ beta (torch.Tensor):
100
+ The beta tensor of shape `[B, T, H]`.
101
+ g0 (torch.Tensor):
102
+ The cumulative sum minus the original one of the gate tensor of shape `[B, T, H]`.
103
+ Default: None
104
+ g (torch.Tensor):
105
+ The cumulative sum of the gate tensor of shape `[B, T, H]`.
106
+ Default: None
107
+ cu_seqlens (torch.LongTensor):
108
+ The cumulative sequence lengths of the input tensor.
109
+ Default: None
110
+ chunk_size (int):
111
+ The chunk size. Default: 64.
112
+ output_dtype (torch.dtype):
113
+ The dtype of the output tensor. Default: `torch.float32`
114
+
115
+ Returns:
116
+ beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size.
117
+ """
118
+ B, T, H, K = k.shape
119
+ BT = chunk_size
120
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
121
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
122
+ A = torch.empty(B, T, H, BT, device=k.device, dtype=output_dtype)
123
+ chunk_scaled_dot_comba_pkt_fwd_kernel[(NT, B * H)](
124
+ k=k,
125
+ p=p,
126
+ beta=beta,
127
+ g0=g0,
128
+ g=g,
129
+ A=A,
130
+ cu_seqlens=cu_seqlens,
131
+ chunk_indices=chunk_indices,
132
+ T=T,
133
+ H=H,
134
+ K=K,
135
+ BT=BT,
136
+ )
137
+ return A
138
+
139
+
140
+ @triton.heuristics({
141
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
142
+ })
143
+ @triton.autotune(
144
+ configs=[
145
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
146
+ for num_warps in [2, 4]
147
+ for num_stages in [2, 3, 4]
148
+ ],
149
+ key=['H', 'K', 'V', 'BT', 'BK', 'BV', 'IS_VARLEN'],
150
+ **autotune_cache_kwargs,
151
+ )
152
+ @triton.jit(do_not_specialize=['T'])
153
+ def prepare_wy_repr_bwd_kernel(
154
+ k,
155
+ v,
156
+ p,
157
+ beta,
158
+ g0,
159
+ g,
160
+ A,
161
+ dw,
162
+ du,
163
+ dk,
164
+ dv,
165
+ dp,
166
+ dbeta,
167
+ dg0,
168
+ dg,
169
+ cu_seqlens,
170
+ chunk_indices,
171
+ T,
172
+ H: tl.constexpr,
173
+ K: tl.constexpr,
174
+ V: tl.constexpr,
175
+ BT: tl.constexpr,
176
+ BK: tl.constexpr,
177
+ BV: tl.constexpr,
178
+ IS_VARLEN: tl.constexpr,
179
+ ):
180
+ i_t, i_bh = tl.program_id(0), tl.program_id(1)
181
+ i_b, i_h = i_bh // H, i_bh % H
182
+ if IS_VARLEN:
183
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
184
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
185
+ T = eos - bos
186
+ else:
187
+ bos, eos = i_b * T, i_b * T + T
188
+
189
+ p_beta = tl.make_block_ptr(beta + (bos*H + i_h), (T,), (H,), (i_t * BT,), (BT,), (0,))
190
+ p_g0 = tl.make_block_ptr(g0 + (bos*H + i_h), (T,), (H,), (i_t * BT,), (BT,), (0,))
191
+ p_g = tl.make_block_ptr(g + (bos*H + i_h), (T,), (H,), (i_t * BT,), (BT,), (0,))
192
+ p_A = tl.make_block_ptr(A + (bos*H + i_h) * BT, (BT, T), (1, H*BT), (0, i_t * BT), (BT, BT), (0, 1))
193
+
194
+ b_A = tl.load(p_A, boundary_check=(0, 1))
195
+ b_beta = tl.load(p_beta, boundary_check=(0,))
196
+ b_g0 = tl.load(p_g0, boundary_check=(0,))
197
+ b_g0_exp = tl.exp(b_g0)
198
+ b_g = tl.load(p_g, boundary_check=(0,))
199
+
200
+ b_dbeta = tl.zeros([BT], dtype=tl.float32)
201
+ b_dA = tl.zeros([BT, BT], dtype=tl.float32)
202
+ b_dg0 = tl.zeros([BT], dtype=tl.float32)
203
+
204
+ for i_k in range(tl.cdiv(K, BK)):
205
+ p_p = tl.make_block_ptr(p + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
206
+ p_dp = tl.make_block_ptr(dp + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
207
+ p_dw = tl.make_block_ptr(dw + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
208
+ b_p = tl.load(p_p, boundary_check=(0, 1))
209
+ b_p_beta_g0 = (b_p * b_beta[:, None] * b_g0_exp[:, None]).to(b_p.dtype)
210
+ b_dw = tl.load(p_dw, boundary_check=(0, 1))
211
+ b_dA += tl.dot(b_dw, tl.trans(b_p_beta_g0))
212
+ b_dp_beta_g0 = tl.dot(b_A, b_dw)
213
+ b_dp = b_dp_beta_g0 * b_beta[:, None] * b_g0_exp[:, None]
214
+ b_dbeta += tl.sum(b_dp_beta_g0 * b_p * b_g0_exp[:, None], 1)
215
+ b_dg0 += tl.sum(b_dp * b_p, 1)
216
+ tl.store(p_dp, b_dp.to(p_dp.dtype.element_ty), boundary_check=(0, 1))
217
+
218
+ for i_v in range(tl.cdiv(V, BV)):
219
+ p_v = tl.make_block_ptr(v + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
220
+ p_dv = tl.make_block_ptr(dv + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
221
+ p_du = tl.make_block_ptr(du + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
222
+ b_v = tl.load(p_v, boundary_check=(0, 1))
223
+ b_v_beta = (b_v * b_beta[:, None]).to(b_v.dtype)
224
+ b_du = tl.load(p_du, boundary_check=(0, 1))
225
+ b_dA += tl.dot(b_du, tl.trans(b_v_beta))
226
+ b_dv_beta = tl.dot(b_A, b_du)
227
+ b_dv = b_dv_beta * b_beta[:, None]
228
+ b_dbeta += tl.sum(b_dv_beta * b_v, 1)
229
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
230
+
231
+ o_t = i_t * BT + tl.arange(0, BT)
232
+ m_t = o_t < T
233
+ m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
234
+ b_dA = tl.where(m_A, b_dA, 0)
235
+ b_dA = tl.dot(b_dA.to(b_A.dtype), b_A)
236
+ b_dA = tl.dot(b_A, b_dA.to(b_A.dtype))
237
+ b_dA = tl.where(m_A, -b_dA * exp(b_g0[:, None] - b_g[None, :]), 0).to(k.dtype.element_ty)
238
+ b_dA = b_dA.to(k.dtype.element_ty)
239
+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
240
+
241
+ for i_k in range(tl.cdiv(K, BK)):
242
+ p_k = tl.make_block_ptr(k + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
243
+ p_p = tl.make_block_ptr(p + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
244
+ p_dk = tl.make_block_ptr(dk + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
245
+ p_dp = tl.make_block_ptr(dp + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
246
+ b_k = tl.load(p_k, boundary_check=(0, 1))
247
+ b_p = tl.load(p_p, boundary_check=(0, 1))
248
+ b_dp = tl.load(p_dp, boundary_check=(0, 1))
249
+ b_p_beta = (b_p * b_beta[:, None]).to(b_p.dtype)
250
+ b_A += tl.dot(b_p_beta, tl.trans(b_k))
251
+ b_dp_beta = tl.dot(b_dA, b_k)
252
+ b_dbeta += tl.sum(b_dp_beta * b_p, 1)
253
+ b_dk = tl.dot(tl.trans(b_dA), b_p_beta)
254
+ b_dp += b_dp_beta * b_beta[:, None]
255
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
256
+ tl.store(p_dp, b_dp.to(p_dp.dtype.element_ty), boundary_check=(0, 1))
257
+
258
+ b_dA_A = b_dA * b_A
259
+ b_dg0 += tl.sum(b_dA_A, axis=1)
260
+ b_dg = - tl.sum(b_dA_A, axis=0)
261
+ p_dg = tl.make_block_ptr(dg + (bos*H + i_h), (T,), (H,), (i_t * BT,), (BT,), (0,))
262
+ p_dg0 = tl.make_block_ptr(dg0 + (bos*H + i_h), (T,), (H,), (i_t * BT,), (BT,), (0,))
263
+ p_dbeta = tl.make_block_ptr(dbeta + (bos*H + i_h), (T,), (H,), (i_t * BT,), (BT,), (0,))
264
+ tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), boundary_check=(0,))
265
+ tl.store(p_dg0, b_dg0.to(p_dg0.dtype.element_ty), boundary_check=(0,))
266
+ tl.store(p_dbeta, b_dbeta.to(p_dbeta.dtype.element_ty), boundary_check=(0,))
267
+
268
+
269
+ @triton.heuristics({
270
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
271
+ })
272
+ @triton.autotune(
273
+ configs=[
274
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
275
+ for num_warps in [2, 4, 8]
276
+ for num_stages in [2, 3, 4]
277
+ ],
278
+ key=['H', 'K', 'V', 'BT', 'BK', 'BV', 'IS_VARLEN'],
279
+ **autotune_cache_kwargs,
280
+ )
281
+ @triton.jit(do_not_specialize=['T'])
282
+ def recompute_w_u_fwd_kernel(
283
+ k,
284
+ v,
285
+ beta,
286
+ w,
287
+ u,
288
+ A,
289
+ g,
290
+ cu_seqlens,
291
+ chunk_indices,
292
+ T,
293
+ H: tl.constexpr,
294
+ K: tl.constexpr,
295
+ V: tl.constexpr,
296
+ BT: tl.constexpr,
297
+ BK: tl.constexpr,
298
+ BV: tl.constexpr,
299
+ IS_VARLEN: tl.constexpr,
300
+ ):
301
+ i_t, i_bh = tl.program_id(0), tl.program_id(1)
302
+ i_b, i_h = i_bh // H, i_bh % H
303
+ if IS_VARLEN:
304
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
305
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
306
+ T = eos - bos
307
+ else:
308
+ bos, eos = i_b * T, i_b * T + T
309
+ p_beta = tl.make_block_ptr(beta + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
310
+ p_g = tl.make_block_ptr(g + (bos*H + i_h), (T,), (H,), (i_t * BT,), (BT,), (0,))
311
+ p_A = tl.make_block_ptr(A + (bos*H + i_h) * BT, (T, BT), (H*BT, 1), (i_t * BT, 0), (BT, BT), (1, 0))
312
+ b_beta = tl.load(p_beta, boundary_check=(0,))
313
+ b_A = tl.load(p_A, boundary_check=(0, 1))
314
+ b_g = tl.exp(tl.load(p_g, boundary_check=(0,)))
315
+
316
+ for i_v in range(tl.cdiv(V, BV)):
317
+ p_v = tl.make_block_ptr(v + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
318
+ p_u = tl.make_block_ptr(u + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
319
+ b_v = tl.load(p_v, boundary_check=(0, 1))
320
+ b_vb = (b_v * b_beta[:, None]).to(b_v.dtype)
321
+ b_u = tl.dot(b_A, b_vb, allow_tf32=False)
322
+ tl.store(p_u, b_u.to(p_u.dtype.element_ty), boundary_check=(0, 1))
323
+
324
+ for i_k in range(tl.cdiv(K, BK)):
325
+ p_k = tl.make_block_ptr(k + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
326
+ p_w = tl.make_block_ptr(w + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
327
+ b_k = tl.load(p_k, boundary_check=(0, 1))
328
+ b_kb = (b_k * b_beta[:, None] * b_g[:, None]).to(b_k.dtype)
329
+ b_w = tl.dot(b_A, b_kb)
330
+ tl.store(p_w, b_w.to(p_w.dtype.element_ty), boundary_check=(0, 1))
331
+
332
+
333
+ def recompute_w_u_fwd(
334
+ k: torch.Tensor,
335
+ v: torch.Tensor,
336
+ beta: torch.Tensor,
337
+ g_cumsum: torch.Tensor,
338
+ A: torch.Tensor,
339
+ cu_seqlens: torch.LongTensor | None,
340
+ ) -> tuple[torch.Tensor, torch.Tensor]:
341
+ B, T, H, K, V = *k.shape, v.shape[-1]
342
+ BT = A.shape[-1]
343
+
344
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
345
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
346
+ BK = 64
347
+ BV = 64
348
+
349
+ u = torch.empty_like(v)
350
+ w = torch.empty_like(k)
351
+ recompute_w_u_fwd_kernel[(NT, B*H)](
352
+ k=k,
353
+ v=v,
354
+ beta=beta,
355
+ w=w,
356
+ u=u,
357
+ A=A,
358
+ g=g_cumsum,
359
+ cu_seqlens=cu_seqlens,
360
+ chunk_indices=chunk_indices,
361
+ T=T,
362
+ H=H,
363
+ K=K,
364
+ V=V,
365
+ BT=BT,
366
+ BK=BK,
367
+ BV=BV,
368
+ )
369
+ return w, u
370
+
371
+
372
+ def prepare_wy_repr_bwd(
373
+ k: torch.Tensor,
374
+ v: torch.Tensor,
375
+ p: torch.Tensor,
376
+ g0: torch.Tensor,
377
+ g: torch.Tensor,
378
+ beta: torch.Tensor,
379
+ A: torch.Tensor,
380
+ dw: torch.Tensor,
381
+ du: torch.Tensor,
382
+ cu_seqlens: torch.LongTensor | None,
383
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
384
+ B, T, H, K, V = *k.shape, v.shape[-1]
385
+ BT = 64
386
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
387
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
388
+ CONST_TILING = 64 if check_shared_mem() else 32
389
+ BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
390
+ BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
391
+
392
+ dk = torch.empty_like(k)
393
+ dv = torch.empty_like(v)
394
+ dp = torch.empty_like(p)
395
+ dbeta = torch.empty_like(beta)
396
+ dg0 = torch.empty_like(g0)
397
+ dg = torch.empty_like(g)
398
+ prepare_wy_repr_bwd_kernel[(NT, B * H)](
399
+ k=k,
400
+ v=v,
401
+ p=p,
402
+ beta=beta,
403
+ g0=g0,
404
+ g=g,
405
+ A=A,
406
+ dw=dw,
407
+ du=du,
408
+ dk=dk,
409
+ dv=dv,
410
+ dp=dp,
411
+ dbeta=dbeta,
412
+ dg0=dg0,
413
+ dg=dg,
414
+ cu_seqlens=cu_seqlens,
415
+ chunk_indices=chunk_indices,
416
+ T=T,
417
+ H=H,
418
+ K=K,
419
+ V=V,
420
+ BT=BT,
421
+ BK=BK,
422
+ BV=BV,
423
+ )
424
+ return dk, dv, dp, dbeta, dg0, dg
code/flash-linear-attention/fla/ops/common/__init__.py ADDED
File without changes
code/flash-linear-attention/fla/ops/common/chunk_delta_h.py ADDED
@@ -0,0 +1,533 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils import prepare_chunk_indices, prepare_chunk_offsets
9
+ from fla.ops.utils.op import exp
10
+ from fla.utils import autotune_cache_kwargs, check_shared_mem, is_nvidia_hopper, use_cuda_graph
11
+
12
+ NUM_WARPS = [2, 4] if is_nvidia_hopper else [2, 4, 8, 16]
13
+
14
+
15
+ @triton.heuristics({
16
+ 'USE_G': lambda args: args['g'] is not None,
17
+ 'USE_GK': lambda args: args['gk'] is not None,
18
+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
19
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
20
+ 'SAVE_NEW_VALUE': lambda args: args['v_new'] is not None,
21
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
22
+ })
23
+ @triton.autotune(
24
+ configs=[
25
+ triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
26
+ for num_warps in [2, 4]
27
+ for num_stages in [2, 3, 4]
28
+ for BV in [32, 64]
29
+ ],
30
+ key=['H', 'K', 'V', 'BT'],
31
+ use_cuda_graph=use_cuda_graph,
32
+ **autotune_cache_kwargs,
33
+ )
34
+ @triton.jit(do_not_specialize=['T'])
35
+ def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
36
+ k,
37
+ v,
38
+ w,
39
+ v_new,
40
+ g,
41
+ gk,
42
+ h,
43
+ h0,
44
+ ht,
45
+ cu_seqlens,
46
+ chunk_offsets,
47
+ T,
48
+ H: tl.constexpr,
49
+ K: tl.constexpr,
50
+ V: tl.constexpr,
51
+ BT: tl.constexpr,
52
+ BV: tl.constexpr,
53
+ USE_G: tl.constexpr,
54
+ USE_GK: tl.constexpr,
55
+ USE_INITIAL_STATE: tl.constexpr,
56
+ STORE_FINAL_STATE: tl.constexpr,
57
+ SAVE_NEW_VALUE: tl.constexpr,
58
+ IS_VARLEN: tl.constexpr,
59
+ ):
60
+ i_v, i_nh = tl.program_id(0), tl.program_id(1)
61
+ i_n, i_h = i_nh // H, i_nh % H
62
+ if IS_VARLEN:
63
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
64
+ T = eos - bos
65
+ NT = tl.cdiv(T, BT)
66
+ boh = tl.load(chunk_offsets + i_n).to(tl.int32)
67
+ else:
68
+ bos, eos = i_n * T, i_n * T + T
69
+ NT = tl.cdiv(T, BT)
70
+ boh = i_n * NT
71
+
72
+ # [BK, BV]
73
+ b_h1 = tl.zeros([64, BV], dtype=tl.float32)
74
+ if K > 64:
75
+ b_h2 = tl.zeros([64, BV], dtype=tl.float32)
76
+ if K > 128:
77
+ b_h3 = tl.zeros([64, BV], dtype=tl.float32)
78
+ if K > 192:
79
+ b_h4 = tl.zeros([64, BV], dtype=tl.float32)
80
+
81
+ # calculate offset
82
+ h += ((boh * H + i_h) * K*V).to(tl.int64)
83
+ v += ((bos * H + i_h) * V).to(tl.int64)
84
+ k += ((bos * H + i_h) * K).to(tl.int64)
85
+ w += ((bos * H + i_h) * K).to(tl.int64)
86
+ if SAVE_NEW_VALUE:
87
+ v_new += ((bos * H + i_h) * V).to(tl.int64)
88
+ stride_v = H*V
89
+ stride_h = H*K*V
90
+ stride_k = H*K
91
+ if USE_INITIAL_STATE:
92
+ h0 = h0 + i_nh * K*V
93
+ if STORE_FINAL_STATE:
94
+ ht = ht + i_nh * K*V
95
+
96
+ # load initial state
97
+ if USE_INITIAL_STATE:
98
+ p_h0_1 = tl.make_block_ptr(h0, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
99
+ b_h1 += tl.load(p_h0_1, boundary_check=(0, 1)).to(tl.float32)
100
+ if K > 64:
101
+ p_h0_2 = tl.make_block_ptr(h0, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
102
+ b_h2 += tl.load(p_h0_2, boundary_check=(0, 1)).to(tl.float32)
103
+ if K > 128:
104
+ p_h0_3 = tl.make_block_ptr(h0, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
105
+ b_h3 += tl.load(p_h0_3, boundary_check=(0, 1)).to(tl.float32)
106
+ if K > 192:
107
+ p_h0_4 = tl.make_block_ptr(h0, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
108
+ b_h4 += tl.load(p_h0_4, boundary_check=(0, 1)).to(tl.float32)
109
+
110
+ # main recurrence
111
+ for i_t in range(NT):
112
+ p_h1 = tl.make_block_ptr(h + i_t * stride_h, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
113
+ tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), boundary_check=(0, 1))
114
+ if K > 64:
115
+ p_h2 = tl.make_block_ptr(h + i_t * stride_h, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
116
+ tl.store(p_h2, b_h2.to(p_h2.dtype.element_ty), boundary_check=(0, 1))
117
+ if K > 128:
118
+ p_h3 = tl.make_block_ptr(h + i_t * stride_h, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
119
+ tl.store(p_h3, b_h3.to(p_h3.dtype.element_ty), boundary_check=(0, 1))
120
+ if K > 192:
121
+ p_h4 = tl.make_block_ptr(h + i_t * stride_h, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
122
+ tl.store(p_h4, b_h4.to(p_h4.dtype.element_ty), boundary_check=(0, 1))
123
+
124
+ p_w = tl.make_block_ptr(w, (T, K), (stride_k, 1), (i_t * BT, 0), (BT, 64), (1, 0))
125
+ b_w = tl.load(p_w, boundary_check=(0, 1))
126
+ b_v = tl.dot(b_w, b_h1.to(b_w.dtype))
127
+ if K > 64:
128
+ p_w = tl.make_block_ptr(w, (T, K), (stride_k, 1), (i_t * BT, 64), (BT, 64), (1, 0))
129
+ b_w = tl.load(p_w, boundary_check=(0, 1))
130
+ b_v += tl.dot(b_w, b_h2.to(b_w.dtype))
131
+ if K > 128:
132
+ p_w = tl.make_block_ptr(w, (T, K), (stride_k, 1), (i_t * BT, 128), (BT, 64), (1, 0))
133
+ b_w = tl.load(p_w, boundary_check=(0, 1))
134
+ b_v += tl.dot(b_w, b_h3.to(b_w.dtype))
135
+ if K > 192:
136
+ p_w = tl.make_block_ptr(w, (T, K), (stride_k, 1), (i_t * BT, 192), (BT, 64), (1, 0))
137
+ b_w = tl.load(p_w, boundary_check=(0, 1))
138
+ b_v += tl.dot(b_w, b_h4.to(b_w.dtype))
139
+ p_v = tl.make_block_ptr(v, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
140
+ b_v = tl.load(p_v, boundary_check=(0, 1)) - b_v
141
+
142
+ if SAVE_NEW_VALUE:
143
+ p_v = tl.make_block_ptr(v_new, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
144
+ tl.store(p_v, b_v.to(p_v.dtype.element_ty), boundary_check=(0, 1))
145
+
146
+ last_idx = min((i_t + 1) * BT, T) - 1
147
+ if USE_G:
148
+ m_t = (i_t * BT + tl.arange(0, BT)) < T
149
+ b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
150
+ p_g = tl.make_block_ptr(g + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
151
+ b_g = tl.load(p_g, boundary_check=(0,))
152
+ b_v = b_v * tl.where(m_t, exp(b_g_last - b_g), 0)[:, None]
153
+ b_g_last = exp(b_g_last)
154
+ b_h1 *= b_g_last
155
+ if K > 64:
156
+ b_h2 *= b_g_last
157
+ if K > 128:
158
+ b_h3 *= b_g_last
159
+ if K > 192:
160
+ b_h4 *= b_g_last
161
+
162
+ if USE_GK:
163
+ o_k1 = tl.arange(0, 64)
164
+ b_gk_last1 = tl.load(gk + (bos + last_idx) * H*K + i_h * K + o_k1, mask=(o_k1 < K), other=0.)
165
+ b_h1 *= exp(b_gk_last1)[:, None]
166
+ if K > 64:
167
+ o_k2 = 64 + o_k1
168
+ b_gk_last2 = tl.load(gk + (bos + last_idx) * H*K + i_h * K + o_k2, mask=(o_k2 < K), other=0.)
169
+ b_h2 *= exp(b_gk_last2)[:, None]
170
+ if K > 128:
171
+ o_k3 = 128 + o_k1
172
+ b_gk_last3 = tl.load(gk + (bos + last_idx) * H*K + i_h * K + o_k3, mask=(o_k3 < K), other=0.)
173
+ b_h3 *= exp(b_gk_last3)[:, None]
174
+ if K > 192:
175
+ o_k4 = 192 + o_k1
176
+ b_gk_last4 = tl.load(gk + (bos + last_idx) * H*K + i_h * K + o_k4, mask=(o_k4 < K), other=0.)
177
+ b_h4 *= exp(b_gk_last4)[:, None]
178
+ b_v = b_v.to(k.dtype.element_ty)
179
+
180
+ p_k = tl.make_block_ptr(k, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1))
181
+ b_k = tl.load(p_k, boundary_check=(0, 1))
182
+ b_h1 += tl.dot(b_k, b_v)
183
+ if K > 64:
184
+ p_k = tl.make_block_ptr(k, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1))
185
+ b_k = tl.load(p_k, boundary_check=(0, 1))
186
+ b_h2 += tl.dot(b_k, b_v)
187
+ if K > 128:
188
+ p_k = tl.make_block_ptr(k, (K, T), (1, stride_k), (128, i_t * BT), (64, BT), (0, 1))
189
+ b_k = tl.load(p_k, boundary_check=(0, 1))
190
+ b_h3 += tl.dot(b_k, b_v)
191
+ if K > 192:
192
+ p_k = tl.make_block_ptr(k, (K, T), (1, stride_k), (192, i_t * BT), (64, BT), (0, 1))
193
+ b_k = tl.load(p_k, boundary_check=(0, 1))
194
+ b_h4 += tl.dot(b_k, b_v)
195
+ # epilogue
196
+ if STORE_FINAL_STATE:
197
+ p_ht = tl.make_block_ptr(ht, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
198
+ tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
199
+ if K > 64:
200
+ p_ht = tl.make_block_ptr(ht, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
201
+ tl.store(p_ht, b_h2.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
202
+ if K > 128:
203
+ p_ht = tl.make_block_ptr(ht, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
204
+ tl.store(p_ht, b_h3.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
205
+ if K > 192:
206
+ p_ht = tl.make_block_ptr(ht, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
207
+ tl.store(p_ht, b_h4.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
208
+
209
+
210
+ @triton.heuristics({
211
+ 'USE_G': lambda args: args['g'] is not None,
212
+ 'USE_GK': lambda args: args['gk'] is not None,
213
+ 'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
214
+ 'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
215
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
216
+ })
217
+ @triton.autotune(
218
+ configs=[
219
+ triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
220
+ for num_warps in [2, 4]
221
+ for num_stages in ([4, 3, 2] if check_shared_mem('ampere') else [1])
222
+ for BV in [64, 32]
223
+ ],
224
+ key=['H', 'K', 'V', 'BT', 'BV', 'USE_G'],
225
+ use_cuda_graph=use_cuda_graph,
226
+ **autotune_cache_kwargs,
227
+ )
228
+ @triton.jit(do_not_specialize=['T'])
229
+ def chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64(
230
+ q,
231
+ k,
232
+ w,
233
+ g,
234
+ gk,
235
+ dht,
236
+ dh0,
237
+ do,
238
+ dh,
239
+ dv,
240
+ dv2,
241
+ cu_seqlens,
242
+ chunk_offsets,
243
+ scale,
244
+ T,
245
+ H: tl.constexpr,
246
+ K: tl.constexpr,
247
+ V: tl.constexpr,
248
+ BT: tl.constexpr,
249
+ BV: tl.constexpr,
250
+ USE_G: tl.constexpr,
251
+ USE_GK: tl.constexpr,
252
+ USE_INITIAL_STATE: tl.constexpr,
253
+ USE_FINAL_STATE_GRADIENT: tl.constexpr,
254
+ IS_VARLEN: tl.constexpr,
255
+ ):
256
+ i_v, i_nh = tl.program_id(0), tl.program_id(1)
257
+ i_n, i_h = i_nh // H, i_nh % H
258
+ if IS_VARLEN:
259
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
260
+ T = eos - bos
261
+ NT = tl.cdiv(T, BT)
262
+ boh = tl.load(chunk_offsets + i_n).to(tl.int32)
263
+ else:
264
+ bos, eos = i_n * T, i_n * T + T
265
+ NT = tl.cdiv(T, BT)
266
+ boh = i_n * NT
267
+
268
+ # [BK, BV]
269
+ b_dh1 = tl.zeros([64, BV], dtype=tl.float32)
270
+ if K > 64:
271
+ b_dh2 = tl.zeros([64, BV], dtype=tl.float32)
272
+ if K > 128:
273
+ b_dh3 = tl.zeros([64, BV], dtype=tl.float32)
274
+ if K > 192:
275
+ b_dh4 = tl.zeros([64, BV], dtype=tl.float32)
276
+
277
+ # calculate offset
278
+ q += ((bos * H + i_h) * K).to(tl.int64)
279
+ k += ((bos * H + i_h) * K).to(tl.int64)
280
+ w += ((bos * H + i_h) * K).to(tl.int64)
281
+ do += ((bos * H + i_h) * V).to(tl.int64)
282
+ dv += ((bos * H + i_h) * V).to(tl.int64)
283
+ dv2 += ((bos * H + i_h) * V).to(tl.int64)
284
+ dh += ((boh * H + i_h) * K*V).to(tl.int64)
285
+ if USE_GK:
286
+ gk += ((bos * H + i_h) * K).to(tl.int64)
287
+
288
+ stride_v = H*V
289
+ stride_h = H*K*V
290
+ stride_k = H*K
291
+ if USE_INITIAL_STATE:
292
+ dh0 += i_nh * K*V
293
+ if USE_FINAL_STATE_GRADIENT:
294
+ dht += i_nh * K*V
295
+
296
+ if USE_FINAL_STATE_GRADIENT:
297
+ p_dht1 = tl.make_block_ptr(dht, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
298
+ b_dh1 += tl.load(p_dht1, boundary_check=(0, 1))
299
+ if K > 64:
300
+ p_dht2 = tl.make_block_ptr(dht, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
301
+ b_dh2 += tl.load(p_dht2, boundary_check=(0, 1))
302
+ if K > 128:
303
+ p_dht3 = tl.make_block_ptr(dht, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
304
+ b_dh3 += tl.load(p_dht3, boundary_check=(0, 1))
305
+ if K > 192:
306
+ p_dht4 = tl.make_block_ptr(dht, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
307
+ b_dh4 += tl.load(p_dht4, boundary_check=(0, 1))
308
+
309
+ for i_t in range(NT - 1, -1, -1):
310
+ p_dh1 = tl.make_block_ptr(dh + i_t*stride_h, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
311
+ tl.store(p_dh1, b_dh1.to(p_dh1.dtype.element_ty), boundary_check=(0, 1))
312
+ if K > 64:
313
+ p_dh2 = tl.make_block_ptr(dh + i_t*stride_h, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
314
+ tl.store(p_dh2, b_dh2.to(p_dh2.dtype.element_ty), boundary_check=(0, 1))
315
+ if K > 128:
316
+ p_dh3 = tl.make_block_ptr(dh + i_t*stride_h, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
317
+ tl.store(p_dh3, b_dh3.to(p_dh3.dtype.element_ty), boundary_check=(0, 1))
318
+ if K > 192:
319
+ p_dh4 = tl.make_block_ptr(dh + i_t*stride_h, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
320
+ tl.store(p_dh4, b_dh4.to(p_dh4.dtype.element_ty), boundary_check=(0, 1))
321
+
322
+ last_idx = min((i_t + 1) * BT, T) - 1
323
+ if USE_G:
324
+ bg_last = tl.load(g + (bos + last_idx) * H + i_h)
325
+ bg_last_exp = exp(bg_last)
326
+ p_g = tl.make_block_ptr(g + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
327
+ b_g = tl.load(p_g, boundary_check=(0,))
328
+ b_g_exp = exp(b_g)
329
+
330
+ p_dv = tl.make_block_ptr(dv, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
331
+ p_dv2 = tl.make_block_ptr(dv2, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
332
+ p_do = tl.make_block_ptr(do, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
333
+
334
+ b_do = tl.load(p_do, boundary_check=(0, 1))
335
+
336
+ # Update dv
337
+ p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 0), (BT, 64), (1, 0))
338
+ b_k = tl.load(p_k, boundary_check=(0, 1))
339
+ if USE_GK:
340
+ o_k1 = tl.arange(0, 64)
341
+ b_gk_last1 = tl.load(gk + last_idx * H*K + o_k1, mask=(o_k1 < K), other=0.)
342
+ b_dv = tl.dot(b_k, b_dh1.to(b_k.dtype))
343
+
344
+ if K > 64:
345
+ p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 64), (BT, 64), (1, 0))
346
+ b_k = tl.load(p_k, boundary_check=(0, 1))
347
+ if USE_GK:
348
+ o_k2 = 64 + o_k1
349
+ b_gk_last2 = tl.load(gk + last_idx * H*K + o_k2, mask=(o_k2 < K), other=0.)
350
+ b_dv += tl.dot(b_k, b_dh2.to(b_k.dtype))
351
+
352
+ if K > 128:
353
+ p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 128), (BT, 64), (1, 0))
354
+ b_k = tl.load(p_k, boundary_check=(0, 1))
355
+ if USE_GK:
356
+ o_k3 = 128 + o_k1
357
+ b_gk_last3 = tl.load(gk + last_idx * H*K + o_k3, mask=(o_k3 < K), other=0.)
358
+ b_dv += tl.dot(b_k, b_dh3.to(b_k.dtype))
359
+
360
+ if K > 192:
361
+ p_k = tl.make_block_ptr(k, (T, K), (stride_k, 1), (i_t * BT, 192), (BT, 64), (1, 0))
362
+ b_k = tl.load(p_k, boundary_check=(0, 1))
363
+ if USE_GK:
364
+ o_k4 = 192 + o_k1
365
+ b_gk_last4 = tl.load(gk + last_idx * H*K + o_k4, mask=(o_k4 < K), other=0.)
366
+ b_dv += tl.dot(b_k, b_dh4.to(b_k.dtype))
367
+
368
+ if USE_G:
369
+ m_t = (i_t * BT + tl.arange(0, BT)) < T
370
+ b_dv *= tl.where(m_t, exp(bg_last - b_g), 0)[:, None]
371
+ b_dv += tl.load(p_dv, boundary_check=(0, 1))
372
+
373
+ tl.store(p_dv2, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
374
+ # Update dh
375
+ p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1))
376
+ p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1))
377
+ b_w = tl.load(p_w, boundary_check=(0, 1))
378
+ b_q = tl.load(p_q, boundary_check=(0, 1))
379
+ if USE_G:
380
+ b_dh1 *= bg_last_exp
381
+ b_q = b_q * b_g_exp[None, :]
382
+ if USE_GK:
383
+ b_dh1 *= exp(b_gk_last1[:, None])
384
+ b_dh1 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
385
+ if K > 64:
386
+ p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1))
387
+ p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1))
388
+ b_q = tl.load(p_q, boundary_check=(0, 1))
389
+ b_w = tl.load(p_w, boundary_check=(0, 1))
390
+ if USE_G:
391
+ b_dh2 *= bg_last_exp
392
+ b_q = b_q * b_g_exp[None, :]
393
+ if USE_GK:
394
+ b_dh2 *= exp(b_gk_last2[:, None])
395
+ b_dh2 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
396
+ if K > 128:
397
+ p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (128, i_t * BT), (64, BT), (0, 1))
398
+ p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (128, i_t * BT), (64, BT), (0, 1))
399
+ b_q = tl.load(p_q, boundary_check=(0, 1))
400
+ b_w = tl.load(p_w, boundary_check=(0, 1))
401
+ if USE_G:
402
+ b_dh3 *= bg_last_exp
403
+ b_q = b_q * b_g_exp[None, :]
404
+ if USE_GK:
405
+ b_dh3 *= exp(b_gk_last3[:, None])
406
+ b_dh3 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
407
+ if K > 192:
408
+ p_q = tl.make_block_ptr(q, (K, T), (1, stride_k), (192, i_t * BT), (64, BT), (0, 1))
409
+ p_w = tl.make_block_ptr(w, (K, T), (1, stride_k), (192, i_t * BT), (64, BT), (0, 1))
410
+ b_q = tl.load(p_q, boundary_check=(0, 1))
411
+ b_w = tl.load(p_w, boundary_check=(0, 1))
412
+ if USE_G:
413
+ b_dh4 *= bg_last_exp
414
+ b_q = b_q * b_g_exp[None, :]
415
+ if USE_GK:
416
+ b_dh4 *= exp(b_gk_last4[:, None])
417
+ b_dh4 += tl.dot(b_q.to(b_q.dtype), b_do.to(b_q.dtype)) * scale - tl.dot(b_w, b_dv.to(b_w.dtype))
418
+
419
+ if USE_INITIAL_STATE:
420
+ p_dh0 = tl.make_block_ptr(dh0, (K, V), (V, 1), (0, i_v * BV), (64, BV), (1, 0))
421
+ tl.store(p_dh0, b_dh1.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
422
+ if K > 64:
423
+ p_dh1 = tl.make_block_ptr(dh0, (K, V), (V, 1), (64, i_v * BV), (64, BV), (1, 0))
424
+ tl.store(p_dh1, b_dh2.to(p_dh1.dtype.element_ty), boundary_check=(0, 1))
425
+ if K > 128:
426
+ p_dh2 = tl.make_block_ptr(dh0, (K, V), (V, 1), (128, i_v * BV), (64, BV), (1, 0))
427
+ tl.store(p_dh2, b_dh3.to(p_dh2.dtype.element_ty), boundary_check=(0, 1))
428
+ if K > 192:
429
+ p_dh3 = tl.make_block_ptr(dh0, (K, V), (V, 1), (192, i_v * BV), (64, BV), (1, 0))
430
+ tl.store(p_dh3, b_dh4.to(p_dh3.dtype.element_ty), boundary_check=(0, 1))
431
+
432
+
433
+ def chunk_gated_delta_rule_fwd_h(
434
+ k: torch.Tensor,
435
+ w: torch.Tensor,
436
+ u: torch.Tensor,
437
+ g: torch.Tensor | None = None,
438
+ gk: torch.Tensor | None = None,
439
+ initial_state: torch.Tensor | None = None,
440
+ output_final_state: bool = False,
441
+ chunk_size: int = 64, # SY: remove this argument and force chunk size 64?
442
+ save_new_value: bool = True,
443
+ cu_seqlens: torch.LongTensor | None = None,
444
+ ) -> tuple[torch.Tensor, torch.Tensor]:
445
+ B, T, H, K, V = *k.shape, u.shape[-1]
446
+ BT = chunk_size
447
+
448
+ chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) if cu_seqlens is not None else None
449
+ # N: the actual number of sequences in the batch with either equal or variable lengths
450
+ if cu_seqlens is None:
451
+ N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
452
+ else:
453
+ N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
454
+ assert K <= 256, "current kernel does not support head dimension larger than 256."
455
+
456
+ h = k.new_empty(B, NT, H, K, V)
457
+ final_state = k.new_empty(N, H, K, V, dtype=torch.float32) if output_final_state else None
458
+
459
+ v_new = torch.empty_like(u) if save_new_value else None
460
+ def grid(meta): return (triton.cdiv(V, meta['BV']), N*H)
461
+ chunk_gated_delta_rule_fwd_kernel_h_blockdim64[grid](
462
+ k=k,
463
+ v=u,
464
+ w=w,
465
+ v_new=v_new,
466
+ g=g,
467
+ gk=gk,
468
+ h=h,
469
+ h0=initial_state,
470
+ ht=final_state,
471
+ cu_seqlens=cu_seqlens,
472
+ chunk_offsets=chunk_offsets,
473
+ T=T,
474
+ H=H,
475
+ K=K,
476
+ V=V,
477
+ BT=BT,
478
+ )
479
+ return h, v_new, final_state
480
+
481
+
482
+ def chunk_gated_delta_rule_bwd_dhu(
483
+ q: torch.Tensor,
484
+ k: torch.Tensor,
485
+ w: torch.Tensor,
486
+ do: torch.Tensor,
487
+ dv: torch.Tensor,
488
+ g: torch.Tensor | None = None,
489
+ gk: torch.Tensor | None = None,
490
+ h0: torch.Tensor | None = None,
491
+ dht: torch.Tensor | None = None,
492
+ scale: float | None = None,
493
+ cu_seqlens: torch.LongTensor | None = None,
494
+ chunk_size: int = 64, # SY: remove this argument and force chunk size 64?
495
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
496
+ B, T, H, K, V = *q.shape, do.shape[-1]
497
+ # N: the actual number of sequences in the batch with either equal or variable lengths
498
+ BT = 64
499
+ assert K <= 256, "current kernel does not support head dimension being larger than 256."
500
+
501
+ chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) if cu_seqlens is not None else None
502
+ if cu_seqlens is None:
503
+ N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
504
+ else:
505
+ N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
506
+
507
+ dh = q.new_empty(B, NT, H, K, V)
508
+ dh0 = torch.empty_like(h0, dtype=torch.float32) if h0 is not None else None
509
+ dv2 = torch.empty_like(dv)
510
+
511
+ def grid(meta): return (triton.cdiv(V, meta['BV']), N*H)
512
+ chunk_gated_delta_rule_bwd_kernel_dhu_blockdim64[grid](
513
+ q=q,
514
+ k=k,
515
+ w=w,
516
+ g=g,
517
+ gk=gk,
518
+ dht=dht,
519
+ dh0=dh0,
520
+ do=do,
521
+ dh=dh,
522
+ dv=dv,
523
+ dv2=dv2,
524
+ cu_seqlens=cu_seqlens,
525
+ chunk_offsets=chunk_offsets,
526
+ scale=scale,
527
+ T=T,
528
+ H=H,
529
+ K=K,
530
+ V=V,
531
+ BT=BT,
532
+ )
533
+ return dh, dh0, dv2
code/flash-linear-attention/fla/ops/common/chunk_h.py ADDED
@@ -0,0 +1,394 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils import prepare_chunk_offsets
9
+ from fla.ops.utils.op import exp
10
+ from fla.utils import autotune_cache_kwargs, check_shared_mem
11
+
12
+ BKV_LIST = [32, 64] if check_shared_mem() else [16, 32]
13
+
14
+
15
+ @triton.heuristics({
16
+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
17
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
18
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
19
+ 'HAS_MIXED_PRECISION': lambda args: args['gk'] is not None and args['k'].dtype != args['gk'].dtype,
20
+ })
21
+ @triton.autotune(
22
+ configs=[
23
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
24
+ for BK in BKV_LIST
25
+ for BV in BKV_LIST
26
+ for num_warps in [1, 2, 4, 8]
27
+ for num_stages in [2, 3, 4]
28
+ ],
29
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
30
+ **autotune_cache_kwargs,
31
+ )
32
+ @triton.jit(do_not_specialize=['T'])
33
+ def chunk_fwd_kernel_h(
34
+ k,
35
+ v,
36
+ h,
37
+ g,
38
+ g_gamma,
39
+ gk,
40
+ gv,
41
+ h0,
42
+ ht,
43
+ cu_seqlens,
44
+ split_offsets,
45
+ T,
46
+ H: tl.constexpr,
47
+ K: tl.constexpr,
48
+ V: tl.constexpr,
49
+ BT: tl.constexpr,
50
+ BS: tl.constexpr,
51
+ BK: tl.constexpr,
52
+ BV: tl.constexpr,
53
+ USE_G: tl.constexpr,
54
+ USE_G_GAMMA: tl.constexpr,
55
+ USE_GK: tl.constexpr,
56
+ USE_GV: tl.constexpr,
57
+ USE_INITIAL_STATE: tl.constexpr,
58
+ STORE_FINAL_STATE: tl.constexpr,
59
+ IS_VARLEN: tl.constexpr,
60
+ HAS_MIXED_PRECISION: tl.constexpr = False,
61
+ ):
62
+ i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
63
+ i_n, i_h = i_nh // H, i_nh % H
64
+ if IS_VARLEN:
65
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
66
+ T = eos - bos
67
+ NT, NS = tl.cdiv(T, BT), tl.cdiv(T, BS)
68
+ boh = tl.load(split_offsets + i_n).to(tl.int32)
69
+ else:
70
+ bos, eos = i_n * T, i_n * T + T
71
+ NT, NS = tl.cdiv(T, BT), tl.cdiv(T, BS)
72
+ boh = i_n * NS
73
+ NTS = BS // BT
74
+
75
+ if USE_G_GAMMA:
76
+ # decay rate given the head index
77
+ b_gamma = tl.load(g_gamma + i_h)
78
+ b_g = b_gamma * (tl.arange(0, BT) + 1)
79
+
80
+ # [BK, BV]
81
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
82
+ if USE_INITIAL_STATE:
83
+ p_h0 = tl.make_block_ptr(h0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
84
+ b_h = tl.load(p_h0, boundary_check=(0, 1)).to(tl.float32)
85
+
86
+ for i_t in range(NT):
87
+ i_s = i_t // NTS
88
+ p_k = tl.make_block_ptr(k + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
89
+ p_v = tl.make_block_ptr(v + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
90
+
91
+ o_h = ((boh + i_s) * H + i_h).to(tl.int64) * K*V
92
+ p_h = tl.make_block_ptr(h + o_h, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
93
+
94
+ if i_t % NTS == 0:
95
+ tl.store(p_h, b_h.to(p_h.dtype.element_ty), boundary_check=(0, 1))
96
+ # [BK, BT]
97
+ b_k = tl.load(p_k, boundary_check=(0, 1))
98
+ # [BT, BV]
99
+ b_v = tl.load(p_v, boundary_check=(0, 1))
100
+ last_idx = min((i_t + 1) * BT, T) - 1
101
+
102
+ # scalar decay
103
+ if USE_G:
104
+ b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
105
+ p_g = g + bos*H + (i_t * BT + tl.arange(0, BT)) * H + i_h
106
+ b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
107
+ b_h *= exp(b_g_last)
108
+ b_v = (b_v * exp(b_g_last - b_g)[:, None]).to(b_v.dtype)
109
+
110
+ if USE_G_GAMMA:
111
+ b_g_last = b_gamma * min(BT, T - i_t * BT)
112
+ b_h *= exp(b_g_last)
113
+ b_v = (b_v * exp(b_g_last - b_g)[:, None]).to(b_v.dtype)
114
+
115
+ # vector decay, h = Diag(gk) @ h
116
+ if USE_GK:
117
+ p_gk = tl.make_block_ptr(gk + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
118
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
119
+
120
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
121
+ b_h *= exp(b_gk_last)[:, None]
122
+
123
+ b_gk = tl.load(p_gk, boundary_check=(0, 1))
124
+ b_k = (b_k * exp(b_gk_last[:, None] - b_gk)).to(b_k.dtype)
125
+
126
+ # vector decay, h = h @ Diag(gv)
127
+ if USE_GV:
128
+ p_gv = tl.make_block_ptr(gv + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
129
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
130
+
131
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
132
+ b_h *= exp(b_gv_last)[None, :]
133
+
134
+ b_gv = tl.load(p_gv, boundary_check=(0, 1))
135
+ b_v = (b_v * exp(b_gv_last[None, :] - b_gv)).to(b_v.dtype)
136
+
137
+ if HAS_MIXED_PRECISION:
138
+ b_h += tl.dot(b_k.to(tl.float32), b_v.to(tl.float32))
139
+ else:
140
+ b_h += tl.dot(b_k, b_v)
141
+
142
+ if STORE_FINAL_STATE:
143
+ p_ht = tl.make_block_ptr(ht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
144
+ tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
145
+
146
+
147
+ @triton.heuristics({
148
+ 'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
149
+ 'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
150
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
151
+ 'HAS_MIXED_PRECISION': lambda args: args['gk'] is not None and args['q'].dtype != args['gk'].dtype,
152
+ })
153
+ @triton.autotune(
154
+ configs=[
155
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
156
+ for BK in BKV_LIST
157
+ for BV in BKV_LIST
158
+ for num_warps in [1, 2, 4, 8]
159
+ for num_stages in [2, 3, 4]
160
+ ],
161
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
162
+ **autotune_cache_kwargs,
163
+ )
164
+ @triton.jit(do_not_specialize=['T'])
165
+ def chunk_bwd_kernel_dh(
166
+ q,
167
+ g,
168
+ g_gamma,
169
+ gk,
170
+ gv,
171
+ do,
172
+ dh,
173
+ dht,
174
+ dh0,
175
+ cu_seqlens,
176
+ split_offsets,
177
+ scale,
178
+ T,
179
+ HQ: tl.constexpr,
180
+ H: tl.constexpr,
181
+ K: tl.constexpr,
182
+ V: tl.constexpr,
183
+ BT: tl.constexpr,
184
+ BS: tl.constexpr,
185
+ BK: tl.constexpr,
186
+ BV: tl.constexpr,
187
+ NG: tl.constexpr,
188
+ USE_G: tl.constexpr,
189
+ USE_G_GAMMA: tl.constexpr,
190
+ USE_GK: tl.constexpr,
191
+ USE_GV: tl.constexpr,
192
+ STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
193
+ USE_FINAL_STATE_GRADIENT: tl.constexpr,
194
+ IS_VARLEN: tl.constexpr,
195
+ HAS_MIXED_PRECISION: tl.constexpr = False,
196
+ ):
197
+ i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
198
+ i_n, i_hq = i_nh // HQ, i_nh % HQ
199
+ i_h = i_hq // NG
200
+ if IS_VARLEN:
201
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
202
+ T = eos - bos
203
+ NT = tl.cdiv(T, BT)
204
+ NS = tl.cdiv(T, BS)
205
+ boh = tl.load(split_offsets + i_n).to(tl.int32)
206
+ else:
207
+ bos, eos = i_n * T, i_n * T + T
208
+ NT = tl.cdiv(T, BT)
209
+ NS = tl.cdiv(T, BS)
210
+ boh = i_n * NS
211
+
212
+ if USE_G_GAMMA:
213
+ b_gamma = tl.load(g_gamma + i_h)
214
+ b_g = b_gamma * (tl.arange(0, BT) + 1)
215
+
216
+ # [BK, BV]
217
+ b_dh = tl.zeros([BK, BV], dtype=tl.float32)
218
+ if USE_FINAL_STATE_GRADIENT:
219
+ p_dht = tl.make_block_ptr(dht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
220
+ b_dh += tl.load(p_dht, boundary_check=(0, 1)).to(tl.float32)
221
+
222
+ for i_t in range(NT - 1, -1, -1):
223
+ i_s = i_t // (BS // BT)
224
+ o_dh = ((boh + i_s) * H + i_h).to(tl.int64) * K*V
225
+ p_dh = tl.make_block_ptr(dh + o_dh, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
226
+
227
+ if i_t % (BS // BT) == 0:
228
+ tl.store(p_dh, b_dh.to(p_dh.dtype.element_ty), boundary_check=(0, 1))
229
+ last_idx = min(i_t * BT + BT, T) - 1
230
+ # [BK, BT]
231
+ p_q = tl.make_block_ptr(q + (bos*HQ + i_hq) * K, (K, T), (1, HQ*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
232
+ p_do = tl.make_block_ptr(do + (bos*HQ + i_hq) * V, (T, V), (HQ*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
233
+ b_q = tl.load(p_q, boundary_check=(0, 1))
234
+ b_q = (b_q * scale).to(b_q.dtype)
235
+ # [BT, BV]
236
+ b_do = tl.load(p_do, boundary_check=(0, 1))
237
+
238
+ if USE_G:
239
+ p_g = g + (bos + i_t * BT + tl.arange(0, BT)) * H + i_h
240
+ b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
241
+ b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
242
+ b_q = (b_q * exp(b_g)[None, :]).to(b_q.dtype)
243
+ b_dh *= exp(b_g_last)
244
+
245
+ if USE_G_GAMMA:
246
+ b_g_last = b_gamma * min(BT, T - i_t * BT)
247
+ b_q = (b_q * exp(b_g)[None, :]).to(b_q.dtype)
248
+ b_dh *= exp(b_g_last)
249
+
250
+ if USE_GK:
251
+ p_gk = tl.make_block_ptr(gk + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
252
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
253
+
254
+ b_gk = tl.load(p_gk, boundary_check=(0, 1))
255
+ b_q = (b_q * exp(b_gk)).to(b_q.dtype)
256
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
257
+ b_dh *= exp(b_gk_last)[:, None]
258
+
259
+ if USE_GV:
260
+ p_gv = tl.make_block_ptr(gv + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
261
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
262
+
263
+ b_gv = tl.load(p_gv, boundary_check=(0, 1))
264
+ b_do = (b_do * exp(b_gv))
265
+
266
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
267
+ b_dh *= exp(b_gv_last)[None, :]
268
+
269
+ if HAS_MIXED_PRECISION:
270
+ b_dh += tl.dot(b_q.to(tl.float32), b_do.to(tl.float32))
271
+ else:
272
+ b_dh += tl.dot(b_q, b_do.to(b_q.dtype))
273
+
274
+ if STORE_INITIAL_STATE_GRADIENT:
275
+ p_dh0 = tl.make_block_ptr(dh0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
276
+ tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
277
+
278
+
279
+ def chunk_fwd_h(
280
+ k: torch.Tensor,
281
+ v: torch.Tensor,
282
+ g: torch.Tensor | None = None,
283
+ g_gamma: torch.Tensor | None = None,
284
+ gk: torch.Tensor | None = None,
285
+ gv: torch.Tensor | None = None,
286
+ h0: torch.Tensor | None = None,
287
+ output_final_state: bool = False,
288
+ cu_seqlens: torch.Tensor | None = None,
289
+ chunk_size: int = 64,
290
+ split_size: int | None = None,
291
+ states_in_fp32: bool = False,
292
+ ) -> tuple[torch.Tensor, torch.Tensor]:
293
+ B, T, H, K, V = *k.shape, v.shape[-1]
294
+ BT = chunk_size
295
+ BS = BT if split_size is None else split_size
296
+ assert BS % BT == 0, f"The `split_size` (got {BS}) must be a multiple of `chunk_size` {BT}"
297
+ # N: the actual number of sequences in the batch with either equal or variable lengths
298
+ if cu_seqlens is None:
299
+ N, NS, split_offsets = B, triton.cdiv(T, BS), None
300
+ else:
301
+ split_offsets = prepare_chunk_offsets(cu_seqlens, BS)
302
+ N, NS = len(cu_seqlens) - 1, split_offsets[-1].item()
303
+
304
+ h = k.new_empty(B, NS, H, K, V, dtype=k.dtype if not states_in_fp32 else torch.float)
305
+ ht = k.new_empty(N, H, K, V, dtype=torch.float) if output_final_state else None
306
+ def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
307
+ chunk_fwd_kernel_h[grid](
308
+ k=k,
309
+ v=v,
310
+ h=h,
311
+ g=g,
312
+ g_gamma=g_gamma,
313
+ gk=gk,
314
+ gv=gv,
315
+ h0=h0,
316
+ ht=ht,
317
+ cu_seqlens=cu_seqlens,
318
+ split_offsets=split_offsets,
319
+ T=T,
320
+ H=H,
321
+ K=K,
322
+ V=V,
323
+ BT=BT,
324
+ BS=BS,
325
+ USE_G=g is not None,
326
+ USE_G_GAMMA=g_gamma is not None,
327
+ USE_GK=gk is not None,
328
+ USE_GV=gv is not None,
329
+ )
330
+ return h, ht
331
+
332
+
333
+ def chunk_bwd_dh(
334
+ q: torch.Tensor,
335
+ k: torch.Tensor,
336
+ v: torch.Tensor,
337
+ do: torch.Tensor,
338
+ h0: torch.Tensor,
339
+ dht: torch.Tensor,
340
+ scale: float,
341
+ g: torch.Tensor | None = None,
342
+ g_gamma: torch.Tensor | None = None,
343
+ gk: torch.Tensor | None = None,
344
+ gv: torch.Tensor | None = None,
345
+ cu_seqlens: torch.Tensor | None = None,
346
+ chunk_size: int = 64,
347
+ split_size: int | None = None,
348
+ states_in_fp32: bool = False,
349
+ ) -> tuple[torch.Tensor, torch.Tensor]:
350
+ B, T, H, K, V = *k.shape, v.shape[-1]
351
+ HQ = q.shape[2]
352
+ BT = chunk_size
353
+ BS = BT if split_size is None else split_size
354
+ assert BS % BT == 0, f"The `split_size` (got {BS}) must be a multiple of `chunk_size` {BT}"
355
+ # N: the actual number of sequences in the batch with either equal or variable lengths
356
+ # NG: number of groups in GQA
357
+ if cu_seqlens is None:
358
+ N, NS, split_offsets = B, triton.cdiv(T, BS), None
359
+ else:
360
+ split_offsets = prepare_chunk_offsets(cu_seqlens, BS)
361
+ N, NS = len(cu_seqlens) - 1, split_offsets[-1].item()
362
+ NG = HQ // H
363
+
364
+ dh = k.new_empty(B, NS, HQ, K, V, dtype=k.dtype if not states_in_fp32 else torch.float)
365
+ dh0 = torch.empty_like(h0, dtype=torch.float) if h0 is not None else None
366
+
367
+ def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
368
+ chunk_bwd_kernel_dh[grid](
369
+ q=q,
370
+ g=g,
371
+ g_gamma=g_gamma,
372
+ gk=gk,
373
+ gv=gv,
374
+ do=do,
375
+ dh=dh,
376
+ dht=dht,
377
+ dh0=dh0,
378
+ cu_seqlens=cu_seqlens,
379
+ split_offsets=split_offsets,
380
+ scale=scale,
381
+ T=T,
382
+ HQ=HQ,
383
+ H=H,
384
+ K=K,
385
+ V=V,
386
+ BT=BT,
387
+ BS=BS,
388
+ NG=NG,
389
+ USE_G=g is not None,
390
+ USE_G_GAMMA=g_gamma is not None,
391
+ USE_GK=gk is not None,
392
+ USE_GV=gv is not None,
393
+ )
394
+ return dh, dh0
code/flash-linear-attention/fla/ops/common/chunk_h_parallel.py ADDED
@@ -0,0 +1,554 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+ """
4
+ Fully parallelized state passing.
5
+ """
6
+
7
+
8
+ import torch
9
+ import triton
10
+ import triton.language as tl
11
+
12
+ from fla.ops.utils import prepare_chunk_indices, prepare_chunk_offsets
13
+ from fla.ops.utils.op import exp
14
+ from fla.utils import autotune_cache_kwargs
15
+
16
+
17
+ @triton.heuristics({
18
+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
19
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
20
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
21
+ })
22
+ @triton.autotune(
23
+ configs=[
24
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
25
+ for BK in [32, 64, 128]
26
+ for BV in [32, 64, 128]
27
+ for num_warps in [2, 4, 8]
28
+ for num_stages in [2, 3, 4]
29
+ ],
30
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
31
+ **autotune_cache_kwargs,
32
+ )
33
+ @triton.jit(do_not_specialize=['T'])
34
+ def chunk_fwd_kernel_h_parallel(
35
+ k,
36
+ v,
37
+ h,
38
+ g,
39
+ gk,
40
+ gv,
41
+ h0,
42
+ ht,
43
+ cu_seqlens,
44
+ chunk_indices,
45
+ T,
46
+ H: tl.constexpr,
47
+ K: tl.constexpr,
48
+ V: tl.constexpr,
49
+ BT: tl.constexpr,
50
+ BK: tl.constexpr,
51
+ BV: tl.constexpr,
52
+ USE_G: tl.constexpr,
53
+ USE_GK: tl.constexpr,
54
+ USE_GV: tl.constexpr,
55
+ USE_INITIAL_STATE: tl.constexpr,
56
+ STORE_FINAL_STATE: tl.constexpr,
57
+ IS_VARLEN: tl.constexpr,
58
+ ):
59
+ i_kv, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
60
+
61
+ NV = tl.cdiv(V, BV)
62
+ # i_b: batch index
63
+ # i_h: head index
64
+ # i_n: sequence index
65
+ # i_t: chunk index within current sequence
66
+ # i_tg: (global) chunk index across all sequences
67
+ i_k, i_v = i_kv // NV, i_kv % NV
68
+ i_b, i_h = i_bh // H, i_bh % H
69
+ if IS_VARLEN:
70
+ i_tg = i_t
71
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
72
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
73
+ T = eos - bos
74
+ NT = tl.cdiv(T, BT)
75
+ else:
76
+ bos, eos = i_b * T, i_b * T + T
77
+ NT = tl.cdiv(T, BT)
78
+ i_n, i_tg = i_b, i_b * NT + i_t
79
+ i_nh = i_n * H + i_h
80
+
81
+ p_k = tl.make_block_ptr(k + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
82
+ p_v = tl.make_block_ptr(v + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
83
+ p_h = tl.make_block_ptr(h + (i_tg * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
84
+
85
+ if i_t == 0:
86
+ if USE_INITIAL_STATE:
87
+ p_h0 = tl.make_block_ptr(h0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
88
+ b_h = tl.load(p_h0, boundary_check=(0, 1)).to(tl.float32)
89
+ else:
90
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
91
+ tl.store(p_h, b_h.to(p_h.dtype.element_ty), boundary_check=(0, 1))
92
+
93
+ # [BK, BT]
94
+ b_k = tl.load(p_k, boundary_check=(0, 1))
95
+ # [BT, BV]
96
+ b_v = tl.load(p_v, boundary_check=(0, 1))
97
+
98
+ last_idx = min(i_t * BT + BT, T) - 1
99
+ # scalar decay
100
+ if USE_G:
101
+ b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
102
+ p_g = g + bos*H + (i_t * BT + tl.arange(0, BT)) * H + i_h
103
+ b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
104
+ b_v = (b_v * exp(b_g_last - b_g)[:, None]).to(b_v.dtype)
105
+
106
+ # vector decay, h = Diag(gk) @ h
107
+ if USE_GK:
108
+ p_gk = tl.make_block_ptr(gk + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
109
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
110
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
111
+
112
+ b_gk = tl.load(p_gk, boundary_check=(0, 1))
113
+ b_k = (b_k * exp(b_gk_last[:, None] - b_gk)).to(b_k.dtype)
114
+
115
+ # vector decay, h = h @ Diag(gv)
116
+ if USE_GV:
117
+ p_gv = tl.make_block_ptr(gv + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
118
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
119
+
120
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
121
+ b_gv = tl.load(p_gv, boundary_check=(0, 1))
122
+ b_v = (b_v * exp(b_gv_last[None, :] - b_gv)).to(b_v.dtype)
123
+
124
+ b_h = tl.dot(b_k, b_v)
125
+ if i_t < NT - 1:
126
+ p_h = tl.make_block_ptr(h + ((i_tg + 1) * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
127
+ tl.store(p_h, b_h.to(p_h.dtype.element_ty), boundary_check=(0, 1))
128
+ elif STORE_FINAL_STATE:
129
+ p_ht = tl.make_block_ptr(ht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
130
+ tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
131
+
132
+
133
+ @triton.heuristics({
134
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
135
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
136
+ })
137
+ @triton.autotune(
138
+ configs=[
139
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
140
+ for BK in [32, 64, 128]
141
+ for BV in [32, 64, 128]
142
+ for num_warps in [2, 4, 8, 16]
143
+ for num_stages in [2, 3]
144
+ ],
145
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
146
+ **autotune_cache_kwargs,
147
+ )
148
+ @triton.jit(do_not_specialize=['T'])
149
+ def chunk_fwd_kernel_h_reduction(
150
+ h,
151
+ g,
152
+ gk,
153
+ gv,
154
+ kvt,
155
+ ht,
156
+ cu_seqlens,
157
+ chunk_offsets,
158
+ T,
159
+ H: tl.constexpr,
160
+ K: tl.constexpr,
161
+ V: tl.constexpr,
162
+ BT: tl.constexpr,
163
+ BK: tl.constexpr,
164
+ BV: tl.constexpr,
165
+ USE_G: tl.constexpr,
166
+ USE_GK: tl.constexpr,
167
+ USE_GV: tl.constexpr,
168
+ STORE_FINAL_STATE: tl.constexpr,
169
+ IS_VARLEN: tl.constexpr,
170
+ ):
171
+ i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
172
+ i_n, i_h = i_nh // H, i_nh % H
173
+ if IS_VARLEN:
174
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
175
+ T = eos - bos
176
+ NT = tl.cdiv(T, BT)
177
+ boh = tl.load(chunk_offsets + i_n).to(tl.int32)
178
+ else:
179
+ bos, eos = i_n * T, i_n * T + T
180
+ NT = tl.cdiv(T, BT)
181
+ boh = i_n * NT
182
+
183
+ # [BK, BV]
184
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
185
+ for i_t in range(NT):
186
+ p_h = tl.make_block_ptr(h + ((boh + i_t) * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
187
+ b_h += tl.load(p_h, boundary_check=(0, 1)).to(tl.float32)
188
+ if i_t > 0:
189
+ tl.store(p_h, b_h.to(p_h.dtype.element_ty), boundary_check=(0, 1))
190
+
191
+ last_idx = min(i_t * BT + BT, T) - 1
192
+ # scalar decay
193
+ if USE_G:
194
+ b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
195
+ b_h *= exp(b_g_last)
196
+
197
+ # vector decay, h = Diag(gk) @ h
198
+ if USE_GK:
199
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
200
+
201
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
202
+ b_h *= exp(b_gk_last)[:, None]
203
+
204
+ # vector decay, h = h @ Diag(gv)
205
+ if USE_GV:
206
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
207
+
208
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
209
+ b_h *= exp(b_gv_last)[None, :]
210
+
211
+ if STORE_FINAL_STATE:
212
+ p_kvt = tl.make_block_ptr(kvt + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
213
+ p_ht = tl.make_block_ptr(ht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
214
+ b_h += tl.load(p_kvt, boundary_check=(0, 1)).to(tl.float32)
215
+ tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
216
+
217
+
218
+ @triton.heuristics({
219
+ 'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
220
+ 'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
221
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
222
+ })
223
+ @triton.autotune(
224
+ configs=[
225
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
226
+ for BK in [32, 64, 128]
227
+ for BV in [32, 64, 128]
228
+ for num_warps in [2, 4, 8]
229
+ for num_stages in [2, 3, 4]
230
+ ],
231
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
232
+ **autotune_cache_kwargs,
233
+ )
234
+ @triton.jit(do_not_specialize=['T'])
235
+ def chunk_bwd_kernel_dh_parallel(
236
+ q,
237
+ g,
238
+ gk,
239
+ gv,
240
+ do,
241
+ dh,
242
+ dht,
243
+ dh0,
244
+ cu_seqlens,
245
+ chunk_indices,
246
+ scale,
247
+ T,
248
+ HQ: tl.constexpr,
249
+ H: tl.constexpr,
250
+ K: tl.constexpr,
251
+ V: tl.constexpr,
252
+ BT: tl.constexpr,
253
+ BK: tl.constexpr,
254
+ BV: tl.constexpr,
255
+ NG: tl.constexpr,
256
+ USE_G: tl.constexpr,
257
+ USE_GK: tl.constexpr,
258
+ USE_GV: tl.constexpr,
259
+ STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
260
+ USE_FINAL_STATE_GRADIENT: tl.constexpr,
261
+ IS_VARLEN: tl.constexpr,
262
+ ):
263
+ i_kv, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
264
+
265
+ NV = tl.cdiv(V, BV)
266
+ i_k, i_v = i_kv // NV, i_kv % NV
267
+ i_b, i_hq = i_bh // HQ, i_bh % HQ
268
+ i_h = i_hq // NG
269
+ if IS_VARLEN:
270
+ i_tg = i_t
271
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
272
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
273
+ T = eos - bos
274
+ NT = tl.cdiv(T, BT)
275
+ else:
276
+ bos, eos = i_b * T, i_b * T + T
277
+ NT = tl.cdiv(T, BT)
278
+ i_n, i_tg = i_b, i_b * NT + i_t
279
+ i_nh = i_n * HQ + i_hq
280
+
281
+ p_q = tl.make_block_ptr(q + (bos*HQ + i_hq) * K, (K, T), (1, HQ*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
282
+ p_do = tl.make_block_ptr(do + (bos*HQ + i_hq) * V, (T, V), (HQ*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
283
+ p_dh = tl.make_block_ptr(dh + (i_tg * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
284
+
285
+ if i_t == NT - 1:
286
+ if USE_FINAL_STATE_GRADIENT:
287
+ p_dht = tl.make_block_ptr(dht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
288
+ b_dh = tl.load(p_dht, boundary_check=(0, 1)).to(tl.float32)
289
+ else:
290
+ b_dh = tl.zeros([BK, BV], dtype=tl.float32)
291
+ tl.store(p_dh, b_dh.to(p_dh.dtype.element_ty), boundary_check=(0, 1))
292
+
293
+ # [BK, BT]
294
+ b_q = tl.load(p_q, boundary_check=(0, 1))
295
+ b_q = (b_q * scale).to(b_q.dtype)
296
+ # [BT, BV]
297
+ b_do = tl.load(p_do, boundary_check=(0, 1))
298
+
299
+ if USE_G:
300
+ p_g = g + (bos + i_t * BT + tl.arange(0, BT)) * H + i_h
301
+ b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
302
+ b_q = (b_q * exp(b_g)[None, :]).to(b_q.dtype)
303
+
304
+ if USE_GK:
305
+ p_gk = tl.make_block_ptr(gk + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
306
+ b_gk = tl.load(p_gk, boundary_check=(0, 1))
307
+ b_q = (b_q * exp(b_gk)).to(b_q.dtype)
308
+
309
+ if USE_GV:
310
+ p_gv = tl.make_block_ptr(gv + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
311
+ b_gv = tl.load(p_gv, boundary_check=(0, 1))
312
+ b_do = (b_do * exp(b_gv)).to(b_do.dtype)
313
+
314
+ b_dh = tl.dot(b_q, b_do)
315
+ if i_t > 0:
316
+ p_dh = tl.make_block_ptr(dh + ((i_tg - 1) * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
317
+ tl.store(p_dh, b_dh.to(p_dh.dtype.element_ty), boundary_check=(0, 1))
318
+ elif STORE_INITIAL_STATE_GRADIENT:
319
+ p_dh0 = tl.make_block_ptr(dh0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
320
+ tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
321
+
322
+
323
+ @triton.heuristics({
324
+ 'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
325
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
326
+ })
327
+ @triton.autotune(
328
+ configs=[
329
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
330
+ for BK in [32, 64, 128]
331
+ for BV in [32, 64, 128]
332
+ for num_warps in [2, 4, 8, 16]
333
+ for num_stages in [2, 3]
334
+ ],
335
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
336
+ **autotune_cache_kwargs,
337
+ )
338
+ @triton.jit(do_not_specialize=['T'])
339
+ def chunk_bwd_kernel_dh_reduction(
340
+ g,
341
+ gk,
342
+ gv,
343
+ dh,
344
+ doq0,
345
+ dh0,
346
+ cu_seqlens,
347
+ chunk_offsets,
348
+ T,
349
+ HQ: tl.constexpr,
350
+ H: tl.constexpr,
351
+ K: tl.constexpr,
352
+ V: tl.constexpr,
353
+ BT: tl.constexpr,
354
+ BK: tl.constexpr,
355
+ BV: tl.constexpr,
356
+ NG: tl.constexpr,
357
+ USE_G: tl.constexpr,
358
+ USE_GK: tl.constexpr,
359
+ USE_GV: tl.constexpr,
360
+ STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
361
+ IS_VARLEN: tl.constexpr,
362
+ ):
363
+ i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
364
+ i_n, i_hq = i_nh // HQ, i_nh % HQ
365
+ i_h = i_hq // NG
366
+ if IS_VARLEN:
367
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
368
+ T = eos - bos
369
+ NT = tl.cdiv(T, BT)
370
+ boh = tl.load(chunk_offsets + i_n).to(tl.int32)
371
+ else:
372
+ bos, eos = i_n * T, i_n * T + T
373
+ NT = tl.cdiv(T, BT)
374
+ boh = i_n * NT
375
+
376
+ # [BK, BV]
377
+ b_dh = tl.zeros([BK, BV], dtype=tl.float32)
378
+ for i_t in range(NT - 1, -1, -1):
379
+ p_dh = tl.make_block_ptr(dh + ((boh+i_t) * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
380
+ b_dh += tl.load(p_dh, boundary_check=(0, 1)).to(tl.float32)
381
+ if i_t < NT - 1:
382
+ tl.store(p_dh, b_dh.to(p_dh.dtype.element_ty), boundary_check=(0, 1))
383
+
384
+ last_idx = min(i_t * BT + BT, T) - 1
385
+ if USE_G:
386
+ b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
387
+ b_dh *= exp(b_g_last)
388
+
389
+ if USE_GK:
390
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
391
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
392
+ b_dh *= exp(b_gk_last)[:, None]
393
+
394
+ if USE_GV:
395
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
396
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
397
+ b_dh *= exp(b_gv_last)[None, :]
398
+
399
+ if STORE_INITIAL_STATE_GRADIENT:
400
+ p_doq0 = tl.make_block_ptr(doq0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
401
+ p_dh0 = tl.make_block_ptr(dh0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
402
+ b_dh += tl.load(p_doq0, boundary_check=(0, 1)).to(tl.float32)
403
+ tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
404
+
405
+
406
+ def chunk_fwd_h(
407
+ k: torch.Tensor,
408
+ v: torch.Tensor,
409
+ g: torch.Tensor,
410
+ gk: torch.Tensor,
411
+ gv: torch.Tensor,
412
+ h0: torch.Tensor,
413
+ output_final_state: bool,
414
+ states_in_fp32: bool = False,
415
+ cu_seqlens: torch.Tensor | None = None,
416
+ chunk_size: int = 64,
417
+ ) -> tuple[torch.Tensor, torch.Tensor]:
418
+ B, T, H, K, V = *k.shape, v.shape[-1]
419
+ BT = chunk_size
420
+
421
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
422
+ # N: the actual number of sequences in the batch with either equal or variable lengths
423
+ if cu_seqlens is None:
424
+ N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
425
+ else:
426
+ N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
427
+
428
+ h = k.new_empty(B, NT, H, K, V, dtype=torch.float)
429
+ ht = k.new_empty(N, H, K, V, dtype=torch.float) if output_final_state else None
430
+ def grid(meta): return (triton.cdiv(K, meta['BK']) * triton.cdiv(V, meta['BV']), NT, B * H)
431
+ chunk_fwd_kernel_h_parallel[grid](
432
+ k=k,
433
+ v=v,
434
+ h=h,
435
+ g=g,
436
+ gk=gk,
437
+ gv=gv,
438
+ h0=h0,
439
+ ht=ht,
440
+ cu_seqlens=cu_seqlens,
441
+ chunk_indices=chunk_indices,
442
+ T=T,
443
+ H=H,
444
+ K=K,
445
+ V=V,
446
+ BT=BT,
447
+ USE_G=g is not None,
448
+ USE_GK=gk is not None,
449
+ USE_GV=gv is not None,
450
+ )
451
+ kvt, ht = ht, (torch.empty_like(ht) if output_final_state else None)
452
+ def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
453
+ chunk_fwd_kernel_h_reduction[grid](
454
+ h=h,
455
+ g=g,
456
+ gk=gk,
457
+ gv=gv,
458
+ kvt=kvt,
459
+ ht=ht,
460
+ cu_seqlens=cu_seqlens,
461
+ chunk_offsets=chunk_offsets,
462
+ T=T,
463
+ H=H,
464
+ K=K,
465
+ V=V,
466
+ BT=BT,
467
+ USE_G=g is not None,
468
+ USE_GK=gk is not None,
469
+ USE_GV=gv is not None,
470
+ )
471
+ h = h.to(k.dtype) if not states_in_fp32 else h
472
+ return h, ht
473
+
474
+
475
+ def chunk_bwd_dh(
476
+ q: torch.Tensor,
477
+ k: torch.Tensor,
478
+ v: torch.Tensor,
479
+ g: torch.Tensor,
480
+ gk: torch.Tensor,
481
+ gv: torch.Tensor,
482
+ do: torch.Tensor,
483
+ h0: torch.Tensor,
484
+ dht: torch.Tensor,
485
+ scale: float,
486
+ states_in_fp32: bool = False,
487
+ cu_seqlens: torch.Tensor | None = None,
488
+ chunk_size: int = 64,
489
+ ) -> tuple[torch.Tensor, torch.Tensor]:
490
+ B, T, H, K, V = *k.shape, v.shape[-1]
491
+ HQ = q.shape[2]
492
+ BT = chunk_size
493
+
494
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
495
+ # N: the actual number of sequences in the batch with either equal or variable lengths
496
+ # NG: number of groups in GQA
497
+ if cu_seqlens is None:
498
+ N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
499
+ else:
500
+ N, NT, chunk_offsets = len(cu_seqlens) - 1, len(chunk_indices), prepare_chunk_offsets(cu_seqlens, BT)
501
+ NG = HQ // H
502
+
503
+ dh = k.new_empty(B, NT, HQ, K, V, dtype=k.dtype if not states_in_fp32 else torch.float)
504
+ dh0 = torch.empty_like(h0, dtype=torch.float) if h0 is not None else None
505
+
506
+ def grid(meta): return (triton.cdiv(K, meta['BK']) * triton.cdiv(V, meta['BV']), NT, B * HQ)
507
+ chunk_bwd_kernel_dh_parallel[grid](
508
+ q=q,
509
+ g=g,
510
+ gk=gk,
511
+ gv=gv,
512
+ do=do,
513
+ dh=dh,
514
+ dht=dht,
515
+ dh0=dh0,
516
+ cu_seqlens=cu_seqlens,
517
+ chunk_indices=chunk_indices,
518
+ scale=scale,
519
+ T=T,
520
+ HQ=HQ,
521
+ H=H,
522
+ K=K,
523
+ V=V,
524
+ BT=BT,
525
+ NG=NG,
526
+ USE_G=g is not None,
527
+ USE_GK=gk is not None,
528
+ USE_GV=gv is not None,
529
+ )
530
+
531
+ doq0, dh0 = dh0, (torch.empty_like(dh0) if dh0 is not None else None)
532
+ def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * HQ)
533
+ chunk_bwd_kernel_dh_reduction[grid](
534
+ g=g,
535
+ gk=gk,
536
+ gv=gv,
537
+ dh=dh,
538
+ doq0=doq0,
539
+ dh0=dh0,
540
+ cu_seqlens=cu_seqlens,
541
+ chunk_offsets=chunk_offsets,
542
+ T=T,
543
+ HQ=HQ,
544
+ H=H,
545
+ K=K,
546
+ V=V,
547
+ BT=BT,
548
+ NG=NG,
549
+ USE_G=g is not None,
550
+ USE_GK=gk is not None,
551
+ USE_GV=gv is not None,
552
+ )
553
+ dh = dh.to(q.dtype) if not states_in_fp32 else dh
554
+ return dh, dh0
code/flash-linear-attention/fla/ops/common/chunk_h_split.py ADDED
@@ -0,0 +1,599 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils.op import exp
9
+ from fla.utils import autotune_cache_kwargs
10
+
11
+
12
+ @triton.heuristics({
13
+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
14
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
15
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
16
+ })
17
+ @triton.autotune(
18
+ configs=[
19
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
20
+ for BK in [32, 64]
21
+ for BV in [32, 64]
22
+ for num_warps in [2, 4, 8]
23
+ for num_stages in [2, 3]
24
+ ],
25
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
26
+ **autotune_cache_kwargs,
27
+ )
28
+ @triton.jit(do_not_specialize=['T'])
29
+ def chunk_fwd_kernel_h_split(
30
+ k,
31
+ v,
32
+ g,
33
+ gk,
34
+ gv,
35
+ hs,
36
+ hr,
37
+ h0,
38
+ ht,
39
+ cu_seqlens,
40
+ split_indices,
41
+ T,
42
+ S: tl.constexpr,
43
+ H: tl.constexpr,
44
+ K: tl.constexpr,
45
+ V: tl.constexpr,
46
+ BT: tl.constexpr,
47
+ BK: tl.constexpr,
48
+ BV: tl.constexpr,
49
+ USE_G: tl.constexpr,
50
+ USE_GK: tl.constexpr,
51
+ USE_GV: tl.constexpr,
52
+ USE_INITIAL_STATE: tl.constexpr,
53
+ STORE_FINAL_STATE: tl.constexpr,
54
+ IS_VARLEN: tl.constexpr,
55
+ ):
56
+ # handle one split at a time
57
+ # i_h: head index
58
+ # i_n: sequence index
59
+ # i_s: local split index inside a sequence
60
+ i_k, i_v, i_sh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
61
+ i_ss, i_h = i_sh // H, i_sh % H
62
+ if IS_VARLEN:
63
+ i_n, i_s = tl.load(split_indices + i_ss * 2).to(tl.int32), tl.load(split_indices + i_ss * 2 + 1).to(tl.int32)
64
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
65
+ T = eos - bos
66
+ NS = tl.cdiv(T, S)
67
+ else:
68
+ NS = tl.cdiv(T, S)
69
+ i_n, i_s = i_ss // NS, i_ss % NS
70
+ bos, eos = i_n * T, i_n * T + T
71
+ i_nh = i_n * H + i_h
72
+
73
+ # [BK, BV]
74
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
75
+ # for the first split, we directly store the state as the final result
76
+ if i_s == 0:
77
+ if USE_INITIAL_STATE:
78
+ p_h0 = tl.make_block_ptr(h0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
79
+ b_h += tl.load(p_h0, boundary_check=(0, 1)).to(tl.float32)
80
+ p_hr = tl.make_block_ptr(hr + i_sh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
81
+ tl.store(p_hr, b_h.to(p_hr.dtype.element_ty), boundary_check=(0, 1))
82
+ for i_t in range(tl.cdiv(i_s * S, BT), tl.cdiv(min(i_s * S + S, T), BT)):
83
+ p_k = tl.make_block_ptr(k + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
84
+ p_v = tl.make_block_ptr(v + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
85
+ # [BK, BT]
86
+ b_k = tl.load(p_k, boundary_check=(0, 1))
87
+ # [BT, BV]
88
+ b_v = tl.load(p_v, boundary_check=(0, 1))
89
+ last_idx = min(i_t * BT + BT, T) - 1
90
+
91
+ # scalar decay
92
+ if USE_G:
93
+ b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
94
+ p_g = g + bos*H + (i_t * BT + tl.arange(0, BT)) * H + i_h
95
+ b_h *= exp(b_g_last)
96
+ b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
97
+ b_v = (b_v * exp(b_g_last - b_g)[:, None]).to(b_v.dtype)
98
+
99
+ # vector decay, h = Diag(gk) @ h
100
+ if USE_GK:
101
+ p_gk = tl.make_block_ptr(gk + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
102
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
103
+
104
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
105
+ b_h *= exp(b_gk_last)[:, None]
106
+
107
+ b_gk = tl.load(p_gk, boundary_check=(0, 1))
108
+ b_k = (b_k * exp(b_gk_last[:, None] - b_gk)).to(b_k.dtype)
109
+
110
+ # vector decay, h = h @ Diag(gv)
111
+ if USE_GV:
112
+ p_gv = tl.make_block_ptr(gv + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
113
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
114
+
115
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
116
+ b_h *= exp(b_gv_last)[None, :]
117
+
118
+ b_gv = tl.load(p_gv, boundary_check=(0, 1))
119
+ b_v = (b_v * exp(b_gv_last[None, :] - b_gv)).to(b_v.dtype)
120
+
121
+ b_h += tl.dot(b_k, b_v)
122
+
123
+ # if there are more than one splits, we store the result to (unreduced) hs
124
+ # otherwise, we store the result to ht as the final state
125
+ if NS > 1:
126
+ p_hs = tl.make_block_ptr(hs + i_sh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
127
+ tl.store(p_hs, b_h.to(p_hs.dtype.element_ty), boundary_check=(0, 1))
128
+ elif STORE_FINAL_STATE:
129
+ p_ht = tl.make_block_ptr(ht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
130
+ tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
131
+
132
+
133
+ @triton.heuristics({
134
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
135
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
136
+ })
137
+ @triton.autotune(
138
+ configs=[
139
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
140
+ for BK in [32, 64]
141
+ for BV in [32, 64]
142
+ for num_warps in [2, 4, 8]
143
+ for num_stages in [2, 3, 4]
144
+ ],
145
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
146
+ **autotune_cache_kwargs,
147
+ )
148
+ @triton.jit(do_not_specialize=['T'])
149
+ def chunk_fwd_kernel_h_reduction(
150
+ g,
151
+ gk,
152
+ gv,
153
+ hs,
154
+ hr,
155
+ ht,
156
+ cu_seqlens,
157
+ split_offsets,
158
+ T,
159
+ S: tl.constexpr,
160
+ H: tl.constexpr,
161
+ K: tl.constexpr,
162
+ V: tl.constexpr,
163
+ BT: tl.constexpr,
164
+ BK: tl.constexpr,
165
+ BV: tl.constexpr,
166
+ USE_G: tl.constexpr,
167
+ USE_GK: tl.constexpr,
168
+ USE_GV: tl.constexpr,
169
+ STORE_FINAL_STATE: tl.constexpr,
170
+ IS_VARLEN: tl.constexpr,
171
+ ):
172
+ i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
173
+ i_n, i_h = i_nh // H, i_nh % H
174
+ if IS_VARLEN:
175
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
176
+ T = eos - bos
177
+ NS = tl.cdiv(T, S)
178
+ boh = tl.load(split_offsets + i_n).to(tl.int32)
179
+ else:
180
+ bos, eos = i_n * T, i_n * T + T
181
+ NS = tl.cdiv(T, S)
182
+ boh = i_n * NS
183
+
184
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
185
+ # skip the first split
186
+ for i_s in range(1, NS):
187
+ p_hs = tl.make_block_ptr(hs + ((boh + i_s-1) * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
188
+ p_hr = tl.make_block_ptr(hr + ((boh + i_s) * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
189
+ b_h += tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
190
+ tl.store(p_hr, b_h.to(p_hr.dtype.element_ty), boundary_check=(0, 1))
191
+
192
+ for i_t in range(tl.cdiv(i_s * S, BT), tl.cdiv(min(i_s * S + S, T), BT)):
193
+ last_idx = min(i_t * BT + BT, T) - 1
194
+ # scalar decay
195
+ if USE_G:
196
+ b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
197
+ b_h *= exp(b_g_last)
198
+
199
+ # vector decay, h = Diag(gk) @ h
200
+ if USE_GK:
201
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
202
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
203
+ b_h *= exp(b_gk_last)[:, None]
204
+
205
+ # vector decay, h = h @ Diag(gv)
206
+ if USE_GV:
207
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
208
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
209
+ b_h *= exp(b_gv_last)[None, :]
210
+
211
+ if NS > 1:
212
+ if STORE_FINAL_STATE:
213
+ p_hs = tl.make_block_ptr(hs + ((boh + NS-1) * H + i_h)*K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
214
+ p_ht = tl.make_block_ptr(ht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
215
+ b_h += tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
216
+ tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
217
+
218
+
219
+ @triton.heuristics({
220
+ 'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
221
+ 'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
222
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
223
+ })
224
+ @triton.autotune(
225
+ configs=[
226
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
227
+ for BK in [32, 64]
228
+ for BV in [32, 64]
229
+ for num_warps in [2, 4, 8]
230
+ for num_stages in [2, 3]
231
+ ],
232
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
233
+ **autotune_cache_kwargs,
234
+ )
235
+ @triton.jit(do_not_specialize=['T'])
236
+ def chunk_bwd_kernel_dh_split(
237
+ q,
238
+ g,
239
+ gk,
240
+ gv,
241
+ do,
242
+ dht,
243
+ dhs,
244
+ dhr,
245
+ dh0,
246
+ cu_seqlens,
247
+ split_indices,
248
+ scale,
249
+ T,
250
+ S: tl.constexpr,
251
+ HQ: tl.constexpr,
252
+ H: tl.constexpr,
253
+ K: tl.constexpr,
254
+ V: tl.constexpr,
255
+ BT: tl.constexpr,
256
+ BK: tl.constexpr,
257
+ BV: tl.constexpr,
258
+ NG: tl.constexpr,
259
+ USE_G: tl.constexpr,
260
+ USE_GK: tl.constexpr,
261
+ USE_GV: tl.constexpr,
262
+ USE_FINAL_STATE_GRADIENT: tl.constexpr,
263
+ STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
264
+ IS_VARLEN: tl.constexpr,
265
+ ):
266
+ # handle one split at a time
267
+ # i_h: head index
268
+ # i_n: sequence index
269
+ # i_s: local split index inside a sequence
270
+ i_k, i_v, i_sh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
271
+ i_ss, i_hq = i_sh // HQ, i_sh % HQ
272
+ if IS_VARLEN:
273
+ i_n, i_s = tl.load(split_indices + i_ss * 2).to(tl.int32), tl.load(split_indices + i_ss * 2 + 1).to(tl.int32)
274
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
275
+ T = eos - bos
276
+ NS = tl.cdiv(T, S)
277
+ else:
278
+ NS = tl.cdiv(T, S)
279
+ i_n, i_s = i_ss // NS, i_ss % NS
280
+ bos, eos = i_n * T, i_n * T + T
281
+ i_nh = i_n * HQ + i_hq
282
+ i_h = i_hq // NG
283
+
284
+ # [BK, BV]
285
+ b_dh = tl.zeros([BK, BV], dtype=tl.float32)
286
+ if i_s == NS - 1:
287
+ if USE_FINAL_STATE_GRADIENT:
288
+ p_dht = tl.make_block_ptr(dht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
289
+ b_dh += tl.load(p_dht, boundary_check=(0, 1)).to(tl.float32)
290
+ p_dhr = tl.make_block_ptr(dhr + i_sh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
291
+ tl.store(p_dhr, b_dh.to(p_dhr.dtype.element_ty), boundary_check=(0, 1))
292
+
293
+ for i_t in range(tl.cdiv(min(i_s * S + S, T), BT) - 1, tl.cdiv(i_s * S, BT) - 1, -1):
294
+ p_q = tl.make_block_ptr(q + (bos*HQ + i_hq) * K, (K, T), (1, HQ*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
295
+ p_do = tl.make_block_ptr(do + (bos*HQ + i_hq) * V, (T, V), (HQ*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
296
+
297
+ b_q = tl.load(p_q, boundary_check=(0, 1))
298
+ b_q = (b_q * scale).to(b_q.dtype)
299
+ # [BT, BV]
300
+ b_do = tl.load(p_do, boundary_check=(0, 1))
301
+
302
+ last_idx = min(i_t * BT + BT, T) - 1
303
+ if USE_G:
304
+ p_g = g + (bos + i_t * BT + tl.arange(0, BT)) * H + i_h
305
+ b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
306
+ b_g = tl.load(p_g, mask=(i_t * BT + tl.arange(0, BT) < T), other=0.)
307
+ b_q = (b_q * exp(b_g)[None, :]).to(b_q.dtype)
308
+ b_dh *= exp(b_g_last)
309
+
310
+ if USE_GK:
311
+ p_gk = tl.make_block_ptr(gk + (bos*H + i_h) * K, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
312
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
313
+
314
+ b_gk = tl.load(p_gk, boundary_check=(0, 1))
315
+ b_q = (b_q * exp(b_gk)).to(b_q.dtype)
316
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
317
+ b_dh *= exp(b_gk_last)[:, None]
318
+
319
+ if USE_GV:
320
+ p_gv = tl.make_block_ptr(gv + (bos*H + i_h) * V, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
321
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
322
+
323
+ b_gv = tl.load(p_gv, boundary_check=(0, 1))
324
+ b_do = (b_do * exp(b_gv)).to(b_do.dtype)
325
+
326
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
327
+ b_dh *= exp(b_gv_last)[None, :]
328
+
329
+ b_dh += tl.dot(b_q, b_do)
330
+
331
+ if NS > 1:
332
+ p_dhs = tl.make_block_ptr(dhs + i_sh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
333
+ tl.store(p_dhs, b_dh.to(p_dhs.dtype.element_ty), boundary_check=(0, 1))
334
+ elif STORE_INITIAL_STATE_GRADIENT:
335
+ p_dh0 = tl.make_block_ptr(dh0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
336
+ tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
337
+
338
+
339
+ @triton.heuristics({
340
+ 'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
341
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
342
+ })
343
+ @triton.autotune(
344
+ configs=[
345
+ triton.Config({'BK': BK, 'BV': BV}, num_warps=num_warps, num_stages=num_stages)
346
+ for BK in [32, 64]
347
+ for BV in [32, 64]
348
+ for num_warps in [2, 4, 8]
349
+ for num_stages in [2, 3, 4]
350
+ ],
351
+ key=['BT', 'USE_G', 'USE_GK', 'USE_GV'],
352
+ **autotune_cache_kwargs,
353
+ )
354
+ @triton.jit(do_not_specialize=['T'])
355
+ def chunk_bwd_kernel_dh_reduction(
356
+ g,
357
+ gk,
358
+ gv,
359
+ dhs,
360
+ dhr,
361
+ dh0,
362
+ cu_seqlens,
363
+ split_offsets,
364
+ T,
365
+ S: tl.constexpr,
366
+ H: tl.constexpr,
367
+ HQ: tl.constexpr,
368
+ K: tl.constexpr,
369
+ V: tl.constexpr,
370
+ BT: tl.constexpr,
371
+ BK: tl.constexpr,
372
+ BV: tl.constexpr,
373
+ NG: tl.constexpr,
374
+ USE_G: tl.constexpr,
375
+ USE_GK: tl.constexpr,
376
+ USE_GV: tl.constexpr,
377
+ STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
378
+ IS_VARLEN: tl.constexpr,
379
+ ):
380
+ i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
381
+ i_n, i_hq = i_nh // HQ, i_nh % HQ
382
+ i_h = i_hq // NG
383
+ if IS_VARLEN:
384
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
385
+ T = eos - bos
386
+ NS = tl.cdiv(T, S)
387
+ boh = tl.load(split_offsets + i_n).to(tl.int32)
388
+ else:
389
+ bos, eos = i_n * T, i_n * T + T
390
+ NS = tl.cdiv(T, S)
391
+ boh = i_n * NS
392
+
393
+ b_dh = tl.zeros([BK, BV], dtype=tl.float32)
394
+ for i_s in range(NS - 2, -1, -1):
395
+ p_dhs = tl.make_block_ptr(dhs + ((boh+i_s+1) * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
396
+ p_dhr = tl.make_block_ptr(dhr + ((boh+i_s) * H + i_h) * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
397
+ b_dh += tl.load(p_dhs, boundary_check=(0, 1)).to(tl.float32)
398
+ tl.store(p_dhr, b_dh.to(p_dhr.dtype.element_ty), boundary_check=(0, 1))
399
+
400
+ for i_t in range(tl.cdiv(min(i_s * S + S, T), BT) - 1, tl.cdiv(i_s * S, BT) - 1, -1):
401
+ last_idx = min(i_t * BT + BT, T) - 1
402
+ # scalar decay
403
+ if USE_G:
404
+ b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
405
+ b_dh *= exp(b_g_last)
406
+
407
+ if USE_GK:
408
+ p_gk_last = gk + (bos + last_idx) * H*K + i_h * K + i_k * BK + tl.arange(0, BK)
409
+ b_gk_last = tl.load(p_gk_last, mask=(i_k * BK + tl.arange(0, BK) < K), other=0.)
410
+ b_dh *= exp(b_gk_last)[:, None]
411
+
412
+ if USE_GV:
413
+ p_gv_last = gv + (bos + last_idx) * H*V + i_h * V + i_v * BV + tl.arange(0, BV)
414
+ b_gv_last = tl.load(p_gv_last, mask=(i_v * BV + tl.arange(0, BV) < V), other=0.)
415
+ b_dh *= exp(b_gv_last)[None, :]
416
+
417
+ if NS > 1:
418
+ if STORE_INITIAL_STATE_GRADIENT:
419
+ p_dhs = tl.make_block_ptr(dhs + (boh * H + i_h)*K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
420
+ p_dh0 = tl.make_block_ptr(dh0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
421
+ b_dh += tl.load(p_dhs, boundary_check=(0, 1)).to(tl.float32)
422
+ tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
423
+
424
+
425
+ def chunk_fwd_h(
426
+ k: torch.Tensor,
427
+ v: torch.Tensor,
428
+ g: torch.Tensor,
429
+ gk: torch.Tensor,
430
+ gv: torch.Tensor,
431
+ h0: torch.Tensor,
432
+ output_final_state: bool,
433
+ cu_seqlens: torch.LongTensor | None = None,
434
+ split_offsets: torch.LongTensor | None = None,
435
+ split_indices: torch.LongTensor | None = None,
436
+ chunk_size: int = 64,
437
+ split_size: int = 256,
438
+ states_in_fp32: bool = True,
439
+ ) -> tuple[torch.Tensor, torch.Tensor]:
440
+ B, T, H, K, V = *k.shape, v.shape[-1]
441
+ # B: batch size
442
+ # N: the actual number of sequences in the batch
443
+ # H: number of heads
444
+ # T: sequence length, can be variable across sequences
445
+ # S: split size, a multiple of chunk size
446
+ # BT: chunk size
447
+ S, BT = split_size, chunk_size
448
+ assert S % BT == 0, f"The `split_size` (got {S}) must be a multiple of `chunk_size` {BT}"
449
+ if cu_seqlens is None:
450
+ N = B
451
+ NS = N * triton.cdiv(T, S)
452
+ else:
453
+ N = len(cu_seqlens) - 1
454
+ NS = split_offsets[-1]
455
+
456
+ # unreduced kv states per split
457
+ hs = k.new_empty(NS, H, K, V, dtype=torch.float)
458
+ # reduced states per split
459
+ hr = k.new_empty(NS, H, K, V, dtype=torch.float if states_in_fp32 else k.dtype)
460
+ ht = k.new_empty(N, H, K, V, dtype=torch.float) if output_final_state else None
461
+ # parallelized over splits
462
+ def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), NS * H)
463
+ chunk_fwd_kernel_h_split[grid](
464
+ k=k,
465
+ v=v,
466
+ g=g,
467
+ gk=gk,
468
+ gv=gv,
469
+ hs=hs,
470
+ hr=hr,
471
+ h0=h0,
472
+ ht=ht,
473
+ cu_seqlens=cu_seqlens,
474
+ split_indices=split_indices,
475
+ T=T,
476
+ S=S,
477
+ H=H,
478
+ K=K,
479
+ V=V,
480
+ BT=BT,
481
+ USE_G=g is not None,
482
+ USE_GK=gk is not None,
483
+ USE_GV=gv is not None,
484
+ )
485
+ def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * H)
486
+ chunk_fwd_kernel_h_reduction[grid](
487
+ g=g,
488
+ gk=gk,
489
+ gv=gv,
490
+ hs=hs,
491
+ hr=hr,
492
+ ht=ht,
493
+ cu_seqlens=cu_seqlens,
494
+ split_offsets=split_offsets,
495
+ T=T,
496
+ S=S,
497
+ H=H,
498
+ K=K,
499
+ V=V,
500
+ BT=BT,
501
+ USE_G=g is not None,
502
+ USE_GK=gk is not None,
503
+ USE_GV=gv is not None,
504
+ )
505
+ return hr, ht
506
+
507
+
508
+ def chunk_bwd_dh(
509
+ q: torch.Tensor,
510
+ k: torch.Tensor,
511
+ v: torch.Tensor,
512
+ g: torch.Tensor,
513
+ gk: torch.Tensor,
514
+ gv: torch.Tensor,
515
+ do: torch.Tensor,
516
+ h0: torch.Tensor,
517
+ dht: torch.Tensor,
518
+ scale: float,
519
+ cu_seqlens: torch.Tensor | None = None,
520
+ split_offsets: torch.Tensor | None = None,
521
+ split_indices: torch.Tensor | None = None,
522
+ chunk_size: int = 64,
523
+ split_size: int = 256,
524
+ states_in_fp32: bool = True,
525
+ ) -> tuple[torch.Tensor, torch.Tensor]:
526
+ B, T, H, K, V = *k.shape, v.shape[-1]
527
+ HQ = q.shape[2]
528
+ # B: batch size
529
+ # N: the actual number of sequences in the batch
530
+ # H: number of heads
531
+ # T: sequence length, can be variable across sequences
532
+ # S: split size, a multiple of chunk size
533
+ # BT: chunk size
534
+ S, BT = max(chunk_size, min(split_size, triton.next_power_of_2(T))), chunk_size
535
+ assert S % BT == 0, f"The `split_size` (got {S}) must be a multiple of `chunk_size` {BT}"
536
+ if cu_seqlens is None:
537
+ N = B
538
+ NS = N * triton.cdiv(T, S)
539
+ else:
540
+ N = len(cu_seqlens) - 1
541
+ NS = split_offsets[-1]
542
+ # number of groups in GQA
543
+ NG = HQ // H
544
+
545
+ dhs = q.new_empty(NS, HQ, K, V, dtype=torch.float)
546
+ dhr = q.new_empty(NS, HQ, K, V, dtype=torch.float if states_in_fp32 else k.dtype)
547
+ dh0 = torch.empty_like(h0, dtype=torch.float) if h0 is not None else None
548
+
549
+ # parallelized over splits
550
+ def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), NS * HQ)
551
+ chunk_bwd_kernel_dh_split[grid](
552
+ q=q,
553
+ g=g,
554
+ gk=gk,
555
+ gv=gv,
556
+ do=do,
557
+ dht=dht,
558
+ dhs=dhs,
559
+ dhr=dhr,
560
+ dh0=dh0,
561
+ cu_seqlens=cu_seqlens,
562
+ split_indices=split_indices,
563
+ scale=scale,
564
+ T=T,
565
+ S=S,
566
+ HQ=HQ,
567
+ H=H,
568
+ K=K,
569
+ V=V,
570
+ BT=BT,
571
+ NG=NG,
572
+ USE_G=g is not None,
573
+ USE_GK=gk is not None,
574
+ USE_GV=gv is not None,
575
+ )
576
+
577
+ def grid(meta): return (triton.cdiv(K, meta['BK']), triton.cdiv(V, meta['BV']), N * HQ)
578
+ chunk_bwd_kernel_dh_reduction[grid](
579
+ g=g,
580
+ gk=gk,
581
+ gv=gv,
582
+ dhs=dhs,
583
+ dhr=dhr,
584
+ dh0=dh0,
585
+ cu_seqlens=cu_seqlens,
586
+ split_offsets=split_offsets,
587
+ T=T,
588
+ S=S,
589
+ HQ=HQ,
590
+ H=H,
591
+ K=K,
592
+ V=V,
593
+ BT=BT,
594
+ NG=NG,
595
+ USE_G=g is not None,
596
+ USE_GK=gk is not None,
597
+ USE_GV=gv is not None,
598
+ )
599
+ return dhr, dh0
code/flash-linear-attention/fla/ops/common/chunk_o.py ADDED
@@ -0,0 +1,689 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils import prepare_chunk_indices
9
+ from fla.ops.utils.op import exp
10
+ from fla.utils import autotune_cache_kwargs, check_shared_mem, is_nvidia_hopper
11
+
12
+ BKV_LIST = [64, 128] if check_shared_mem() else [32, 64]
13
+ NUM_WARPS = [2, 4] if is_nvidia_hopper else [2, 4, 8]
14
+
15
+
16
+ @triton.heuristics({
17
+ 'USE_G': lambda args: args['g'] is not None,
18
+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
19
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
20
+ })
21
+ @triton.autotune(
22
+ configs=[
23
+ triton.Config({'BK': 128, 'BV': 128}, num_warps=8, num_stages=3),
24
+ triton.Config({'BK': 64, 'BV': 64}, num_warps=4, num_stages=3),
25
+ triton.Config({'BK': 32, 'BV': 32}, num_warps=2, num_stages=3),
26
+ ],
27
+ key=['H', 'K', 'V', 'BT'],
28
+ **autotune_cache_kwargs,
29
+ )
30
+ @triton.jit(do_not_specialize=['T'])
31
+ def chunk_fwd_kernel_o(
32
+ q,
33
+ k,
34
+ v,
35
+ h,
36
+ g,
37
+ g_gamma,
38
+ o,
39
+ cu_seqlens,
40
+ chunk_indices,
41
+ scale,
42
+ T,
43
+ H: tl.constexpr,
44
+ K: tl.constexpr,
45
+ V: tl.constexpr,
46
+ BT: tl.constexpr,
47
+ BK: tl.constexpr,
48
+ BV: tl.constexpr,
49
+ USE_G: tl.constexpr,
50
+ USE_G_GAMMA: tl.constexpr,
51
+ IS_VARLEN: tl.constexpr,
52
+ ):
53
+ i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
54
+ i_b, i_h = i_bh // H, i_bh % H
55
+
56
+ if IS_VARLEN:
57
+ i_tg = i_t
58
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
59
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
60
+ T = eos - bos
61
+ NT = tl.cdiv(T, BT)
62
+ else:
63
+ NT = tl.cdiv(T, BT)
64
+ i_tg = i_b * NT + i_t
65
+ bos, eos = i_b * T, i_b * T + T
66
+
67
+ # offset calculation
68
+ q += (bos * H + i_h) * K
69
+ k += (bos * H + i_h) * K
70
+ v += (bos * H + i_h) * V
71
+ o += (bos * H + i_h) * V
72
+ h += (i_tg * H + i_h).to(tl.int64) * K*V
73
+
74
+ b_o = tl.zeros([BT, BV], dtype=tl.float32)
75
+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
76
+
77
+ for i_k in range(tl.cdiv(K, BK)):
78
+ p_q = tl.make_block_ptr(q, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
79
+ p_k = tl.make_block_ptr(k, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
80
+ p_h = tl.make_block_ptr(h, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
81
+ # [BT, BK]
82
+ b_q = tl.load(p_q, boundary_check=(0, 1))
83
+ # [BK, BT]
84
+ b_k = tl.load(p_k, boundary_check=(0, 1))
85
+ # [BK, BV]
86
+ b_h = tl.load(p_h, boundary_check=(0, 1))
87
+
88
+ # [BT, BK] @ [BK, BV] -> [BT, BV]
89
+ b_o += tl.dot(b_q, b_h)
90
+ # [BT, BK] @ [BK, BT] -> [BT, BT]
91
+ b_A += tl.dot(b_q, b_k)
92
+
93
+ if USE_G:
94
+ g += bos * H + i_h
95
+ p_g = tl.make_block_ptr(g, (T,), (H,), (i_t * BT,), (BT,), (0,))
96
+ b_g = tl.load(p_g, boundary_check=(0,))
97
+ b_o = b_o * exp(b_g)[:, None]
98
+ b_A = b_A * exp(b_g[:, None] - b_g[None, :])
99
+
100
+ if USE_G_GAMMA:
101
+ b_gamma = tl.load(g_gamma + i_h)
102
+ b_g = b_gamma * (tl.arange(0, BT) + 1)
103
+ b_o = b_o * exp(b_g)[:, None]
104
+ b_A = b_A * exp(b_g[:, None] - b_g[None, :])
105
+
106
+ o_t = i_t * BT + tl.arange(0, BT)
107
+ m_t = o_t < T
108
+ m_A = (o_t[:, None] >= o_t[None, :]) & (m_t[:, None] & m_t)
109
+ b_A = tl.where(m_A, b_A, 0)
110
+
111
+ p_v = tl.make_block_ptr(v, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
112
+ p_o = tl.make_block_ptr(o, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
113
+
114
+ b_v = tl.load(p_v, boundary_check=(0, 1))
115
+ # to fix mma -> mma layout conversion
116
+ # already solved by triton v3.2 or higher
117
+ b_o = b_o * scale + tl.dot(b_A.to(b_v.dtype), b_v) * scale
118
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
119
+
120
+
121
+ @triton.heuristics({
122
+ 'USE_G': lambda args: args['g'] is not None,
123
+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
124
+ 'USE_DW': lambda args: args['dw'] is not None,
125
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
126
+ })
127
+ @triton.autotune(
128
+ configs=[
129
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
130
+ for num_warps in NUM_WARPS
131
+ for num_stages in [2, 3, 4]
132
+ ],
133
+ key=['H', 'K', 'V', 'BT', 'BK', 'BV', 'USE_G', 'USE_G_GAMMA', 'USE_DW'],
134
+ **autotune_cache_kwargs,
135
+ )
136
+ @triton.jit(do_not_specialize=['T'])
137
+ def chunk_bwd_kernel_dqkwg(
138
+ q,
139
+ k,
140
+ v,
141
+ g,
142
+ g_gamma,
143
+ h,
144
+ do,
145
+ dh,
146
+ dq,
147
+ dk,
148
+ dw,
149
+ dv,
150
+ dg,
151
+ cu_seqlens,
152
+ chunk_indices,
153
+ scale,
154
+ B: tl.constexpr,
155
+ T,
156
+ H: tl.constexpr,
157
+ K: tl.constexpr,
158
+ V: tl.constexpr,
159
+ BT: tl.constexpr,
160
+ BK: tl.constexpr,
161
+ BV: tl.constexpr,
162
+ USE_G: tl.constexpr,
163
+ USE_G_GAMMA: tl.constexpr,
164
+ USE_DW: tl.constexpr,
165
+ IS_VARLEN: tl.constexpr,
166
+ ):
167
+ i_k, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
168
+ i_b, i_h = i_bh // H, i_bh % H
169
+
170
+ all = B * T
171
+ if IS_VARLEN:
172
+ i_tg = i_t
173
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
174
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
175
+ T = eos - bos
176
+ NT = tl.cdiv(T, BT)
177
+ else:
178
+ NT = tl.cdiv(T, BT)
179
+ i_tg = i_b * NT + i_t
180
+ bos, eos = i_b * T, i_b * T + T
181
+
182
+ # offset calculation
183
+ v += (bos * H + i_h) * V
184
+ do += (bos * H + i_h) * V
185
+ h += (i_tg * H + i_h).to(tl.int64) * K*V
186
+ dh += (i_tg * H + i_h).to(tl.int64) * K*V
187
+ q += (bos * H + i_h) * K
188
+ k += (bos * H + i_h) * K
189
+ dq += (bos * H + i_h) * K
190
+ dk += (bos * H + i_h) * K
191
+
192
+ # for delta rule only
193
+ if USE_DW:
194
+ dw += (bos * H + i_h) * K
195
+ dv += (bos * H + i_h) * V
196
+
197
+ if USE_G:
198
+ dg += i_k * all * H
199
+ b_dg_last = tl.zeros([1], dtype=tl.float32) if USE_G else None
200
+ if USE_G_GAMMA:
201
+ b_gamma = tl.load(g_gamma + i_h)
202
+ b_g = b_gamma * (tl.arange(0, BT) + 1)
203
+ b_g_last = b_gamma * min(BT, T - i_t * BT)
204
+ b_dq = tl.zeros([BT, BK], dtype=tl.float32)
205
+ b_dk = tl.zeros([BT, BK], dtype=tl.float32)
206
+ b_ds = tl.zeros([BT, BT], dtype=tl.float32)
207
+ b_dw = tl.zeros([BT, BK], dtype=tl.float32) if USE_DW else None
208
+
209
+ for i_v in range(tl.cdiv(V, BV)):
210
+ p_v = tl.make_block_ptr(v, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
211
+ p_do = tl.make_block_ptr(do, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
212
+ p_h = tl.make_block_ptr(h, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))
213
+ p_dh = tl.make_block_ptr(dh, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))
214
+ # [BT, BV]
215
+ b_v = tl.load(p_v, boundary_check=(0, 1))
216
+ b_do = tl.load(p_do, boundary_check=(0, 1))
217
+ # [BV, BK]
218
+ b_h = tl.load(p_h, boundary_check=(0, 1))
219
+ b_dh = tl.load(p_dh, boundary_check=(0, 1))
220
+ if USE_G:
221
+ b_dg_last += (tl.sum(b_h * b_dh))
222
+ # [BT, BV] @ [BV, BT] -> [BT, BT]
223
+ b_ds += tl.dot(b_do, tl.trans(b_v))
224
+ # [BT, BV] @ [BV, BK] -> [BT, BK]
225
+ b_dq += tl.dot(b_do, b_h.to(b_do.dtype))
226
+ # [BT, BV] @ [BV, BK] -> [BT, BK]
227
+ b_dk += tl.dot(b_v, b_dh.to(b_v.dtype))
228
+ if USE_DW:
229
+ p_dv = tl.make_block_ptr(dv, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
230
+ b_dv = tl.load(p_dv, boundary_check=(0, 1))
231
+ b_dw += tl.dot(b_dv.to(b_v.dtype), b_h.to(b_v.dtype))
232
+
233
+ if USE_DW:
234
+ p_dw = tl.make_block_ptr(dw, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
235
+ tl.store(p_dw, -b_dw.to(p_dw.dtype.element_ty), boundary_check=(0, 1))
236
+
237
+ tl.debug_barrier()
238
+ p_q = tl.make_block_ptr(q, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
239
+ p_k = tl.make_block_ptr(k, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
240
+ b_q = tl.load(p_q, boundary_check=(0, 1))
241
+ b_k = tl.load(p_k, boundary_check=(0, 1))
242
+
243
+ p_dq = tl.make_block_ptr(dq, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
244
+ p_dk = tl.make_block_ptr(dk, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
245
+
246
+ o_t = i_t * BT + tl.arange(0, BT)
247
+ m_t = o_t < T
248
+ m_A = (o_t[:, None] >= o_t[None, :]) & (m_t[:, None] & m_t)
249
+ if USE_G:
250
+ b_dg = tl.zeros([BT], dtype=tl.float32)
251
+ g += bos * H + i_h
252
+ dg += bos * H + i_h
253
+ p_g = tl.make_block_ptr(g, (T,), (H,), (i_t * BT,), (BT,), (0,))
254
+ b_g = tl.load(p_g, boundary_check=(0,))
255
+ b_g_last = tl.load(g + (min(i_t * BT + BT, T) - 1) * H)
256
+ b_dg_last *= exp(b_g_last)
257
+
258
+ b_dq = b_dq * exp(b_g)[:, None] * scale
259
+ b_dg += tl.sum(b_dq * b_q, axis=1)
260
+
261
+ b_dk = b_dk * tl.where(m_t, exp(-b_g + b_g_last), 0)[:, None]
262
+ b_dg -= tl.sum(b_k * b_dk, axis=1)
263
+ b_dg_last += tl.sum(b_dk * b_k)
264
+
265
+ b_ds = tl.where(m_A, b_ds * exp(b_g[:, None] - b_g[None, :]), 0) * scale
266
+ b_ds2 = b_ds * tl.dot(b_q, tl.trans(b_k))
267
+ b_dg += tl.sum(b_ds2, axis=1)
268
+ b_dg -= tl.sum(b_ds2, axis=0)
269
+
270
+ b_ds = b_ds.to(b_k.dtype)
271
+ # [BT, BK]
272
+ b_dq += tl.dot(b_ds, b_k)
273
+ b_dk += tl.dot(tl.trans(b_ds), b_q)
274
+ p_dg = tl.make_block_ptr(dg, (T,), (H,), (i_t * BT,), (BT,), (0,))
275
+ # (SY 09/21) revcumsum in a separate kernel due to strange triton compiler issue
276
+ # b_dg = tl.dot(tl.where(o_t[:, None] <= o_t[None, :], 1., 0.), b_dg, allow_tf32=False) + b_dg_last)
277
+ b_dg = tl.where(o_t < min(i_t * BT + BT, T) - 1, b_dg, b_dg + b_dg_last)
278
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
279
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
280
+ tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), boundary_check=(0,))
281
+
282
+ elif USE_G_GAMMA:
283
+ b_dq = b_dq * exp(b_g)[:, None] * scale
284
+ b_dk = b_dk * tl.where(m_t, exp(-b_g + b_g_last), 0)[:, None]
285
+ b_ds = tl.where(m_A, b_ds * exp(b_g[:, None] - b_g[None, :]), 0) * scale
286
+ b_ds = b_ds.to(b_k.dtype)
287
+ # [BT, BK]
288
+ b_dq += tl.dot(b_ds, b_k)
289
+ b_dk += tl.dot(tl.trans(b_ds), b_q)
290
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
291
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
292
+
293
+ else:
294
+ b_ds = tl.where(m_A, b_ds, 0)
295
+ b_ds = b_ds.to(b_k.dtype)
296
+ b_dq += tl.dot(b_ds, b_k)
297
+ b_dk += tl.dot(tl.trans(b_ds), b_q) * scale
298
+ b_dq *= scale
299
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
300
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
301
+
302
+
303
+ @triton.heuristics({
304
+ 'USE_G': lambda args: args['g'] is not None,
305
+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
306
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
307
+ })
308
+ @triton.autotune(
309
+ configs=[
310
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
311
+ for num_warps in NUM_WARPS
312
+ for num_stages in [2, 3, 4]
313
+ ],
314
+ key=['H', 'K', 'V', 'BT', 'BK', 'BV', 'USE_G', 'USE_G_GAMMA'],
315
+ **autotune_cache_kwargs,
316
+ )
317
+ @triton.jit(do_not_specialize=['T'])
318
+ def chunk_bwd_kernel_dv(
319
+ q,
320
+ k,
321
+ g,
322
+ g_gamma,
323
+ do,
324
+ dv,
325
+ dh,
326
+ cu_seqlens,
327
+ chunk_indices,
328
+ scale,
329
+ T,
330
+ H: tl.constexpr,
331
+ K: tl.constexpr,
332
+ V: tl.constexpr,
333
+ BT: tl.constexpr,
334
+ BK: tl.constexpr,
335
+ BV: tl.constexpr,
336
+ USE_G: tl.constexpr,
337
+ USE_G_GAMMA: tl.constexpr,
338
+ IS_VARLEN: tl.constexpr,
339
+ ):
340
+ i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
341
+ i_b, i_h = i_bh // H, i_bh % H
342
+ if IS_VARLEN:
343
+ i_tg = i_t
344
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
345
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
346
+ T = eos - bos
347
+ NT = tl.cdiv(T, BT)
348
+ else:
349
+ NT = tl.cdiv(T, BT)
350
+ i_tg = i_b * NT + i_t
351
+ bos, eos = i_b * T, i_b * T + T
352
+
353
+ b_dv = tl.zeros([BT, BV], dtype=tl.float32)
354
+
355
+ # offset calculation
356
+ q += (bos * H + i_h) * K
357
+ k += (bos * H + i_h) * K
358
+ do += (bos * H + i_h) * V
359
+ dv += (bos * H + i_h) * V
360
+ dh += (i_tg * H + i_h).to(tl.int64) * K*V
361
+
362
+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
363
+ for i_k in range(tl.cdiv(K, BK)):
364
+ p_k = tl.make_block_ptr(k, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
365
+ p_q = tl.make_block_ptr(q, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
366
+ b_q = tl.load(p_q, boundary_check=(0, 1))
367
+ b_k = tl.load(p_k, boundary_check=(0, 1))
368
+ b_A += tl.dot(b_k, b_q)
369
+ p_dh = tl.make_block_ptr(dh, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
370
+ b_dh = tl.load(p_dh, boundary_check=(0, 1))
371
+ b_dv += tl.dot(b_k, b_dh.to(b_k.dtype))
372
+
373
+ o_t = i_t * BT + tl.arange(0, BT)
374
+ m_t = o_t < T
375
+ if USE_G:
376
+ g += bos * H + i_h
377
+ p_g = tl.make_block_ptr(g, (T,), (H,), (i_t * BT,), (BT,), (0,))
378
+ b_g = tl.load(p_g, boundary_check=(0,))
379
+ b_g_last = tl.load(g + (min(i_t * BT + BT, T) - 1) * H)
380
+ if USE_G_GAMMA:
381
+ b_gamma = tl.load(g_gamma + i_h)
382
+ b_g = b_gamma * (tl.arange(0, BT) + 1)
383
+ b_g_last = b_gamma * min(BT, T - i_t * BT)
384
+
385
+ m_A = (o_t[:, None] <= o_t[None, :]) & (m_t[:, None] & m_t)
386
+ if USE_G or USE_G_GAMMA:
387
+ b_A = tl.where(m_A, b_A * exp(b_g[None, :] - b_g[:, None]) * scale, 0).to(do.dtype.element_ty)
388
+ b_dv *= tl.where(m_t, exp(-b_g + b_g_last), 0)[:, None]
389
+ else:
390
+ b_A = tl.where(m_A, b_A * scale, 0).to(do.dtype.element_ty)
391
+ p_do = tl.make_block_ptr(do, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
392
+ p_dv = tl.make_block_ptr(dv, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
393
+ b_do = tl.load(p_do, boundary_check=(0, 1))
394
+ b_dv += tl.dot(b_A.to(b_do.dtype), b_do)
395
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
396
+
397
+
398
+ @triton.heuristics({
399
+ 'USE_G': lambda args: args['g'] is not None,
400
+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
401
+ 'USE_A': lambda args: args['A'] is not None,
402
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
403
+ })
404
+ @triton.autotune(
405
+ configs=[
406
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
407
+ for num_warps in NUM_WARPS
408
+ for num_stages in [2, 3, 4]
409
+ ],
410
+ key=['H', 'K', 'V', 'BT', 'BK', 'BV', 'USE_G'],
411
+ **autotune_cache_kwargs,
412
+ )
413
+ @triton.jit(do_not_specialize=['T'])
414
+ def chunk_bwd_kernel_dv_local(
415
+ q,
416
+ k,
417
+ g,
418
+ g_gamma,
419
+ A,
420
+ do,
421
+ dv,
422
+ cu_seqlens,
423
+ chunk_indices,
424
+ scale,
425
+ T,
426
+ H: tl.constexpr,
427
+ K: tl.constexpr,
428
+ V: tl.constexpr,
429
+ BT: tl.constexpr,
430
+ BK: tl.constexpr,
431
+ BV: tl.constexpr,
432
+ USE_G: tl.constexpr,
433
+ USE_G_GAMMA: tl.constexpr,
434
+ USE_A: tl.constexpr,
435
+ IS_VARLEN: tl.constexpr,
436
+ ):
437
+ i_t, i_bh = tl.program_id(0), tl.program_id(1)
438
+ i_b, i_h = i_bh // H, i_bh % H
439
+ if IS_VARLEN:
440
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
441
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
442
+ T = eos - bos
443
+ else:
444
+ bos, eos = i_b * T, i_b * T + T
445
+
446
+ # offset calculation
447
+ q += (bos * H + i_h) * K
448
+ k += (bos * H + i_h) * K
449
+ do += (bos * H + i_h) * V
450
+ dv += (bos * H + i_h) * V
451
+
452
+ if USE_A:
453
+ p_A = tl.make_block_ptr(A + (bos * H + i_h) * BT, (BT, T), (1, H*BT), (0, i_t * BT), (BT, BT), (0, 1))
454
+ b_A = tl.load(p_A, boundary_check=(0, 1))
455
+ else:
456
+ if USE_G:
457
+ g += bos * H + i_h
458
+ p_g = tl.make_block_ptr(g, (T,), (H,), (i_t * BT,), (BT,), (0,))
459
+ b_g = tl.load(p_g, boundary_check=(0,))
460
+ if USE_G_GAMMA:
461
+ b_gamma = tl.load(g_gamma + i_h)
462
+ b_g = b_gamma * (tl.arange(0, BT) + 1)
463
+
464
+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
465
+ for i_k in range(tl.cdiv(K, BK)):
466
+ p_k = tl.make_block_ptr(k, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
467
+ p_q = tl.make_block_ptr(q, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
468
+
469
+ b_k = tl.load(p_k, boundary_check=(0, 1))
470
+ b_q = tl.load(p_q, boundary_check=(0, 1))
471
+ b_A += tl.dot(b_k, b_q) * scale
472
+ if USE_G or USE_G_GAMMA:
473
+ b_A *= exp(b_g[None, :] - b_g[:, None])
474
+
475
+ o_t = i_t * BT + tl.arange(0, BT)
476
+ m_t = o_t < T
477
+ m_A = (o_t[:, None] <= o_t[None, :]) & (m_t[:, None] & m_t)
478
+ b_A = tl.where(m_A, b_A, 0).to(do.dtype.element_ty)
479
+
480
+ for i_v in range(tl.cdiv(V, BV)):
481
+ p_do = tl.make_block_ptr(do, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
482
+ p_dv = tl.make_block_ptr(dv, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
483
+ b_do = tl.load(p_do, boundary_check=(0, 1))
484
+ b_dv = tl.dot(b_A.to(b_do.dtype), b_do)
485
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
486
+
487
+
488
+ def chunk_fwd_o(
489
+ q: torch.Tensor,
490
+ k: torch.Tensor,
491
+ v: torch.Tensor,
492
+ h: torch.Tensor,
493
+ g: torch.Tensor | None = None,
494
+ g_gamma: torch.Tensor | None = None,
495
+ scale: float | None = None,
496
+ cu_seqlens: torch.LongTensor | None = None,
497
+ chunk_size: int = 64,
498
+ ) -> torch.Tensor:
499
+ B, T, H, K, V = *q.shape, v.shape[-1]
500
+ BT = chunk_size
501
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
502
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
503
+ if scale is None:
504
+ scale = k.shape[-1] ** -0.5
505
+
506
+ o = torch.empty_like(v)
507
+ def grid(meta): return (triton.cdiv(V, meta['BV']), NT, B * H)
508
+ chunk_fwd_kernel_o[grid](
509
+ q=q,
510
+ k=k,
511
+ v=v,
512
+ h=h,
513
+ g=g,
514
+ g_gamma=g_gamma,
515
+ o=o,
516
+ cu_seqlens=cu_seqlens,
517
+ chunk_indices=chunk_indices,
518
+ scale=scale,
519
+ T=T,
520
+ H=H,
521
+ K=K,
522
+ V=V,
523
+ BT=BT,
524
+ )
525
+ return o
526
+
527
+
528
+ def chunk_bwd_dv(
529
+ q: torch.Tensor,
530
+ k: torch.Tensor,
531
+ do: torch.Tensor,
532
+ dh: torch.Tensor,
533
+ g: torch.Tensor | None = None,
534
+ g_gamma: torch.Tensor | None = None,
535
+ scale: float | None = None,
536
+ cu_seqlens: torch.LongTensor | None = None,
537
+ chunk_size: int = 64,
538
+ ) -> torch.Tensor:
539
+ B, T, H, K, V = *k.shape, do.shape[-1]
540
+ BT = chunk_size
541
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
542
+ # H100 can have larger block size
543
+ if check_shared_mem('hopper', k.device.index):
544
+ CONST_TILING = 128
545
+ elif check_shared_mem:
546
+ CONST_TILING = 64
547
+ else:
548
+ CONST_TILING = 32
549
+ BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
550
+ BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
551
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
552
+ NV = triton.cdiv(V, BV)
553
+ if scale is None:
554
+ scale = k.shape[-1] ** -0.5
555
+
556
+ dv = torch.empty_like(do)
557
+ grid = (NV, NT, B * H)
558
+ chunk_bwd_kernel_dv[grid](
559
+ q=q,
560
+ k=k,
561
+ g=g,
562
+ g_gamma=g_gamma,
563
+ do=do,
564
+ dv=dv,
565
+ dh=dh,
566
+ cu_seqlens=cu_seqlens,
567
+ chunk_indices=chunk_indices,
568
+ scale=scale,
569
+ T=T,
570
+ H=H,
571
+ K=K,
572
+ V=V,
573
+ BT=BT,
574
+ BK=BK,
575
+ BV=BV,
576
+ )
577
+ return dv
578
+
579
+
580
+ def chunk_bwd_dv_local(
581
+ q: torch.Tensor,
582
+ k: torch.Tensor,
583
+ do: torch.Tensor,
584
+ g: torch.Tensor | None = None,
585
+ g_gamma: torch.Tensor | None = None,
586
+ A: torch.Tensor | None = None,
587
+ scale: float = None,
588
+ cu_seqlens: torch.LongTensor | None = None,
589
+ chunk_size: int = 64,
590
+ ) -> torch.Tensor:
591
+ B, T, H, K, V = *k.shape, do.shape[-1]
592
+ BT = chunk_size
593
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
594
+ # H100 can have larger block size
595
+ if check_shared_mem('hopper', k.device.index):
596
+ CONST_TILING = 128
597
+ elif check_shared_mem:
598
+ CONST_TILING = 64
599
+ else:
600
+ CONST_TILING = 32
601
+ BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
602
+ BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
603
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
604
+
605
+ dv = torch.empty_like(do)
606
+ grid = (NT, B * H)
607
+ chunk_bwd_kernel_dv_local[grid](
608
+ q=q,
609
+ k=k,
610
+ g=g,
611
+ g_gamma=g_gamma,
612
+ A=A,
613
+ do=do,
614
+ dv=dv,
615
+ cu_seqlens=cu_seqlens,
616
+ chunk_indices=chunk_indices,
617
+ scale=scale,
618
+ T=T,
619
+ H=H,
620
+ K=K,
621
+ V=V,
622
+ BT=BT,
623
+ BK=BK,
624
+ BV=BV,
625
+ )
626
+ return dv
627
+
628
+
629
+ def chunk_bwd_dqkwg(
630
+ q: torch.Tensor,
631
+ k: torch.Tensor,
632
+ v: torch.Tensor,
633
+ do: torch.Tensor,
634
+ h: torch.Tensor,
635
+ dh: torch.Tensor,
636
+ w: torch.Tensor | None = None,
637
+ g: torch.Tensor | None = None,
638
+ g_gamma: torch.Tensor | None = None,
639
+ dv: torch.Tensor | None = None,
640
+ scale: float | None = None,
641
+ cu_seqlens: torch.LongTensor | None = None,
642
+ chunk_size: int = 64,
643
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
644
+
645
+ B, T, H, K, V = *k.shape, v.shape[-1]
646
+ BT = chunk_size
647
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
648
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
649
+
650
+ CONST_TILING = 64 if check_shared_mem() else 32
651
+ BK = min(max(triton.next_power_of_2(K), 16), CONST_TILING)
652
+ BV = min(max(triton.next_power_of_2(V), 16), CONST_TILING)
653
+ NK = triton.cdiv(K, BK)
654
+ dq = torch.empty_like(q)
655
+ dk = torch.empty_like(k)
656
+ dg = torch.empty(NK, *g.shape, dtype=torch.float32, device=g.device) if g is not None else None
657
+ dw = torch.empty_like(w) if w is not None else None
658
+
659
+ grid = (NK, NT, B * H)
660
+ chunk_bwd_kernel_dqkwg[grid](
661
+ q=q,
662
+ k=k,
663
+ v=v,
664
+ g=g,
665
+ g_gamma=g_gamma,
666
+ h=h,
667
+ do=do,
668
+ dh=dh,
669
+ dw=dw,
670
+ dq=dq,
671
+ dk=dk,
672
+ dv=dv,
673
+ dg=dg,
674
+ cu_seqlens=cu_seqlens,
675
+ chunk_indices=chunk_indices,
676
+ scale=scale,
677
+ B=B,
678
+ T=T,
679
+ H=H,
680
+ K=K,
681
+ V=V,
682
+ BT=BT,
683
+ BK=BK,
684
+ BV=BV,
685
+ )
686
+
687
+ if dg is not None:
688
+ dg = dg.sum(0)
689
+ return dq, dk, dw, dg
code/flash-linear-attention/fla/ops/common/chunk_scaled_dot_kkt.py ADDED
@@ -0,0 +1,124 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils import prepare_chunk_indices
9
+ from fla.ops.utils.op import exp
10
+ from fla.utils import autotune_cache_kwargs
11
+
12
+
13
+ @triton.heuristics({
14
+ 'USE_G': lambda args: args['g'] is not None,
15
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
16
+ })
17
+ @triton.autotune(
18
+ configs=[
19
+ triton.Config({'BK': BK}, num_warps=num_warps, num_stages=num_stages)
20
+ for BK in [32, 64, 128]
21
+ for num_warps in [2, 4, 8]
22
+ for num_stages in [2, 3, 4]
23
+ ],
24
+ key=['H', 'K', 'BT', 'IS_VARLEN'],
25
+ **autotune_cache_kwargs,
26
+ )
27
+ @triton.jit(do_not_specialize=['T'])
28
+ def chunk_scaled_dot_kkt_fwd_kernel(
29
+ k,
30
+ g,
31
+ beta,
32
+ A,
33
+ cu_seqlens,
34
+ chunk_indices,
35
+ T,
36
+ H: tl.constexpr,
37
+ K: tl.constexpr,
38
+ BT: tl.constexpr,
39
+ BK: tl.constexpr,
40
+ IS_VARLEN: tl.constexpr,
41
+ USE_G: tl.constexpr,
42
+ ):
43
+ i_t, i_bh = tl.program_id(0), tl.program_id(1)
44
+ i_b, i_h = i_bh // H, i_bh % H
45
+ if IS_VARLEN:
46
+ i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32)
47
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
48
+ T = eos - bos
49
+ else:
50
+ bos, eos = i_b * T, i_b * T + T
51
+ o_t = i_t * BT + tl.arange(0, BT)
52
+ m_t = o_t < T
53
+
54
+ p_b = tl.make_block_ptr(beta + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
55
+ b_b = tl.load(p_b, boundary_check=(0,))
56
+
57
+ b_A = tl.zeros([BT, BT], dtype=tl.float32)
58
+ for i_k in range(tl.cdiv(K, BK)):
59
+ p_k = tl.make_block_ptr(k + (bos*H + i_h) * K, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
60
+ b_k = tl.load(p_k, boundary_check=(0, 1))
61
+ b_A += tl.dot(b_k, tl.trans(b_k))
62
+
63
+ if USE_G:
64
+ p_g = tl.make_block_ptr(g + bos*H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
65
+ b_g = tl.load(p_g, boundary_check=(0,))
66
+ b_g_diff = b_g[:, None] - b_g[None, :]
67
+ b_A *= exp(b_g_diff)
68
+ b_A *= b_b[:, None]
69
+
70
+ m_A = (o_t[:, None] > o_t[None, :]) & (m_t[:, None] & m_t)
71
+ b_A = tl.where(m_A, b_A, 0)
72
+ p_A = tl.make_block_ptr(A + (bos*H + i_h) * BT, (T, BT), (BT*H, 1), (i_t * BT, 0), (BT, BT), (1, 0))
73
+ tl.store(p_A, b_A.to(p_A.dtype.element_ty), boundary_check=(0, 1))
74
+
75
+
76
+ def chunk_scaled_dot_kkt_fwd(
77
+ k: torch.Tensor,
78
+ g: torch.Tensor | None = None,
79
+ beta: torch.Tensor | None = None,
80
+ cu_seqlens: torch.LongTensor | None = None,
81
+ chunk_size: int = 64,
82
+ output_dtype: torch.dtype = torch.float32,
83
+ ) -> torch.Tensor:
84
+ r"""
85
+ Compute beta * K * K^T.
86
+
87
+ Args:
88
+ k (torch.Tensor):
89
+ The key tensor of shape `[B, T, H, K]`.
90
+ beta (torch.Tensor):
91
+ The beta tensor of shape `[B, T, H]`.
92
+ g (torch.Tensor):
93
+ The cumulative sum of the gate tensor of shape `[B, T, H]`. Default: `None`.
94
+ gk (torch.Tensor):
95
+ The cumulative sum of the gate tensor of shape `[B, T, H, K]` applied to the key tensor. Default: `None`.
96
+ cu_seqlens (torch.LongTensor):
97
+ The cumulative sequence lengths of the input tensor.
98
+ Default: None
99
+ chunk_size (int):
100
+ The chunk size. Default: 64.
101
+ output_dtype (torch.dtype):
102
+ The dtype of the output tensor. Default: `torch.float32`
103
+
104
+ Returns:
105
+ beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size.
106
+ """
107
+ B, T, H, K = k.shape
108
+ BT = chunk_size
109
+ chunk_indices = prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
110
+ NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
111
+ A = torch.empty(B, T, H, BT, device=k.device, dtype=output_dtype)
112
+ chunk_scaled_dot_kkt_fwd_kernel[(NT, B * H)](
113
+ k=k,
114
+ g=g,
115
+ beta=beta,
116
+ A=A,
117
+ cu_seqlens=cu_seqlens,
118
+ chunk_indices=chunk_indices,
119
+ T=T,
120
+ H=H,
121
+ K=K,
122
+ BT=BT,
123
+ )
124
+ return A
code/flash-linear-attention/fla/ops/common/fused_chunk.py ADDED
@@ -0,0 +1,636 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils import chunk_local_cumsum
9
+ from fla.ops.utils.op import exp
10
+ from fla.utils import (
11
+ autocast_custom_bwd,
12
+ autocast_custom_fwd,
13
+ autotune_cache_kwargs,
14
+ check_shared_mem,
15
+ input_guard,
16
+ is_nvidia_hopper,
17
+ )
18
+
19
+ BKV_LIST = [64, 128] if check_shared_mem() else [32, 64]
20
+ NUM_WARPS = [2, 4] if is_nvidia_hopper else [2, 4, 8]
21
+
22
+
23
+ @triton.heuristics({
24
+ 'USE_G': lambda args: args['g'] is not None,
25
+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
26
+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
27
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
28
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
29
+ })
30
+ @triton.autotune(
31
+ configs=[
32
+ triton.Config({'BV': BV}, num_warps=num_warps, num_stages=num_stages)
33
+ for BV in BKV_LIST
34
+ for num_warps in NUM_WARPS
35
+ for num_stages in [2, 3, 4]
36
+ ],
37
+ key=['H', 'K', 'V', 'BT'],
38
+ **autotune_cache_kwargs,
39
+ )
40
+ @triton.jit(do_not_specialize=['T'])
41
+ def fused_chunk_fwd_kernel(
42
+ q,
43
+ k,
44
+ v,
45
+ g,
46
+ g_gamma,
47
+ o,
48
+ h0,
49
+ ht,
50
+ cu_seqlens,
51
+ scale,
52
+ T,
53
+ B: tl.constexpr,
54
+ H: tl.constexpr,
55
+ K: tl.constexpr,
56
+ V: tl.constexpr,
57
+ BT: tl.constexpr,
58
+ BK: tl.constexpr,
59
+ BV: tl.constexpr,
60
+ USE_G: tl.constexpr,
61
+ USE_G_GAMMA: tl.constexpr,
62
+ USE_INITIAL_STATE: tl.constexpr,
63
+ STORE_FINAL_STATE: tl.constexpr,
64
+ IS_VARLEN: tl.constexpr,
65
+ ):
66
+ i_v, i_k, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
67
+ i_n, i_h = i_nh // H, i_nh % H
68
+
69
+ all = B * T
70
+ if IS_VARLEN:
71
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
72
+ T = eos - bos
73
+ else:
74
+ bos, eos = i_n * T, i_n * T + T
75
+ NT = tl.cdiv(T, BT)
76
+
77
+ o_i = tl.arange(0, BT)
78
+
79
+ if USE_G_GAMMA:
80
+ # decay rate given the head index
81
+ b_gamma = tl.load(g_gamma + i_h)
82
+ b_g = b_gamma * (o_i + 1)
83
+ b_g_last = b_gamma * BT
84
+ b_gq = exp(b_g)
85
+ b_gk = exp(b_g_last - b_g)
86
+ b_gn = exp(b_g_last)
87
+
88
+ # [BT, BT]
89
+ m_s = o_i[:, None] >= o_i[None, :]
90
+
91
+ q = q + (bos*H + i_h) * K
92
+ k = k + (bos*H + i_h) * K
93
+ v = v + (bos*H + i_h) * V
94
+ o = o + (i_k * all + bos).to(tl.int64) * H*V + i_h * V
95
+
96
+ # [BK, BV]
97
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
98
+ if USE_INITIAL_STATE:
99
+ p_h = tl.make_block_ptr(h0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
100
+ b_h = tl.load(p_h, boundary_check=(0, 1)).to(tl.float32)
101
+
102
+ for i_t in range(0, NT):
103
+ p_q = tl.make_block_ptr(q, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
104
+ p_k = tl.make_block_ptr(k, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
105
+ p_v = tl.make_block_ptr(v, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
106
+ p_o = tl.make_block_ptr(o, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
107
+
108
+ o_t = i_t * BT + tl.arange(0, BT)
109
+ m_t = o_t < T
110
+ # [BT, BK]
111
+ b_q = tl.load(p_q, boundary_check=(0, 1))
112
+ b_q = (b_q * scale).to(b_q.dtype)
113
+ # [BK, BT]
114
+ b_k = tl.load(p_k, boundary_check=(0, 1))
115
+ # [BT, BV]
116
+ b_v = tl.load(p_v, boundary_check=(0, 1))
117
+ last_idx = min(i_t * BT + BT, T) - 1
118
+
119
+ # [BT, BT]
120
+ b_s = tl.dot(b_q, b_k)
121
+
122
+ # scalar decay
123
+ if USE_G:
124
+ p_g = g + (bos + o_t) * H + i_h
125
+ b_g = tl.load(p_g, mask=(o_t < T), other=0.)
126
+ b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
127
+
128
+ b_gq = exp(b_g)
129
+ b_gk = exp(b_g_last - b_g)
130
+ b_gn = exp(b_g_last)
131
+ if USE_G_GAMMA:
132
+ b_g_last = b_gamma * min(BT, T - i_t * BT)
133
+ b_gk = exp(b_g_last - b_g)
134
+ b_gn = exp(b_g_last)
135
+ if USE_G or USE_G_GAMMA:
136
+ b_gs = tl.where(m_s & m_t, exp(b_g[:, None] - b_g[None, :]), 0)
137
+ # [BT, BT]
138
+ b_s *= b_gs
139
+ # [BT, BV]
140
+ b_o = tl.dot(b_s.to(b_q.dtype), b_v) + tl.dot(b_q, b_h.to(b_q.dtype)) * b_gq[:, None]
141
+ b_v = (b_v * b_gk[:, None]).to(b_v.dtype)
142
+ b_h *= b_gn
143
+ else:
144
+ # [BT, BT]
145
+ b_s *= m_s & m_t
146
+ # [BT, BV]
147
+ b_o = tl.dot(b_s.to(b_q.dtype), b_v) + tl.dot(b_q, b_h.to(b_q.dtype))
148
+
149
+ b_h += tl.dot(b_k, b_v)
150
+
151
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
152
+
153
+ if STORE_FINAL_STATE:
154
+ p_ht = tl.make_block_ptr(ht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
155
+ tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
156
+
157
+
158
+ @triton.heuristics({
159
+ 'USE_G': lambda args: args['g'] is not None,
160
+ 'USE_G_GAMMA': lambda args: args['g_gamma'] is not None,
161
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
162
+ 'USE_INITIAL_STATE': lambda args: args['dh0'] is not None,
163
+ 'USE_FINAL_STATE': lambda args: args['dht'] is not None,
164
+ })
165
+ @triton.autotune(
166
+ configs=[
167
+ triton.Config({}, num_warps=num_warps, num_stages=num_stages)
168
+ for num_warps in NUM_WARPS
169
+ for num_stages in [2, 3, 4]
170
+ ],
171
+ key=['H', 'K', 'V', 'BT'],
172
+ **autotune_cache_kwargs,
173
+ )
174
+ @triton.jit(do_not_specialize=['T'])
175
+ def fused_chunk_bwd_kernel(
176
+ q,
177
+ k,
178
+ v,
179
+ g,
180
+ g_gamma,
181
+ do,
182
+ dq,
183
+ dk,
184
+ dv,
185
+ dg,
186
+ h0,
187
+ dht,
188
+ dh0,
189
+ cu_seqlens,
190
+ scale,
191
+ T,
192
+ B: tl.constexpr,
193
+ H: tl.constexpr,
194
+ K: tl.constexpr,
195
+ V: tl.constexpr,
196
+ BT: tl.constexpr,
197
+ BK: tl.constexpr,
198
+ BV: tl.constexpr,
199
+ USE_G: tl.constexpr,
200
+ USE_G_GAMMA: tl.constexpr,
201
+ IS_VARLEN: tl.constexpr,
202
+ USE_INITIAL_STATE: tl.constexpr,
203
+ USE_FINAL_STATE: tl.constexpr,
204
+ ):
205
+ i_v, i_k, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
206
+ i_n, i_h = i_nh // H, i_nh % H
207
+
208
+ all = B * T
209
+ if IS_VARLEN:
210
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(cu_seqlens + i_n + 1).to(tl.int32)
211
+ T = eos - bos
212
+ else:
213
+ bos, eos = i_n * T, i_n * T + T
214
+ NT = tl.cdiv(T, BT)
215
+ NV = tl.cdiv(V, BV)
216
+
217
+ o_i = tl.arange(0, BT)
218
+ if USE_G_GAMMA:
219
+ b_gamma = tl.load(g_gamma + i_h)
220
+ b_g = b_gamma * (o_i + 1)
221
+ b_g_last = b_gamma * BT
222
+ b_gq = exp(b_g)
223
+ b_gk = exp(b_g_last - b_g)
224
+ b_gn = exp(b_g_last)
225
+
226
+ m_s = o_i[:, None] >= o_i[None, :]
227
+
228
+ q = q + (bos*H + i_h) * K
229
+ k = k + (bos*H + i_h) * K
230
+ v = v + (bos*H + i_h) * V
231
+ do = do + (bos*H + i_h) * V
232
+ dq = dq + (i_v * all + bos).to(tl.int64) * H*K + i_h * K
233
+ dk = dk + (i_v * all + bos).to(tl.int64) * H*K + i_h * K
234
+ dv = dv + (i_k * all + bos).to(tl.int64) * H*V + i_h * V
235
+
236
+ # [BV, BK]
237
+ b_h = tl.zeros([BV, BK], dtype=tl.float32)
238
+ if USE_INITIAL_STATE:
239
+ p_h = tl.make_block_ptr(h0 + i_nh * K*V, (V, K), (1, V), (i_v * BV, i_k * BK), (BV, BK), (0, 1))
240
+ b_h = tl.load(p_h, boundary_check=(0, 1)).to(tl.float32)
241
+
242
+ for i_t in range(0, NT):
243
+ p_q = tl.make_block_ptr(q, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
244
+ p_k = tl.make_block_ptr(k, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
245
+ p_v = tl.make_block_ptr(v, (V, T), (1, H*V), (i_v * BV, i_t * BT), (BV, BT), (0, 1))
246
+ p_do = tl.make_block_ptr(do, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
247
+ p_dq = tl.make_block_ptr(dq, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
248
+
249
+ o_t = i_t * BT + tl.arange(0, BT)
250
+ m_t = o_t < T
251
+ # [BT, BK]
252
+ b_k = tl.load(p_k, boundary_check=(0, 1))
253
+ # [BV, BT]
254
+ b_v = tl.load(p_v, boundary_check=(0, 1))
255
+ # [BT, BV]
256
+ b_do = tl.load(p_do, boundary_check=(0, 1))
257
+ last_idx = min(i_t * BT + BT, T) - 1
258
+
259
+ # [BT, BT]
260
+ b_ds = tl.dot(b_do, b_v) * scale
261
+
262
+ # scalar decay
263
+ if USE_G:
264
+ p_g = g + (bos + o_t) * H + i_h
265
+ b_g = tl.load(p_g, mask=(o_t < T), other=0.)
266
+ b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
267
+
268
+ b_gq = exp(b_g)
269
+ b_gk = exp(b_g_last - b_g)
270
+ b_gn = exp(b_g_last)
271
+
272
+ p_dg = dg + ((i_k * NV + i_v) * all + (bos + o_t)).to(tl.int64) * H + i_h
273
+ # [BT, BT]
274
+ b_gs = tl.where(m_s & m_t, exp(b_g[:, None] - b_g[None, :]), 0)
275
+ b_ds = b_ds * b_gs
276
+ # [BT, BK]
277
+ b_q = tl.load(p_q, boundary_check=(0, 1))
278
+ b_dq = tl.dot(b_ds.to(b_k.dtype), b_k) + tl.dot((b_do * b_gq[:, None] * scale).to(b_k.dtype), b_h.to(b_k.dtype))
279
+ # [BT]
280
+ b_dg_t = tl.sum(b_q * b_dq, 1)
281
+ tl.store(p_dg, b_dg_t.to(p_dg.dtype.element_ty), mask=m_t)
282
+ # [BV, BK]
283
+ b_h = b_h * b_gn + tl.dot(b_v, (b_k * b_gk[:, None]).to(b_k.dtype))
284
+
285
+ elif USE_G_GAMMA:
286
+ b_g_last = b_gamma * min(BT, T - i_t * BT)
287
+ b_gk = exp(b_g_last - b_g)
288
+ b_gn = exp(b_g_last)
289
+
290
+ # [BT, BT]
291
+ b_gs = tl.where(m_s & m_t, exp(b_g[:, None] - b_g[None, :]), 0)
292
+ b_ds = b_ds * b_gs
293
+ # [BT, BK]
294
+ b_q = tl.load(p_q, boundary_check=(0, 1))
295
+ b_dq = tl.dot(b_ds.to(b_k.dtype), b_k) + tl.dot((b_do * b_gq[:, None] * scale).to(b_k.dtype), b_h.to(b_k.dtype))
296
+ # [BV, BK]
297
+ b_h = b_h * b_gn + tl.dot(b_v, (b_k * b_gk[:, None]).to(b_k.dtype))
298
+
299
+ else:
300
+ # [BT, BT]
301
+ b_ds *= m_s & m_t
302
+ # [BT, BK]
303
+ b_dq = tl.dot(b_ds.to(b_k.dtype), b_k) + tl.dot((b_do * scale).to(b_k.dtype), b_h.to(b_k.dtype))
304
+ # [BV, BK]
305
+ b_h += tl.dot(b_v, b_k)
306
+
307
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), boundary_check=(0, 1))
308
+
309
+ # [BK, BV]
310
+ b_dh = tl.zeros([BK, BV], dtype=tl.float32)
311
+ if USE_FINAL_STATE:
312
+ p_dh = tl.make_block_ptr(dht + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
313
+ b_dh += tl.load(p_dh, boundary_check=(0, 1)).to(tl.float32)
314
+
315
+ if USE_G:
316
+ b_dg = tl.zeros([BT], dtype=tl.float32)
317
+ b_dg_last = tl.sum(tl.trans(b_h) * b_dh)
318
+
319
+ # sync threads
320
+ b_h = None
321
+ tl.debug_barrier()
322
+
323
+ for i_t in range(NT - 1, -1, -1):
324
+ p_q = tl.make_block_ptr(q, (K, T), (1, H*K), (i_k * BK, i_t * BT), (BK, BT), (0, 1))
325
+ p_k = tl.make_block_ptr(k, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
326
+ p_v = tl.make_block_ptr(v, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
327
+ p_do = tl.make_block_ptr(do, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
328
+ p_dk = tl.make_block_ptr(dk, (T, K), (H*K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0))
329
+ p_dv = tl.make_block_ptr(dv, (T, V), (H*V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0))
330
+ # [BK, BT]
331
+ b_q = tl.load(p_q, boundary_check=(0, 1))
332
+ # [BT, BK]
333
+ b_k = tl.load(p_k, boundary_check=(0, 1))
334
+ # [BT, BV]
335
+ b_v = tl.load(p_v, boundary_check=(0, 1))
336
+ b_do = tl.load(p_do, boundary_check=(0, 1))
337
+ last_idx = min(i_t * BT + BT, T) - 1
338
+
339
+ o_t = i_t * BT + tl.arange(0, BT)
340
+ m_t = o_t < T
341
+ # [BT, BT]
342
+ b_s = tl.dot(b_k, b_q)
343
+ b_ds = tl.dot(b_v, tl.trans(b_do))
344
+
345
+ if USE_G:
346
+ p_g = g + (bos + o_t) * H + i_h
347
+ p_dg = dg + ((i_k * NV + i_v) * all + (bos + o_t)).to(tl.int64) * H + i_h
348
+ b_g = tl.load(p_g, mask=m_t, other=0.)
349
+ b_g_last = tl.load(g + (bos + last_idx) * H + i_h)
350
+
351
+ b_gq = exp(b_g)
352
+ b_gk = exp(b_g_last - b_g)
353
+ b_gn = exp(b_g_last)
354
+ b_gs = tl.trans(tl.where(m_s & (m_t[:, None] & m_t), exp(b_g[:, None] - b_g[None, :]), 0)) * scale
355
+
356
+ b_s = b_s * b_gs
357
+ b_ds = b_ds * b_gs
358
+
359
+ # [BT, BK]
360
+ b_dk = tl.dot(b_ds.to(b_k.dtype), tl.trans(b_q)) + tl.dot(b_v, tl.trans(b_dh).to(b_v.dtype)) * b_gk[:, None]
361
+
362
+ # [BT]
363
+ b_dg_t = tl.where(m_t, tl.load(p_dg, mask=m_t, other=0.) - tl.sum(b_k * b_dk, 1), 0)
364
+ b_dg_last += tl.sum(b_dg_t, 0)
365
+ b_dg = b_dg_last + b_dg_t - tl.cumsum(b_dg_t, 0)
366
+
367
+ # [BT, BV]
368
+ b_dv = tl.dot(b_s.to(b_do.dtype), b_do) + tl.dot(b_k, b_dh.to(b_k.dtype)) * b_gk[:, None]
369
+ # [BK, BV]
370
+ b_dh = b_dh * b_gn + tl.dot(b_q, (b_do * b_gq[:, None] * scale).to(b_do.dtype))
371
+
372
+ tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty), mask=m_t)
373
+
374
+ elif USE_G_GAMMA:
375
+ b_g_last = b_gamma * min(BT, T - i_t * BT)
376
+ b_gk = exp(b_g_last - b_g)
377
+ b_gn = exp(b_g_last)
378
+ b_gs = tl.trans(tl.where(m_s & (m_t[:, None] & m_t), exp(b_g[:, None] - b_g[None, :]), 0)) * scale
379
+
380
+ b_s = b_s * b_gs
381
+ b_ds = b_ds * b_gs
382
+
383
+ b_dk = tl.dot(b_ds.to(b_k.dtype), tl.trans(b_q)) + tl.dot(b_v, tl.trans(b_dh).to(b_v.dtype)) * b_gk[:, None]
384
+ # [BT, BV]
385
+ b_dv = tl.dot(b_s.to(b_do.dtype), b_do) + tl.dot(b_k, b_dh.to(b_k.dtype)) * b_gk[:, None]
386
+ # [BK, BV]
387
+ b_dh = b_dh * b_gn + tl.dot(b_q, (b_do * b_gq[:, None] * scale).to(b_do.dtype))
388
+
389
+ else:
390
+ mask = tl.trans(m_s & (m_t[:, None] & m_t))
391
+ b_s = tl.where(mask, b_s * scale, 0).to(b_do.dtype)
392
+ b_ds = tl.where(mask, b_ds * scale, 0).to(b_q.dtype)
393
+
394
+ b_dk = tl.dot(b_ds, tl.trans(b_q)) + tl.dot(b_v, tl.trans(b_dh).to(b_v.dtype))
395
+ # [BT, BV]
396
+ b_dv = tl.dot(b_s.to(b_do.dtype), b_do) + tl.dot(b_k, b_dh.to(b_k.dtype))
397
+ # [BK, BV]
398
+ b_dh += tl.dot(b_q, (b_do * scale).to(b_do.dtype))
399
+
400
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), boundary_check=(0, 1))
401
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), boundary_check=(0, 1))
402
+
403
+ if USE_INITIAL_STATE:
404
+ p_dh0 = tl.make_block_ptr(dh0 + i_nh * K*V, (K, V), (V, 1), (i_k * BK, i_v * BV), (BK, BV), (1, 0))
405
+ tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), boundary_check=(0, 1))
406
+
407
+
408
+ def fused_chunk_fwd(
409
+ q: torch.Tensor,
410
+ k: torch.Tensor,
411
+ v: torch.Tensor,
412
+ g: torch.Tensor | None = None,
413
+ g_gamma: torch.Tensor | None = None,
414
+ scale: float | None = None,
415
+ initial_state: torch.Tensor | None = None,
416
+ output_final_state: bool = False,
417
+ cu_seqlens: torch.LongTensor | None = None,
418
+ chunk_size: int = 64,
419
+ ):
420
+ B, T, H, K, V = *q.shape, v.shape[-1]
421
+ BT = chunk_size
422
+ BK = min(max(triton.next_power_of_2(K), 16), 64)
423
+ N = B if cu_seqlens is None else len(cu_seqlens) - 1
424
+ NK = triton.cdiv(K, BK)
425
+
426
+ o = v.new_empty(NK, *v.shape, dtype=torch.float) if NK > 1 else torch.empty_like(v)
427
+ ht = k.new_empty(N, H, K, V, dtype=torch.float) if output_final_state else None
428
+ def grid(meta): return (triton.cdiv(V, meta['BV']), NK, N * H)
429
+ fused_chunk_fwd_kernel[grid](
430
+ q=q,
431
+ k=k,
432
+ v=v,
433
+ g=g,
434
+ g_gamma=g_gamma,
435
+ o=o,
436
+ h0=initial_state,
437
+ ht=ht,
438
+ cu_seqlens=cu_seqlens,
439
+ scale=scale,
440
+ B=B,
441
+ T=T,
442
+ H=H,
443
+ K=K,
444
+ V=V,
445
+ BT=BT,
446
+ BK=BK,
447
+ )
448
+ if NK > 1:
449
+ o = o.sum(0).to(v)
450
+ return o, ht
451
+
452
+
453
+ def fused_chunk_bwd(
454
+ q,
455
+ k,
456
+ v,
457
+ g,
458
+ g_gamma,
459
+ do,
460
+ scale,
461
+ initial_state: torch.Tensor,
462
+ dht: torch.Tensor,
463
+ cu_seqlens: torch.LongTensor | None = None,
464
+ chunk_size: int = 64,
465
+ ):
466
+ B, T, H, K, V = *q.shape, v.shape[-1]
467
+ N = B if cu_seqlens is None else len(cu_seqlens) - 1
468
+ BT = chunk_size
469
+ BK = min(max(triton.next_power_of_2(K), 16), 64)
470
+ BV = min(max(triton.next_power_of_2(V), 16), 64)
471
+ NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
472
+
473
+ dq = q.new_empty(NV, *q.shape, dtype=torch.float) if NV > 1 else torch.empty_like(q)
474
+ dk = k.new_empty(NV, *k.shape, dtype=torch.float) if NV > 1 else torch.empty_like(k)
475
+ dv = v.new_empty(NK, *v.shape, dtype=torch.float) if NK > 1 else torch.empty_like(v)
476
+ dg = g.new_empty(NK*NV, *g.shape, dtype=torch.float) if g is not None else None
477
+ dh0 = torch.empty_like(initial_state) if initial_state is not None else None
478
+
479
+ grid = (NV, NK, N * H)
480
+ fused_chunk_bwd_kernel[grid](
481
+ q=q,
482
+ k=k,
483
+ v=v,
484
+ g=g,
485
+ g_gamma=g_gamma,
486
+ do=do,
487
+ dq=dq,
488
+ dk=dk,
489
+ dv=dv,
490
+ dg=dg,
491
+ h0=initial_state,
492
+ dht=dht,
493
+ dh0=dh0,
494
+ cu_seqlens=cu_seqlens,
495
+ scale=scale,
496
+ T=T,
497
+ B=B,
498
+ H=H,
499
+ K=K,
500
+ V=V,
501
+ BT=BT,
502
+ BK=BK,
503
+ BV=BV,
504
+ )
505
+ dq = dq.sum(0) if NV > 1 else dq
506
+ dk = dk.sum(0) if NV > 1 else dk
507
+ dv = dv.sum(0) if NK > 1 else dv
508
+ if dg is not None:
509
+ dg = dg.sum(0).to(g)
510
+
511
+ return dq, dk, dv, dg, dh0
512
+
513
+
514
+ class FusedChunkFunction(torch.autograd.Function):
515
+
516
+ @staticmethod
517
+ @input_guard
518
+ @autocast_custom_fwd
519
+ def forward(
520
+ ctx,
521
+ q,
522
+ k,
523
+ v,
524
+ g,
525
+ g_gamma,
526
+ scale,
527
+ initial_state,
528
+ output_final_state,
529
+ cu_seqlens,
530
+ ):
531
+ chunk_size = min(64, max(16, triton.next_power_of_2(q.shape[1])))
532
+ g = chunk_local_cumsum(g, chunk_size=chunk_size, cu_seqlens=cu_seqlens) if g is not None else None
533
+ o, ht = fused_chunk_fwd(
534
+ q=q,
535
+ k=k,
536
+ v=v,
537
+ g=g,
538
+ g_gamma=g_gamma,
539
+ scale=scale,
540
+ initial_state=initial_state,
541
+ output_final_state=output_final_state,
542
+ cu_seqlens=cu_seqlens,
543
+ chunk_size=chunk_size,
544
+ )
545
+
546
+ ctx.save_for_backward(q, k, v, g, g_gamma, initial_state)
547
+ ctx.chunk_size = chunk_size
548
+ ctx.scale = scale
549
+ ctx.cu_seqlens = cu_seqlens
550
+ return o.to(q.dtype), ht
551
+
552
+ @staticmethod
553
+ @input_guard
554
+ @autocast_custom_bwd
555
+ def backward(ctx, do, dht=None):
556
+ q, k, v, g, g_gamma, initial_state = ctx.saved_tensors
557
+
558
+ dq, dk, dv, dg, dh0 = fused_chunk_bwd(
559
+ q=q,
560
+ k=k,
561
+ v=v,
562
+ g=g,
563
+ g_gamma=g_gamma,
564
+ do=do,
565
+ scale=ctx.scale,
566
+ initial_state=initial_state,
567
+ dht=dht,
568
+ cu_seqlens=ctx.cu_seqlens,
569
+ chunk_size=ctx.chunk_size,
570
+ )
571
+ if g is not None:
572
+ dg = dg.to(g)
573
+ return dq.to(q), dk.to(k), dv.to(v), dg, None, None, dh0, None, None
574
+
575
+
576
+ def fused_chunk(
577
+ q: torch.Tensor,
578
+ k: torch.Tensor,
579
+ v: torch.Tensor,
580
+ g: torch.Tensor | None = None,
581
+ g_gamma: torch.Tensor | None = None,
582
+ scale: float | None = None,
583
+ initial_state: torch.Tensor | None = None,
584
+ output_final_state: bool = False,
585
+ cu_seqlens: torch.LongTensor | None = None,
586
+ ) -> tuple[torch.Tensor, torch.Tensor]:
587
+ r"""
588
+ Args:
589
+ q (torch.Tensor):
590
+ queries of shape `[B, T, H, K]`.
591
+ k (torch.Tensor):
592
+ keys of shape `[B, T, H, K]`.
593
+ v (torch.Tensor):
594
+ values of shape `[B, T, H, V]`.
595
+ g (torch.Tensor):
596
+ Forget gates of shape `[B, T, H]`.
597
+ Compared to GLA, the gating is head-wise instead of elementwise.
598
+ g_gamma (torch.Tensor):
599
+ Log decay of shape `[H]`.
600
+ Head-wise data-independent decay is used if `g_gamma` is provided.
601
+ Only one of `g` or `g_gamma` should be provided.
602
+ scale (Optional[int]):
603
+ Scale factor for the attention scores.
604
+ If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
605
+ initial_state (Optional[torch.Tensor]):
606
+ Initial state of shape `[N, H, K, V]` for `N` input sequences.
607
+ For equal-length input sequences, `N` equals the batch size `B`.
608
+ Default: `None`.
609
+ output_final_state (Optional[bool]):
610
+ Whether to output the final state of shape `[N, H, K, V]`. Default: `False`.
611
+ cu_seqlens (torch.LongTensor):
612
+ Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
613
+ consistent with the FlashAttention API.
614
+
615
+ Returns:
616
+ o (torch.Tensor):
617
+ Outputs of shape `[B, T, H, V]`.
618
+ final_state (torch.Tensor):
619
+ Final state of shape `[N, H, K, V]` if `output_final_state=True` else `None`.
620
+ """
621
+ if g is not None and g_gamma is not None:
622
+ raise ValueError("Only one of `g` or `g_gamma` should be provided.")
623
+ if scale is None:
624
+ scale = k.shape[-1] ** -0.5
625
+ o, final_state = FusedChunkFunction.apply(
626
+ q,
627
+ k,
628
+ v,
629
+ g,
630
+ g_gamma,
631
+ scale,
632
+ initial_state,
633
+ output_final_state,
634
+ cu_seqlens,
635
+ )
636
+ return o, final_state
code/flash-linear-attention/fla/ops/common/fused_recurrent.py ADDED
@@ -0,0 +1,567 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
2
+
3
+
4
+ import torch
5
+ import triton
6
+ import triton.language as tl
7
+
8
+ from fla.ops.utils.op import exp
9
+ from fla.utils import autocast_custom_bwd, autocast_custom_fwd, autotune_cache_kwargs, input_guard
10
+
11
+
12
+ @triton.heuristics({
13
+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
14
+ 'STORE_FINAL_STATE': lambda args: args['ht'] is not None,
15
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
16
+ })
17
+ @triton.autotune(
18
+ configs=[
19
+ triton.Config({}, num_warps=num_warps)
20
+ for num_warps in [4, 8]
21
+ ],
22
+ key=['BK', 'BV', 'USE_G', 'USE_G_GAMMA', 'USE_GK', 'USE_GV'],
23
+ **autotune_cache_kwargs,
24
+ )
25
+ @triton.jit(do_not_specialize=['B', 'T'])
26
+ def fused_recurrent_fwd_kernel(
27
+ q,
28
+ k,
29
+ v,
30
+ g,
31
+ g_gamma,
32
+ gk,
33
+ gv,
34
+ o,
35
+ h0,
36
+ ht,
37
+ cu_seqlens,
38
+ scale,
39
+ B,
40
+ T,
41
+ H: tl.constexpr,
42
+ K: tl.constexpr,
43
+ V: tl.constexpr,
44
+ BK: tl.constexpr,
45
+ BV: tl.constexpr,
46
+ REVERSE: tl.constexpr,
47
+ USE_G: tl.constexpr,
48
+ USE_G_GAMMA: tl.constexpr,
49
+ USE_GK: tl.constexpr,
50
+ USE_GV: tl.constexpr,
51
+ USE_INITIAL_STATE: tl.constexpr,
52
+ STORE_FINAL_STATE: tl.constexpr,
53
+ IS_VARLEN: tl.constexpr,
54
+ ):
55
+ i_v, i_k, i_nh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
56
+ i_n, i_h = i_nh // H, i_nh % H
57
+
58
+ all = B * T
59
+ if IS_VARLEN:
60
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
61
+ T = eos - bos
62
+ else:
63
+ bos, eos = i_n * T, i_n * T + T
64
+
65
+ o_k = i_k * BK + tl.arange(0, BK)
66
+ o_v = i_v * BV + tl.arange(0, BV)
67
+ p_q = q + (bos + ((T-1) if REVERSE else 0)) * H*K + i_h * K + o_k
68
+ p_k = k + (bos + ((T-1) if REVERSE else 0)) * H*K + i_h * K + o_k
69
+ p_v = v + (bos + ((T-1) if REVERSE else 0)) * H*V + i_h * V + o_v
70
+ p_o = o + ((i_k * all + bos) + ((T-1) if REVERSE else 0)) * H*V + i_h * V + o_v
71
+ if USE_G:
72
+ p_g = g + (bos + ((T-1) if REVERSE else 0)) * H + i_h
73
+ if USE_GK:
74
+ p_gk = gk + (bos + ((T-1) if REVERSE else 0)) * H*K + i_h * K + o_k
75
+ if USE_GV:
76
+ p_gv = gv + (bos + ((T-1) if REVERSE else 0)) * H*V + i_h * V + o_v
77
+ if USE_G_GAMMA:
78
+ b_g_gamma = tl.load(g_gamma + i_h)
79
+
80
+ m_k = o_k < K
81
+ m_v = o_v < V
82
+ m_h = m_k[:, None] & m_v[None, :]
83
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
84
+
85
+ if USE_INITIAL_STATE:
86
+ p_h0 = h0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
87
+ b_h += tl.load(p_h0, mask=m_h, other=0).to(tl.float32)
88
+
89
+ for _ in range(0, T):
90
+ b_q = tl.load(p_q, mask=m_k, other=0).to(tl.float32) * scale
91
+ b_k = tl.load(p_k, mask=m_k, other=0).to(tl.float32)
92
+ b_v = tl.load(p_v, mask=m_v, other=0).to(tl.float32)
93
+ if USE_G:
94
+ b_g = tl.load(p_g).to(tl.float32)
95
+ b_h = b_h * exp(b_g)
96
+ if USE_G_GAMMA:
97
+ b_h = b_h * exp(b_g_gamma)
98
+ if USE_GK:
99
+ b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32)
100
+ b_h = b_h * exp(b_gk[:, None])
101
+ if USE_GV:
102
+ b_gv = tl.load(p_gv, mask=m_v, other=0).to(tl.float32)
103
+ b_h = b_h * exp(b_gv[None, :])
104
+ b_h += b_k[:, None] * b_v[None, :]
105
+ b_o = b_h * b_q[:, None]
106
+ b_o = tl.sum(b_o, axis=0)
107
+ tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=m_v)
108
+ p_q += (-1 if REVERSE else 1) * H*K
109
+ p_k += (-1 if REVERSE else 1) * H*K
110
+ p_v += (-1 if REVERSE else 1) * H*V
111
+ p_o += (-1 if REVERSE else 1) * H*V
112
+ if USE_G:
113
+ p_g += (-1 if REVERSE else 1) * H
114
+ if USE_GK:
115
+ p_gk += (-1 if REVERSE else 1) * H*K
116
+ if USE_GV:
117
+ p_gv += (-1 if REVERSE else 1) * H*V
118
+
119
+ if STORE_FINAL_STATE:
120
+ p_ht = ht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
121
+ tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=m_h)
122
+
123
+
124
+ @triton.heuristics({
125
+ 'USE_INITIAL_STATE': lambda args: args['h0'] is not None,
126
+ 'STORE_INITIAL_STATE_GRADIENT': lambda args: args['dh0'] is not None,
127
+ 'USE_FINAL_STATE_GRADIENT': lambda args: args['dht'] is not None,
128
+ 'IS_VARLEN': lambda args: args['cu_seqlens'] is not None,
129
+ })
130
+ @triton.autotune(
131
+ configs=[
132
+ triton.Config({}, num_warps=num_warps)
133
+ for num_warps in [4]
134
+ ],
135
+ key=['BK', 'BV', 'USE_G', 'USE_G_GAMMA', 'USE_GK', 'USE_GV'],
136
+ **autotune_cache_kwargs,
137
+ )
138
+ @triton.jit(do_not_specialize=['B', 'T'])
139
+ def fused_recurrent_bwd_kernel(
140
+ q,
141
+ k,
142
+ v,
143
+ g,
144
+ g_gamma,
145
+ gk,
146
+ gv,
147
+ o,
148
+ h0,
149
+ do,
150
+ dq,
151
+ dk,
152
+ dv,
153
+ dg,
154
+ dgk,
155
+ dgv,
156
+ dht,
157
+ dh0,
158
+ cu_seqlens,
159
+ scale,
160
+ B,
161
+ T,
162
+ H: tl.constexpr,
163
+ K: tl.constexpr,
164
+ V: tl.constexpr,
165
+ BK: tl.constexpr,
166
+ BV: tl.constexpr,
167
+ REVERSE: tl.constexpr,
168
+ USE_G: tl.constexpr,
169
+ USE_G_GAMMA: tl.constexpr,
170
+ USE_GK: tl.constexpr,
171
+ USE_GV: tl.constexpr,
172
+ USE_INITIAL_STATE: tl.constexpr,
173
+ STORE_INITIAL_STATE_GRADIENT: tl.constexpr,
174
+ USE_FINAL_STATE_GRADIENT: tl.constexpr,
175
+ IS_VARLEN: tl.constexpr,
176
+ ):
177
+ i_v, i_k, i_nh = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
178
+ i_n, i_h = i_nh // H, i_nh % H
179
+
180
+ all = B * T
181
+ if IS_VARLEN:
182
+ bos, eos = tl.load(cu_seqlens + i_n).to(tl.int64), tl.load(cu_seqlens + i_n + 1).to(tl.int64)
183
+ T = eos - bos
184
+ else:
185
+ bos, eos = i_n * T, i_n * T + T
186
+ NV = tl.cdiv(V, BV)
187
+
188
+ o_k = i_k * BK + tl.arange(0, BK)
189
+ o_v = i_v * BV + tl.arange(0, BV)
190
+ m_k = o_k < K
191
+ m_v = o_v < V
192
+ m_h = m_k[:, None] & m_v[None, :]
193
+
194
+ p_k = k + (bos + ((T-1) if REVERSE else 0)) * H*K + i_h * K + o_k
195
+ p_v = v + (bos + ((T-1) if REVERSE else 0)) * H*V + i_h * V + o_v
196
+ p_do = do + (bos + ((T-1) if REVERSE else 0)) * H*V + i_h * V + o_v
197
+ p_dq = dq + ((i_v * all + bos) + ((T-1) if REVERSE else 0)) * H*K + i_h * K + o_k
198
+ if USE_G:
199
+ p_g = g + (bos + ((T-1) if REVERSE else 0)) * H + i_h
200
+ if USE_GK:
201
+ p_gk = gk + (bos + ((T-1) if REVERSE else 0)) * H*K + i_h * K + o_k
202
+ if USE_GV:
203
+ p_gv = gv + (bos + ((T-1) if REVERSE else 0)) * H*V + i_h * V + o_v
204
+ if USE_G_GAMMA:
205
+ b_g_gamma = tl.load(g_gamma + i_h)
206
+
207
+ b_h = tl.zeros([BK, BV], dtype=tl.float32)
208
+ if USE_INITIAL_STATE:
209
+ p_h0 = h0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
210
+ b_h += tl.load(p_h0, mask=m_h, other=0).to(tl.float32)
211
+
212
+ for _ in range(0, T):
213
+ b_k = tl.load(p_k, mask=m_k, other=0).to(tl.float32)
214
+ b_v = tl.load(p_v, mask=m_v, other=0).to(tl.float32)
215
+ b_do = tl.load(p_do, mask=m_v, other=0).to(tl.float32)
216
+ if USE_G:
217
+ b_g = tl.load(p_g).to(tl.float32)
218
+ b_h = b_h * exp(b_g)
219
+ if USE_G_GAMMA:
220
+ b_h = b_h * exp(b_g_gamma)
221
+ if USE_GK:
222
+ b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32)
223
+ b_h = b_h * exp(b_gk[:, None])
224
+ if USE_GV:
225
+ b_gv = tl.load(p_gv, mask=m_v, other=0).to(tl.float32)
226
+ b_h = b_h * exp(b_gv[None, :])
227
+ b_h += b_k[:, None] * b_v[None, :]
228
+ b_dq = b_h * b_do[None, :]
229
+ b_dq = tl.sum(b_dq, axis=1) * scale
230
+ tl.store(p_dq, b_dq.to(p_dq.dtype.element_ty), mask=m_k)
231
+
232
+ p_k += (-1 if REVERSE else 1) * H*K
233
+ p_v += (-1 if REVERSE else 1) * H*V
234
+ p_do += (-1 if REVERSE else 1) * H*V
235
+ p_dq += (-1 if REVERSE else 1) * H*K
236
+ if USE_G:
237
+ p_g += (-1 if REVERSE else 1) * H
238
+ if USE_GK:
239
+ p_gk += (-1 if REVERSE else 1) * H*K
240
+ if USE_GV:
241
+ p_gv += (-1 if REVERSE else 1) * H*V
242
+
243
+ # sync threads
244
+ tl.debug_barrier()
245
+
246
+ p_q = q + (bos + ((T - 1) if not REVERSE else 0)) * H*K + i_h * K + o_k
247
+ p_k = k + (bos + ((T - 1) if not REVERSE else 0)) * H*K + i_h * K + o_k
248
+ p_v = v + (bos + ((T - 1) if not REVERSE else 0)) * H*V + i_h * V + o_v
249
+
250
+ p_do = do + (bos + ((T - 1) if not REVERSE else 0)) * H*V + i_h * V + o_v
251
+ p_dq = dq + ((i_v * all + bos) + ((T - 1) if not REVERSE else 0)) * H*K + i_h * K + o_k
252
+ p_dk = dk + ((i_v * all + bos) + ((T - 1) if not REVERSE else 0)) * H*K + i_h * K + o_k
253
+ p_dv = dv + ((i_k * all + bos) + ((T - 1) if not REVERSE else 0)) * H*V + i_h * V + o_v
254
+ if USE_G:
255
+ p_g = g + (bos + ((T - 1) if not REVERSE else 0)) * H + i_h
256
+ p_dg = dg + ((i_k * NV + i_v) * all + bos + ((T - 1) if not REVERSE else 0)) * H + i_h
257
+ if USE_GK:
258
+ p_gk = gk + (bos + ((T - 1) if not REVERSE else 0)) * H*K + i_h * K + o_k
259
+ p_dgk = dgk + ((i_v * all + bos) + ((T - 1) if not REVERSE else 0)) * H*K + i_h * K + o_k
260
+ if USE_GV:
261
+ p_o = o + (bos + ((T - 1) if not REVERSE else 0)) * H*V + i_h * V + o_v
262
+ p_gv = gv + (bos + ((T - 1) if not REVERSE else 0)) * H*V + i_h * V + o_v
263
+ p_dgv = dgv + ((i_k * all + bos) + ((T - 1) if not REVERSE else 0)) * H*V + i_h * V + o_v
264
+
265
+ b_dh = tl.zeros([BK, BV], dtype=tl.float32)
266
+ if USE_FINAL_STATE_GRADIENT:
267
+ p_dht = dht + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
268
+ b_dh += tl.load(p_dht, mask=m_h, other=0).to(tl.float32)
269
+
270
+ if USE_G:
271
+ b_dg = tl.sum(b_h * b_dh)
272
+ if USE_GK:
273
+ b_dgk = tl.sum(b_h * b_dh, 1)
274
+ if USE_GV:
275
+ b_dgv = tl.sum(b_h * b_dh, 0)
276
+
277
+ for _ in range(T):
278
+ b_q = tl.load(p_q, mask=m_k, other=0).to(tl.float32)
279
+ b_k = tl.load(p_k, mask=m_k, other=0).to(tl.float32)
280
+ b_v = tl.load(p_v, mask=m_v, other=0).to(tl.float32)
281
+ b_do = tl.load(p_do, mask=m_v, other=0).to(tl.float32)
282
+ b_dh += (b_q * scale)[:, None] * b_do[None, :]
283
+ b_dk = tl.sum(b_dh * b_v[None, :], axis=1)
284
+ b_dv = tl.sum(b_dh * b_k[:, None], axis=0)
285
+
286
+ if USE_G:
287
+ b_g = tl.load(p_g).to(tl.float32)
288
+ b_dq = tl.load(p_dq, mask=m_k, other=0).to(tl.float32)
289
+ b_dg += tl.sum(b_q * b_dq - b_k * b_dk)
290
+ b_dh *= exp(b_g)
291
+ tl.store(p_dg, b_dg.to(p_dg.dtype.element_ty))
292
+ if USE_G_GAMMA:
293
+ b_dh *= exp(b_g_gamma)
294
+ if USE_GK:
295
+ b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32)
296
+ b_dq = tl.load(p_dq, mask=m_k, other=0).to(tl.float32)
297
+ b_dgk += b_q * b_dq - b_k * b_dk
298
+ b_dh *= exp(b_gk)[:, None]
299
+ tl.store(p_dgk, b_dgk.to(p_dgk.dtype.element_ty), mask=m_k)
300
+ if USE_GV:
301
+ b_o = tl.load(p_o, mask=m_v, other=0).to(tl.float32)
302
+ b_gv = tl.load(p_gv, mask=m_v, other=0).to(tl.float32)
303
+ if i_k == 0:
304
+ b_dgv += b_o * b_do
305
+ b_dgv -= b_v * b_dv
306
+ b_dh *= exp(b_gv)[None, :]
307
+ tl.store(p_dgv, b_dgv.to(p_dgv.dtype.element_ty), mask=m_v)
308
+
309
+ tl.store(p_dk, b_dk.to(p_dk.dtype.element_ty), mask=m_k)
310
+ tl.store(p_dv, b_dv.to(p_dv.dtype.element_ty), mask=m_v)
311
+
312
+ p_q += (1 if REVERSE else -1) * H*K
313
+ p_k += (1 if REVERSE else -1) * H*K
314
+ p_v += (1 if REVERSE else -1) * H*V
315
+
316
+ p_do += (1 if REVERSE else -1) * H*V
317
+ p_dq += (1 if REVERSE else -1) * H*K
318
+ p_dk += (1 if REVERSE else -1) * H*K
319
+ p_dv += (1 if REVERSE else -1) * H*V
320
+ if USE_G:
321
+ p_g += (1 if REVERSE else -1) * H
322
+ p_dg += (1 if REVERSE else -1) * H
323
+ if USE_GK:
324
+ p_gk += (1 if REVERSE else -1) * H*K
325
+ p_dgk += (1 if REVERSE else -1) * H*K
326
+ if USE_GV:
327
+ p_o += (1 if REVERSE else -1) * H*V
328
+ p_gv += (1 if REVERSE else -1) * H*V
329
+ p_dgv += (1 if REVERSE else -1) * H*V
330
+
331
+ if STORE_INITIAL_STATE_GRADIENT:
332
+ p_dh0 = dh0 + i_nh * K*V + o_k[:, None] * V + o_v[None, :]
333
+ tl.store(p_dh0, b_dh.to(p_dh0.dtype.element_ty), mask=m_h)
334
+
335
+
336
+ def fused_recurrent_fwd(
337
+ q: torch.Tensor,
338
+ k: torch.Tensor,
339
+ v: torch.Tensor,
340
+ g: torch.Tensor | None = None,
341
+ g_gamma: torch.Tensor | None = None,
342
+ gk: torch.Tensor | None = None,
343
+ gv: torch.Tensor | None = None,
344
+ scale: float | None = None,
345
+ initial_state: torch.Tensor | None = None,
346
+ output_final_state: bool = False,
347
+ reverse: bool = False,
348
+ cu_seqlens: torch.LongTensor | None = None,
349
+ ):
350
+ B, T, H, K, V = *k.shape, v.shape[-1]
351
+ N = B if cu_seqlens is None else len(cu_seqlens) - 1
352
+ BK, BV = min(triton.next_power_of_2(K), 64), min(triton.next_power_of_2(V), 64)
353
+ NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
354
+
355
+ h0 = initial_state
356
+ ht = q.new_empty(N, H, K, V, dtype=torch.float32) if output_final_state else None
357
+ o = q.new_empty(NK, *v.shape, dtype=torch.float32)
358
+
359
+ grid = (NV, NK, N * H)
360
+ fused_recurrent_fwd_kernel[grid](
361
+ q=q,
362
+ k=k,
363
+ v=v,
364
+ g=g,
365
+ g_gamma=g_gamma,
366
+ gk=gk,
367
+ gv=gv,
368
+ o=o,
369
+ h0=h0,
370
+ ht=ht,
371
+ cu_seqlens=cu_seqlens,
372
+ scale=scale,
373
+ T=T,
374
+ B=B,
375
+ H=H,
376
+ K=K,
377
+ V=V,
378
+ BK=BK,
379
+ BV=BV,
380
+ USE_G=g is not None,
381
+ USE_G_GAMMA=g_gamma is not None,
382
+ USE_GK=gk is not None,
383
+ USE_GV=gv is not None,
384
+ REVERSE=reverse,
385
+ )
386
+ o = o.sum(0)
387
+ return o, ht
388
+
389
+
390
+ def fused_recurrent_bwd(
391
+ q: torch.Tensor,
392
+ k: torch.Tensor,
393
+ v: torch.Tensor,
394
+ g: torch.Tensor | None = None,
395
+ g_gamma: torch.Tensor | None = None,
396
+ gk: torch.Tensor | None = None,
397
+ gv: torch.Tensor | None = None,
398
+ o: torch.Tensor | None = None,
399
+ do: torch.Tensor | None = None,
400
+ dht: torch.Tensor | None = None,
401
+ scale: float | None = None,
402
+ initial_state: torch.Tensor | None = None,
403
+ reverse: bool = False,
404
+ cu_seqlens: torch.LongTensor | None = None,
405
+ ):
406
+ B, T, H, K, V = *k.shape, v.shape[-1]
407
+ N = B if cu_seqlens is None else len(cu_seqlens) - 1
408
+
409
+ BK, BV = min(triton.next_power_of_2(K), 64), min(triton.next_power_of_2(V), 64)
410
+ NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
411
+
412
+ h0 = initial_state
413
+ dq = q.new_empty(NV, *q.shape, dtype=torch.float32)
414
+ dk = q.new_empty(NV, *k.shape, dtype=torch.float32)
415
+ dv = q.new_empty(NK, *v.shape, dtype=torch.float32)
416
+ dh0 = torch.empty_like(h0) if h0 is not None else None
417
+
418
+ dg, dgk, dgv = None, None, None
419
+ if g is not None:
420
+ dg = g.new_empty(NK*NV, *g.shape, dtype=torch.float32)
421
+ if gk is not None:
422
+ dgk = gk.new_empty(NV, *gk.shape, dtype=torch.float32)
423
+ if gv is not None:
424
+ dgv = gv.new_empty(NK, *gv.shape, dtype=torch.float32)
425
+
426
+ grid = (NV, NK, N * H)
427
+ fused_recurrent_bwd_kernel[grid](
428
+ q=q,
429
+ k=k,
430
+ v=v,
431
+ g=g,
432
+ g_gamma=g_gamma,
433
+ gk=gk,
434
+ gv=gv,
435
+ o=o,
436
+ h0=h0,
437
+ do=do,
438
+ dq=dq,
439
+ dk=dk,
440
+ dv=dv,
441
+ dg=dg,
442
+ dgk=dgk,
443
+ dgv=dgv,
444
+ dht=dht,
445
+ dh0=dh0,
446
+ cu_seqlens=cu_seqlens,
447
+ scale=scale,
448
+ B=B,
449
+ T=T,
450
+ H=H,
451
+ K=K,
452
+ V=V,
453
+ BK=BK,
454
+ BV=BV,
455
+ USE_G=g is not None,
456
+ USE_G_GAMMA=g_gamma is not None,
457
+ USE_GK=gk is not None,
458
+ USE_GV=gv is not None,
459
+ REVERSE=reverse,
460
+ )
461
+ dq = dq.sum(0)
462
+ dk = dk.sum(0)
463
+ dv = dv.sum(0)
464
+ if g is not None:
465
+ dg = dg.sum(0).to(g)
466
+ if gk is not None:
467
+ dgk = dgk.sum(0).to(gk)
468
+ if gv is not None:
469
+ dgv = dgv.sum(0).to(gv)
470
+
471
+ return dq, dk, dv, dg, dgk, dgv, dh0
472
+
473
+
474
+ class FusedRecurrentFunction(torch.autograd.Function):
475
+
476
+ @staticmethod
477
+ @input_guard
478
+ @autocast_custom_fwd
479
+ def forward(
480
+ ctx,
481
+ q: torch.Tensor,
482
+ k: torch.Tensor,
483
+ v: torch.Tensor,
484
+ g: torch.Tensor | None = None,
485
+ g_gamma: torch.Tensor | None = None,
486
+ gk: torch.Tensor | None = None,
487
+ gv: torch.Tensor | None = None,
488
+ scale: float | None = None,
489
+ initial_state: torch.Tensor | None = None,
490
+ output_final_state: bool = False,
491
+ reverse: bool = False,
492
+ cu_seqlens: torch.LongTensor | None = None,
493
+ ):
494
+ o, ht = fused_recurrent_fwd(
495
+ q=q,
496
+ k=k,
497
+ v=v,
498
+ g=g,
499
+ g_gamma=g_gamma,
500
+ gk=gk,
501
+ gv=gv,
502
+ scale=scale,
503
+ initial_state=initial_state,
504
+ output_final_state=output_final_state,
505
+ reverse=reverse,
506
+ cu_seqlens=cu_seqlens,
507
+ )
508
+ ctx.save_for_backward(q, k, v, g, g_gamma, gk, gv, initial_state, o)
509
+ ctx.scale = scale
510
+ ctx.reverse = reverse
511
+ ctx.cu_seqlens = cu_seqlens
512
+ return o.to(q.dtype), ht
513
+
514
+ @staticmethod
515
+ @input_guard
516
+ @autocast_custom_bwd
517
+ def backward(ctx, do, dht):
518
+ q, k, v, g, g_gamma, gk, gv, initial_state, o = ctx.saved_tensors
519
+ dq, dk, dv, dg, dgk, dgv, dh0 = fused_recurrent_bwd(
520
+ q=q,
521
+ k=k,
522
+ v=v,
523
+ g=g,
524
+ g_gamma=g_gamma,
525
+ gk=gk,
526
+ gv=gv,
527
+ o=o,
528
+ do=do,
529
+ dht=dht,
530
+ scale=ctx.scale,
531
+ initial_state=initial_state,
532
+ reverse=ctx.reverse,
533
+ cu_seqlens=ctx.cu_seqlens,
534
+ )
535
+ return dq.to(q.dtype), dk.to(k.dtype), dv.to(v.dtype), dg, None, dgk, dgv, None, dh0, None, None, None
536
+
537
+
538
+ def fused_recurrent(
539
+ q: torch.Tensor,
540
+ k: torch.Tensor,
541
+ v: torch.Tensor,
542
+ g: torch.Tensor | None = None,
543
+ g_gamma: torch.Tensor | None = None,
544
+ gk: torch.Tensor | None = None,
545
+ gv: torch.Tensor | None = None,
546
+ scale: float | None = None,
547
+ initial_state: torch.Tensor | None = None,
548
+ output_final_state: bool = False,
549
+ reverse: bool = False,
550
+ cu_seqlens: torch.LongTensor | None = None,
551
+ ):
552
+ if scale is None:
553
+ scale = k.shape[-1] ** -0.5
554
+ return FusedRecurrentFunction.apply(
555
+ q,
556
+ k,
557
+ v,
558
+ g,
559
+ g_gamma,
560
+ gk,
561
+ gv,
562
+ scale,
563
+ initial_state,
564
+ output_final_state,
565
+ reverse,
566
+ cu_seqlens,
567
+ )
code/flash-linear-attention/fla/ops/delta_rule/README.md ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Chunkwise-form Parallelism of DeltaNet
2
+
3
+ This section expands on the formulation presented in Appendix B of the DeltaNet paper.[^1]
4
+
5
+ To reduce notational clutter, we focus on the first chunk, denoting $\mathbf{S}^r=\mathbf{S}_{[1]}^r$. By partially expanding the recurrence, we have:
6
+ ```math
7
+ \begin{equation}
8
+ \begin{aligned}
9
+ \mathbf{S}^r &= \underbrace{\left(\prod_{i=1}^r \mathbf{I} - \beta^i \bf{k}^i \bf{k}^{i\top} \right)}_{:= \mathbf{P}^r} \cdot\mathbf{S}^{0} + \overbrace{\sum_{i=1}^{r} \underbrace{\left(\prod_{j=i+1}^r \mathbf{I} - \beta^j \bf{k}^j \bf{k}^{j\top} \right)}_{:= \mathbf{P}_{i+1}^r}\beta^i \bf{k}^i\bf{v}^{i\top}}^{:=\mathbf{H}^r} \\
10
+ &=\mathbf{P}^r \cdot \mathbf{S}^{0} + \mathbf{H}^r
11
+ \end{aligned}
12
+ \end{equation}
13
+ ```
14
+
15
+ where $\mathbf{P}_i^r$ involves cumulative products of generalized Householder matrices.
16
+ We abbreviate $\mathbf{P}_1^r$ as $\mathbf{P}^r$.
17
+ This can be optimized using the classical WY representation:
18
+ ```math
19
+ \begin{equation}
20
+ \mathbf{P}^{r} = \mathbf{I} - \sum_{i=1}^{r}\bf{k}^i\bf{w}^{i\top} \in \mathbb{R}^{d_k \times d_k};\qquad
21
+ \bf{w}^r = \beta^r \left(\bf{k}^r - \sum_{i=1}^{r-1} \left(\bf{k}^{r\top}\bf{k}^i \right)\bf{w}^i \right) \in \mathbb{R}^{d_k}
22
+ \end{equation}
23
+ ```
24
+
25
+ We prove this by induction:
26
+ ```math
27
+ \begin{align*}
28
+ \mathbf{P}^{r} &= \prod_{i=1}^r \mathbf{I} - \beta^i \bf{k}^i \bf{k}^{i\top} \\
29
+ &= \left(\mathbf{I} - \beta^r \bf{k}^r \bf{k}^{r\top}\right)\mathbf{P}^{r-1} \\
30
+ &= \left(\mathbf{I} - \beta^r \bf{k}^r \bf{k}^{r\top}\right)\left(\mathbf{I} - \sum_{i=1}^{r-1}\bf{k}^i\bf{w}^{i\top}\right) \\
31
+ &= \mathbf{I} - \sum_{i=1}^{r-1}\bf{k}^i\bf{w}^{i\top} - \beta^r \bf{k}^r \bf{k}^{r\top} + \beta^r\bf{k}^r \bf{k}^{r\top} \left(\sum_{i=1}^{r-1}\bf{k}^i\bf{w}^{i\top}\right) \\
32
+ &= \mathbf{I} - \sum_{i=1}^{r-1}\bf{k}^i\bf{w}^{i\top} - \beta^r \bf{k}^r \left(\bf{k}^{r} - \left(\sum_{i=1}^{r-1}\left(\bf{k}^{r\top} \bf{k}^i\right)\bf{w}^{i}\right) \right)^\top \\
33
+ &= \mathbf{I} - \sum_{i=1}^{r}\bf{k}^i\bf{w}^{i\top}
34
+ \end{align*}
35
+ ```
36
+
37
+ Similarly, $\mathbf{H}^r$ can be represented as:
38
+ ```math
39
+ \begin{equation}
40
+ \mathbf{H}^{r} = \sum_{i=1}^{r} \bf{k}^i \bf{u}^{i\top} \in \mathbb{R}^{d_k \times d_v};\qquad \bf{u}^r = \beta^r \left(\bf{v}^r - \sum_{i=1}^{r-1} \left(\bf{k}^{r\top}\bf{k}^i\right) \bf{u}^i \right)\in \mathbb{R}^{d_v}
41
+ \end{equation}
42
+ ```
43
+
44
+ This can also be proven by induction:
45
+ ```math
46
+ \begin{align*}
47
+ \mathbf{H}^{r} &= \sum_{i=1}^{r} \mathbf{P}_{i+1}^r \beta^i \bf{k}^i \bf{v}^{i\top}\\
48
+ &= \left(\mathbf{I} - \beta^r \bf{k}^r \bf{k}^{r\top}\right) \mathbf{H}^{r-1} + \beta^r \bf{k}^r \bf{v}^{r\top}\\
49
+ &= \sum_{i=1}^{r-1}\bf{k}^i \bf{u}^{i\top} - \beta^r \bf{k}^r \bf{k}^{r\top} \sum_{i=1}^{r-1}\bf{k}^i \bf{u}^{i\top} +\beta^r \bf{k}^r \bf{v}^{r\top}\\
50
+ &= \sum_{i=1}^{r-1}\bf{k}^i \bf{u}^{i\top} + \bf{k}^r \left(\beta^r \bf{v}^{r\top}-\beta^r \bf{k}^{r\top} \sum_{i=1}^{r-1}\bf{k}^i \bf{u}^{i\top}\right) \\
51
+ &= \sum_{i=1}^{r-1}\bf{k}^i \bf{u}^{i\top} + \bf{k}^r \beta^r\left(\bf{v}^{r}-\sum_{i=1}^{r-1}\left(\bf{k}^{r\top}\bf{k}^{i}\right)\bf{u}^{i} \right)^\top \\
52
+ &=\sum_{i=1}^{r} \bf{k}^i \bf{u}^{i\top}
53
+ \end{align*}
54
+ ```
55
+
56
+ In matrix form, $\mathbf{P}$ and $\mathbf{H}$ can be written as:
57
+ ```math
58
+ \begin{equation}
59
+ \mathbf{P}=\mathbf{I}-\mathbf{K}^\top\mathbf{W} \in \mathbb{R}^{d_k \times d_k}, \qquad\mathbf{H}=\mathbf{K}^\top\mathbf{U} \in \mathbb{R}^{d_k\times d_v}
60
+ \end{equation}
61
+ ```
62
+
63
+ Now we can derive the matrix form of $\mathbf{W}$ and $\mathbf{U}$:
64
+ ```math
65
+ \begin{align*}
66
+ \mathbf{W} &= \mathrm{diag}(\beta) \mathbf{K} - \mathrm{tril}(\mathrm{diag}(\beta) \mathbf{K}\mathbf{K}^\top, -1)\mathbf{W}\\
67
+ \left(\mathbf{I} + \mathrm{tril}(\mathrm{diag}(\beta) \mathbf{K}\mathbf{K}^\top, -1)\right) \mathbf{W} &= \mathrm{diag}(\beta) \mathbf{K}
68
+ \end{align*}
69
+ ```
70
+ A similar process holds for $\mathbf{U}$. We can further write $\mathbf{W}$ and $\mathbf{U}$ in matrix form:
71
+ ```math
72
+ \begin{align*}
73
+ \mathbf{T} &= \left(\mathbf{I} + \mathrm{tril}\left(\mathrm{diag}(\beta)\mathbf{K} \mathbf{K}^\top,-1\right)\right)^{-1}\mathrm{diag}\left(\beta\right)\in \mathbb{R}^{C \times C}\\
74
+ \mathbf{W} &= \mathbf{T} \mathbf{K}\in \mathbb{R}^{C \times d_k}\\
75
+ \mathbf{U} &= \mathbf{T}\mathbf{V}\in \mathbb{R}^{C \times d_v}
76
+ \end{align*}
77
+ ```
78
+
79
+ Substituting these back into the original equations yields a hardware-efficient chunkwise algorithm for DeltaNet that leverages matrix multiplications, enabling tensor core based GPU optimization:
80
+ ```math
81
+ \begin{equation}
82
+ \begin{aligned}
83
+ \mathbf{S} &= \mathbf{P}\cdot\mathbf{S}^0 + \mathbf{H} \\
84
+ &= \mathbf{S}^0 + \mathbf{K}^\top (\mathbf{U} -\mathbf{W} \mathbf{S}^0) \in \mathbb{R}^{d_k \times d_v}\\
85
+ \mathbf{O} &= \mathbf{Q} \mathbf{S}^0 + (\mathbf{Q} \mathbf{K}^{\top} \odot \mathbf{M}) \left(\mathbf{U} - \mathbf{W} \mathbf{S}^0\right) \in \mathbb{R}^{C \times d_v}
86
+ \end{aligned}
87
+ \end{equation}
88
+ ```
89
+
90
+ [^1]: https://arxiv.org/abs/2406.06484
code/flash-linear-attention/fla/ops/delta_rule/__init__.py ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+
2
+ from .chunk import chunk_delta_rule
3
+ from .fused_chunk import fused_chunk_delta_rule
4
+ from .fused_recurrent import fused_recurrent_delta_rule
5
+
6
+ __all__ = [
7
+ 'fused_chunk_delta_rule',
8
+ 'fused_recurrent_delta_rule',
9
+ 'chunk_delta_rule',
10
+ ]