Skip to content

Commit a4fb2d7

Browse files
authored
[KDA] Add Paddle training compatibility (#2)
Co-authored-by: huangjiyi <huangjiyi@users.noreply.github.com>
1 parent 0f0f0c9 commit a4fb2d7

23 files changed

Lines changed: 917 additions & 287 deletions

fla/__init__.py

Lines changed: 2 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -5,40 +5,11 @@
55
# For a list of all contributors, visit:
66
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
77

8-
import importlib
98
from pkgutil import extend_path
109

1110
__path__ = extend_path(__path__, __name__)
1211
__version__ = "0.5.2"
1312

14-
__all__: list[str] = []
13+
from fla import modules, ops # noqa: E402
1514

16-
17-
def _import_optional_public_module(module_name: str):
18-
try:
19-
return importlib.import_module(module_name)
20-
except ModuleNotFoundError as exc:
21-
missing = exc.name
22-
# The extension package is optional. Treat its absence, or the absence
23-
# of an external runtime dependency, as the extension being unavailable.
24-
if missing == module_name or (missing is not None and missing.split('.', 1)[0] != 'fla'):
25-
return None
26-
raise
27-
28-
29-
def _export_public_api(module) -> None:
30-
globals()[module.__name__.rsplit('.', maxsplit=1)[-1]] = module
31-
for name in module.__all__:
32-
if name.endswith('Config'):
33-
continue
34-
globals()[name] = getattr(module, name)
35-
__all__.append(name)
36-
37-
38-
_layers = _import_optional_public_module('fla.layers')
39-
_models = _import_optional_public_module('fla.models')
40-
if _layers is not None and _models is not None:
41-
_export_public_api(_layers)
42-
_export_public_api(_models)
43-
44-
del _import_optional_public_module, _export_public_api, _layers, _models
15+
__all__ = ["modules", "ops"]

fla/modules/__init__.py

Lines changed: 3 additions & 44 deletions
Original file line numberDiff line numberDiff line change
@@ -5,48 +5,7 @@
55
# For a list of all contributors, visit:
66
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
77

8-
from fla.modules.convolution import ImplicitLongConvolution, LongConvolution, ShortConvolution
9-
from fla.modules.fused_bitlinear import BitLinear, FusedBitLinear
10-
from fla.modules.fused_cross_entropy import FusedCrossEntropyLoss
11-
from fla.modules.fused_kl_div import FusedKLDivLoss
12-
from fla.modules.fused_linear_cross_entropy import FusedLinearCrossEntropyLoss
13-
from fla.modules.fused_norm_gate import (
14-
FusedLayerNormGated,
15-
FusedLayerNormSwishGate,
16-
FusedLayerNormSwishGateLinear,
17-
FusedRMSNormGated,
18-
FusedRMSNormSwishGate,
19-
FusedRMSNormSwishGateLinear,
20-
)
21-
from fla.modules.l2norm import L2Norm
22-
from fla.modules.layernorm import GroupNorm, GroupNormLinear, LayerNorm, LayerNormLinear, RMSNorm, RMSNormLinear
23-
from fla.modules.mlp import GatedMLP
24-
from fla.modules.rotary import RotaryEmbedding
25-
from fla.modules.token_shift import TokenShift
8+
from fla.modules.conv.short_conv import ShortConvolution
9+
from fla.modules.fused_norm_gate import FusedRMSNormGated
2610

27-
__all__ = [
28-
'BitLinear',
29-
'FusedBitLinear',
30-
'FusedCrossEntropyLoss',
31-
'FusedKLDivLoss',
32-
'FusedLayerNormGated',
33-
'FusedLayerNormSwishGate',
34-
'FusedLayerNormSwishGateLinear',
35-
'FusedLinearCrossEntropyLoss',
36-
'FusedRMSNormGated',
37-
'FusedRMSNormSwishGate',
38-
'FusedRMSNormSwishGateLinear',
39-
'GatedMLP',
40-
'GroupNorm',
41-
'GroupNormLinear',
42-
'ImplicitLongConvolution',
43-
'L2Norm',
44-
'LayerNorm',
45-
'LayerNormLinear',
46-
'LongConvolution',
47-
'RMSNorm',
48-
'RMSNormLinear',
49-
'RotaryEmbedding',
50-
'ShortConvolution',
51-
'TokenShift',
52-
]
11+
__all__ = ["FusedRMSNormGated", "ShortConvolution"]

fla/modules/backends/__init__.py

Lines changed: 2 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -7,11 +7,8 @@
77

88
"""Module-level backends for FLA components such as rotary and cross-entropy."""
99

10-
from fla.modules.backends.triton_ascend import TritonAscendBackend
11-
from fla.ops.backends import BackendRegistry, dispatch
10+
from fla.ops.backends import dispatch
1211

13-
modules_registry = BackendRegistry("modules")
14-
15-
modules_registry.register(TritonAscendBackend())
12+
modules_registry = None
1613

1714
__all__ = ['dispatch', 'modules_registry']

fla/modules/conv/causal_conv1d.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,8 @@
99

1010
import torch
1111

12+
from fla.modules.conv.cp import causal_conv1d_cp
13+
from fla.modules.conv.triton import CausalConv1dFunction
1214
from fla.ops.cp import FLACPContext
1315
from fla.utils import input_guard
1416

@@ -65,11 +67,6 @@ def causal_conv1d(
6567
Tuple of (output, final_state).
6668
If `output_final_state` is `False`, the final state is `None`.
6769
"""
68-
# Import here to avoid circular dependencies
69-
from fla.modules.conv.cp import causal_conv1d_cp
70-
from fla.modules.conv.cuda import causal_conv1d_cuda, fast_causal_conv1d_fn
71-
from fla.modules.conv.triton import CausalConv1dFunction
72-
7370
if cp_context is not None:
7471
assert initial_state is None, "Initial state is not supported for CP"
7572
assert output_final_state is False, "Output final state is not supported for CP"
@@ -98,6 +95,8 @@ def causal_conv1d(
9895
)
9996
return y, final_state
10097
elif backend == 'mix':
98+
from fla.modules.conv.cuda import fast_causal_conv1d_fn
99+
101100
seq_idx = kwargs.get('seq_idx')
102101
return fast_causal_conv1d_fn(
103102
x,
@@ -113,6 +112,8 @@ def causal_conv1d(
113112
seq_idx=seq_idx,
114113
)
115114
elif backend == 'cuda':
115+
from fla.modules.conv.cuda import causal_conv1d_cuda
116+
116117
return causal_conv1d_cuda(
117118
x,
118119
weight,

fla/modules/conv/short_conv.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@
77

88
"""Short convolution implementation for efficient causal convolutions."""
99

10+
from __future__ import annotations
11+
1012
import warnings
1113

1214
import torch

fla/modules/conv/triton/ops.py

Lines changed: 16 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -61,7 +61,7 @@ def causal_conv1d_fwd(
6161
NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT)
6262
NB = triton.cdiv(B*T, 1024)
6363

64-
y = torch.empty_like(x, memory_format=torch.contiguous_format)
64+
y = torch.empty_like(x)
6565

6666
def grid(meta): return (triton.cdiv(D, meta['BD']), NT, B)
6767
causal_conv1d_fwd_kernel[grid](
@@ -378,6 +378,12 @@ def forward(
378378
chunk_indices: torch.LongTensor | None = None,
379379
chunk_size: int = 64,
380380
):
381+
ctx.has_bias = bias is not None
382+
ctx.has_residual = residual is not None
383+
ctx.has_initial_state = initial_state is not None
384+
ctx.has_cu_seqlens = cu_seqlens is not None
385+
ctx.has_cu_seqlens_cpu = cu_seqlens_cpu is not None
386+
ctx.has_chunk_indices = chunk_indices is not None
381387
BT = chunk_size
382388
if cu_seqlens is not None and chunk_indices is None:
383389
chunk_indices = prepare_chunk_indices(cu_seqlens, BT, cu_seqlens_cpu=cu_seqlens_cpu)
@@ -421,4 +427,12 @@ def backward(ctx, dy: torch.Tensor, dht: torch.Tensor | None = None):
421427
chunk_indices=ctx.chunk_indices,
422428
layout_fallback=ctx.layout_fallback,
423429
)
424-
return dx, dw, db, dr, dh0, None, None, None, None, None, None
430+
return (
431+
(dx, dw)
432+
+ ((db,) if ctx.has_bias else ())
433+
+ ((dr,) if ctx.has_residual else ())
434+
+ ((dh0,) if ctx.has_initial_state else ())
435+
+ ((None,) if ctx.has_cu_seqlens else ())
436+
+ ((None,) if ctx.has_cu_seqlens_cpu else ())
437+
+ ((None,) if ctx.has_chunk_indices else ())
438+
)

fla/modules/fused_norm_gate.py

Lines changed: 4 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -656,6 +656,7 @@ def forward(
656656
residual_in_fp32: bool = False,
657657
is_rms_norm: bool = False,
658658
):
659+
ctx.has_bias = bias is not None
659660
x_shape_og = x.shape
660661
g_shape_og = g.shape
661662
# reshape input data into 2D tensor
@@ -716,16 +717,9 @@ def backward(ctx, dy, *args):
716717
x_dtype=ctx.x_dtype,
717718
)
718719
return (
719-
dx.reshape(ctx.x_shape_og),
720-
dg.reshape(ctx.g_shape_og),
721-
dw,
722-
db,
723-
None,
724-
dres_in.reshape(ctx.x_shape_og) if ctx.has_residual else None,
725-
None,
726-
None,
727-
None,
728-
None,
720+
(dx.reshape(ctx.x_shape_og), dg.reshape(ctx.g_shape_og), dw)
721+
+ ((db,) if ctx.has_bias else ())
722+
+ ((dres_in.reshape(ctx.x_shape_og),) if ctx.has_residual else ())
729723
)
730724

731725

fla/ops/__init__.py

Lines changed: 2 additions & 83 deletions
Original file line numberDiff line numberDiff line change
@@ -5,87 +5,6 @@
55
# For a list of all contributors, visit:
66
# https://github.com/fla-org/flash-linear-attention/graphs/contributors
77

8-
from .abc import chunk_abc
9-
from .attn import parallel_attn
10-
from .attnres import fused_attnres
11-
from .based import fused_chunk_based, parallel_based
12-
from .comba import chunk_comba, fused_recurrent_comba
13-
from .delta_rule import chunk_delta_rule, fused_chunk_delta_rule, fused_recurrent_delta_rule
14-
from .forgetting_attn import parallel_forgetting_attn
15-
from .gated_delta_rule import chunk_gated_delta_rule, chunk_gdn, fused_recurrent_gated_delta_rule, fused_recurrent_gdn
16-
from .generalized_delta_rule import (
17-
chunk_dplr_delta_rule,
18-
chunk_iplr_delta_rule,
19-
fused_recurrent_dplr_delta_rule,
20-
fused_recurrent_iplr_delta_rule,
21-
)
22-
from .gla import chunk_gla, fused_chunk_gla, fused_recurrent_gla
23-
from .gsa import chunk_gsa, fused_recurrent_gsa
24-
from .hgrn import fused_recurrent_hgrn
25-
from .kda import chunk_kda, fused_recurrent_kda
26-
from .lightning_attn import chunk_lightning_attn, fused_recurrent_lightning_attn
27-
from .linear_attn import chunk_linear_attn, fused_chunk_linear_attn, fused_recurrent_linear_attn
28-
from .log_linear_attn import chunk_log_linear_attn
29-
from .mesa_net import chunk_mesa_net
30-
from .nsa import parallel_nsa
31-
from .parallax import parallel_parallax
32-
from .path_attn import parallel_path_attn
33-
from .retention import chunk_retention, fused_chunk_retention, fused_recurrent_retention, parallel_retention
34-
from .rwkv6 import chunk_rwkv6, fused_recurrent_rwkv6
35-
from .rwkv7 import chunk_rwkv7, fused_recurrent_rwkv7
36-
from .simple_gla import chunk_simple_gla, fused_chunk_simple_gla, fused_recurrent_simple_gla, parallel_simple_gla
37-
from .wall_attn import parallel_wall_attn, parallel_wall_attn_decode
8+
from fla.ops import cp, kda, utils
389

39-
__all__ = [
40-
'chunk_abc',
41-
'chunk_comba',
42-
'chunk_delta_rule',
43-
'chunk_dplr_delta_rule',
44-
'chunk_gated_delta_rule',
45-
'chunk_gdn',
46-
'chunk_gla',
47-
'chunk_gsa',
48-
'chunk_iplr_delta_rule',
49-
'chunk_kda',
50-
'chunk_lightning_attn',
51-
'chunk_linear_attn',
52-
'chunk_log_linear_attn',
53-
'chunk_mesa_net',
54-
'chunk_retention',
55-
'chunk_rwkv6',
56-
'chunk_rwkv7',
57-
'chunk_simple_gla',
58-
'fused_attnres',
59-
'fused_chunk_based',
60-
'fused_chunk_delta_rule',
61-
'fused_chunk_gla',
62-
'fused_chunk_linear_attn',
63-
'fused_chunk_retention',
64-
'fused_chunk_simple_gla',
65-
'fused_recurrent_comba',
66-
'fused_recurrent_delta_rule',
67-
'fused_recurrent_dplr_delta_rule',
68-
'fused_recurrent_gated_delta_rule',
69-
'fused_recurrent_gdn',
70-
'fused_recurrent_gla',
71-
'fused_recurrent_gsa',
72-
'fused_recurrent_hgrn',
73-
'fused_recurrent_iplr_delta_rule',
74-
'fused_recurrent_kda',
75-
'fused_recurrent_lightning_attn',
76-
'fused_recurrent_linear_attn',
77-
'fused_recurrent_retention',
78-
'fused_recurrent_rwkv6',
79-
'fused_recurrent_rwkv7',
80-
'fused_recurrent_simple_gla',
81-
'parallel_attn',
82-
'parallel_based',
83-
'parallel_forgetting_attn',
84-
'parallel_nsa',
85-
'parallel_parallax',
86-
'parallel_path_attn',
87-
'parallel_retention',
88-
'parallel_simple_gla',
89-
'parallel_wall_attn',
90-
'parallel_wall_attn_decode',
91-
]
10+
__all__ = ["cp", "kda", "utils"]

fla/ops/backends/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -164,6 +164,8 @@ def dispatch(operation: str):
164164
that passes the verifier for the given function call.
165165
"""
166166
def decorator(func: F) -> F:
167+
return func
168+
167169
if _DISPATCH_DISABLED:
168170
return func
169171
func_name = func.__name__

fla/ops/common/gate.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -57,7 +57,7 @@ def fused_beta_sigmoid_bwd_kernel(
5757
@dispatch('common')
5858
def fused_beta_sigmoid_fwd(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
5959
y = torch.empty_like(x, dtype=torch.float32)
60-
n_elements = x.numel()
60+
n_elements = x.shape.numel()
6161
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
6262
fused_beta_sigmoid_fwd_kernel[grid](
6363
x,
@@ -73,7 +73,7 @@ def fused_beta_sigmoid_fwd(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
7373
@dispatch('common')
7474
def fused_beta_sigmoid_bwd(x: torch.Tensor, dy: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
7575
dx = torch.empty_like(x)
76-
n_elements = x.numel()
76+
n_elements = x.shape.numel()
7777
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
7878
fused_beta_sigmoid_bwd_kernel[grid](
7979
x,
@@ -103,7 +103,7 @@ def forward(ctx, x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
103103
def backward(ctx, dy: torch.Tensor):
104104
(x,) = ctx.saved_tensors
105105
dx = fused_beta_sigmoid_bwd(x, dy, ctx.scale)
106-
return dx.type_as(x), None
106+
return (dx.type_as(x),)
107107

108108

109109
def fused_beta_sigmoid(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:

0 commit comments

Comments
 (0)