Add Echo-Memory codebase used for this run (CC BY 4.0, JD Echo Team) (part 3)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- code/flash-linear-attention/fla/models/sse/modeling_sse.py +437 -0
- code/flash-linear-attention/fla/models/transformer/__init__.py +12 -0
- code/flash-linear-attention/fla/models/transformer/configuration_transformer.py +85 -0
- code/flash-linear-attention/fla/models/transformer/modeling_transformer.py +356 -0
- code/flash-linear-attention/fla/models/utils.py +471 -0
- code/flash-linear-attention/fla/modules/__init__.py +32 -0
- code/flash-linear-attention/fla/modules/activations.py +555 -0
- code/flash-linear-attention/fla/modules/convolution.py +1167 -0
- code/flash-linear-attention/fla/modules/feature_map.py +298 -0
- code/flash-linear-attention/fla/modules/fused_bitlinear.py +633 -0
- code/flash-linear-attention/fla/modules/fused_cross_entropy.py +418 -0
- code/flash-linear-attention/fla/modules/fused_kl_div.py +322 -0
- code/flash-linear-attention/fla/modules/fused_linear_cross_entropy.py +630 -0
- code/flash-linear-attention/fla/modules/fused_norm_gate.py +1245 -0
- code/flash-linear-attention/fla/modules/grpo.py +412 -0
- code/flash-linear-attention/fla/modules/l2norm.py +287 -0
- code/flash-linear-attention/fla/modules/l2warp.py +37 -0
- code/flash-linear-attention/fla/modules/layernorm.py +1444 -0
- code/flash-linear-attention/fla/modules/layernorm_gated.py +527 -0
- code/flash-linear-attention/fla/modules/mlp.py +144 -0
- code/flash-linear-attention/fla/modules/parallel.py +53 -0
- code/flash-linear-attention/fla/modules/rotary.py +499 -0
- code/flash-linear-attention/fla/modules/token_shift.py +545 -0
- code/flash-linear-attention/fla/ops/__init__.py +54 -0
- code/flash-linear-attention/fla/ops/abc/__init__.py +6 -0
- code/flash-linear-attention/fla/ops/abc/chunk.py +1115 -0
- code/flash-linear-attention/fla/ops/abc/naive.py +94 -0
- code/flash-linear-attention/fla/ops/attn/__init__.py +6 -0
- code/flash-linear-attention/fla/ops/attn/decoding.py +181 -0
- code/flash-linear-attention/fla/ops/attn/parallel.py +728 -0
- code/flash-linear-attention/fla/ops/based/__init__.py +8 -0
- code/flash-linear-attention/fla/ops/based/fused_chunk.py +371 -0
- code/flash-linear-attention/fla/ops/based/naive.py +70 -0
- code/flash-linear-attention/fla/ops/based/parallel.py +406 -0
- code/flash-linear-attention/fla/ops/comba/__init__.py +7 -0
- code/flash-linear-attention/fla/ops/comba/chunk.py +340 -0
- code/flash-linear-attention/fla/ops/comba/fused_recurrent.py +330 -0
- code/flash-linear-attention/fla/ops/comba/utils.py +174 -0
- code/flash-linear-attention/fla/ops/comba/wy_fast.py +424 -0
- code/flash-linear-attention/fla/ops/common/__init__.py +0 -0
- code/flash-linear-attention/fla/ops/common/chunk_delta_h.py +533 -0
- code/flash-linear-attention/fla/ops/common/chunk_h.py +394 -0
- code/flash-linear-attention/fla/ops/common/chunk_h_parallel.py +554 -0
- code/flash-linear-attention/fla/ops/common/chunk_h_split.py +599 -0
- code/flash-linear-attention/fla/ops/common/chunk_o.py +689 -0
- code/flash-linear-attention/fla/ops/common/chunk_scaled_dot_kkt.py +124 -0
- code/flash-linear-attention/fla/ops/common/fused_chunk.py +636 -0
- code/flash-linear-attention/fla/ops/common/fused_recurrent.py +567 -0
- code/flash-linear-attention/fla/ops/delta_rule/README.md +90 -0
- 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 |
+
]
|