Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 0 additions & 31 deletions fla/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,40 +5,9 @@
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors

import importlib
from pkgutil import extend_path

__path__ = extend_path(__path__, __name__)
__version__ = "0.5.2"

__all__: list[str] = []


def _import_optional_public_module(module_name: str):
try:
return importlib.import_module(module_name)
except ModuleNotFoundError as exc:
missing = exc.name
# The extension package is optional. Treat its absence, or the absence
# of an external runtime dependency, as the extension being unavailable.
if missing == module_name or (missing is not None and missing.split('.', 1)[0] != 'fla'):
return None
raise


def _export_public_api(module) -> None:
globals()[module.__name__.rsplit('.', maxsplit=1)[-1]] = module
for name in module.__all__:
if name.endswith('Config'):
continue
globals()[name] = getattr(module, name)
__all__.append(name)


_layers = _import_optional_public_module('fla.layers')
_models = _import_optional_public_module('fla.models')
if _layers is not None and _models is not None:
_export_public_api(_layers)
_export_public_api(_models)

del _import_optional_public_module, _export_public_api, _layers, _models
47 changes: 3 additions & 44 deletions fla/modules/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,48 +5,7 @@
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors

from fla.modules.convolution import ImplicitLongConvolution, LongConvolution, ShortConvolution
from fla.modules.fused_bitlinear import BitLinear, FusedBitLinear
from fla.modules.fused_cross_entropy import FusedCrossEntropyLoss
from fla.modules.fused_kl_div import FusedKLDivLoss
from fla.modules.fused_linear_cross_entropy import FusedLinearCrossEntropyLoss
from fla.modules.fused_norm_gate import (
FusedLayerNormGated,
FusedLayerNormSwishGate,
FusedLayerNormSwishGateLinear,
FusedRMSNormGated,
FusedRMSNormSwishGate,
FusedRMSNormSwishGateLinear,
)
from fla.modules.l2norm import L2Norm
from fla.modules.layernorm import GroupNorm, GroupNormLinear, LayerNorm, LayerNormLinear, RMSNorm, RMSNormLinear
from fla.modules.mlp import GatedMLP
from fla.modules.rotary import RotaryEmbedding
from fla.modules.token_shift import TokenShift
from fla.modules.conv.short_conv import ShortConvolution
from fla.modules.fused_norm_gate import FusedRMSNormGated

__all__ = [
'BitLinear',
'FusedBitLinear',
'FusedCrossEntropyLoss',
'FusedKLDivLoss',
'FusedLayerNormGated',
'FusedLayerNormSwishGate',
'FusedLayerNormSwishGateLinear',
'FusedLinearCrossEntropyLoss',
'FusedRMSNormGated',
'FusedRMSNormSwishGate',
'FusedRMSNormSwishGateLinear',
'GatedMLP',
'GroupNorm',
'GroupNormLinear',
'ImplicitLongConvolution',
'L2Norm',
'LayerNorm',
'LayerNormLinear',
'LongConvolution',
'RMSNorm',
'RMSNormLinear',
'RotaryEmbedding',
'ShortConvolution',
'TokenShift',
]
__all__ = ["FusedRMSNormGated", "ShortConvolution"]
7 changes: 2 additions & 5 deletions fla/modules/backends/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,11 +7,8 @@

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

from fla.modules.backends.triton_ascend import TritonAscendBackend
from fla.ops.backends import BackendRegistry, dispatch
from fla.ops.backends import dispatch

modules_registry = BackendRegistry("modules")

modules_registry.register(TritonAscendBackend())
modules_registry = None

__all__ = ['dispatch', 'modules_registry']
13 changes: 8 additions & 5 deletions fla/modules/conv/causal_conv1d.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,12 +65,9 @@ def causal_conv1d(
Tuple of (output, final_state).
If `output_final_state` is `False`, the final state is `None`.
"""
# Import here to avoid circular dependencies
from fla.modules.conv.cp import causal_conv1d_cp
from fla.modules.conv.cuda import causal_conv1d_cuda, fast_causal_conv1d_fn
from fla.modules.conv.triton import CausalConv1dFunction

if cp_context is not None:
from fla.modules.conv.cp import causal_conv1d_cp

assert initial_state is None, "Initial state is not supported for CP"
assert output_final_state is False, "Output final state is not supported for CP"
output = causal_conv1d_cp(
Expand All @@ -84,6 +81,8 @@ def causal_conv1d(
return output, None

if backend == 'triton':
from fla.modules.conv.triton import CausalConv1dFunction

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

这里的 lazy import 要注意下


y, final_state = CausalConv1dFunction.apply(
x,
weight,
Expand All @@ -98,6 +97,8 @@ def causal_conv1d(
)
return y, final_state
elif backend == 'mix':
from fla.modules.conv.cuda import fast_causal_conv1d_fn

seq_idx = kwargs.get('seq_idx')
return fast_causal_conv1d_fn(
x,
Expand All @@ -113,6 +114,8 @@ def causal_conv1d(
seq_idx=seq_idx,
)
elif backend == 'cuda':
from fla.modules.conv.cuda import causal_conv1d_cuda

return causal_conv1d_cuda(
x,
weight,
Expand Down
2 changes: 2 additions & 0 deletions fla/modules/conv/short_conv.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@

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

from __future__ import annotations

import warnings

import torch
Expand Down
18 changes: 16 additions & 2 deletions fla/modules/conv/triton/ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,7 +61,7 @@ def causal_conv1d_fwd(
NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT)
NB = triton.cdiv(B*T, 1024)

y = torch.empty_like(x, memory_format=torch.contiguous_format)
y = torch.empty_like(x)

def grid(meta): return (triton.cdiv(D, meta['BD']), NT, B)
causal_conv1d_fwd_kernel[grid](
Expand Down Expand Up @@ -378,6 +378,12 @@ def forward(
chunk_indices: torch.LongTensor | None = None,
chunk_size: int = 64,
):
ctx.has_bias = bias is not None
ctx.has_residual = residual is not None
ctx.has_initial_state = initial_state is not None
ctx.has_cu_seqlens = cu_seqlens is not None
ctx.has_cu_seqlens_cpu = cu_seqlens_cpu is not None
ctx.has_chunk_indices = chunk_indices is not None
BT = chunk_size
if cu_seqlens is not None and chunk_indices is None:
chunk_indices = prepare_chunk_indices(cu_seqlens, BT, cu_seqlens_cpu=cu_seqlens_cpu)
Expand Down Expand Up @@ -421,4 +427,12 @@ def backward(ctx, dy: torch.Tensor, dht: torch.Tensor | None = None):
chunk_indices=ctx.chunk_indices,
layout_fallback=ctx.layout_fallback,
)
return dx, dw, db, dr, dh0, None, None, None, None, None, None
return (
(dx, dw)
+ ((db,) if ctx.has_bias else ())
+ ((dr,) if ctx.has_residual else ())
+ ((dh0,) if ctx.has_initial_state else ())
+ ((None,) if ctx.has_cu_seqlens else ())
+ ((None,) if ctx.has_cu_seqlens_cpu else ())
+ ((None,) if ctx.has_chunk_indices else ())
)
14 changes: 4 additions & 10 deletions fla/modules/fused_norm_gate.py
Original file line number Diff line number Diff line change
Expand Up @@ -656,6 +656,7 @@ def forward(
residual_in_fp32: bool = False,
is_rms_norm: bool = False,
):
ctx.has_bias = bias is not None
x_shape_og = x.shape
g_shape_og = g.shape
# reshape input data into 2D tensor
Expand Down Expand Up @@ -716,16 +717,9 @@ def backward(ctx, dy, *args):
x_dtype=ctx.x_dtype,
)
return (
dx.reshape(ctx.x_shape_og),
dg.reshape(ctx.g_shape_og),
dw,
db,
None,
dres_in.reshape(ctx.x_shape_og) if ctx.has_residual else None,
None,
None,
None,
None,
(dx.reshape(ctx.x_shape_og), dg.reshape(ctx.g_shape_og), dw)
+ ((db,) if ctx.has_bias else ())
+ ((dres_in.reshape(ctx.x_shape_og),) if ctx.has_residual else ())
)


Expand Down
85 changes: 1 addition & 84 deletions fla/ops/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,87 +5,4 @@
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors

from .abc import chunk_abc
from .attn import parallel_attn
from .attnres import fused_attnres
from .based import fused_chunk_based, parallel_based
from .comba import chunk_comba, fused_recurrent_comba
from .delta_rule import chunk_delta_rule, fused_chunk_delta_rule, fused_recurrent_delta_rule
from .forgetting_attn import parallel_forgetting_attn
from .gated_delta_rule import chunk_gated_delta_rule, chunk_gdn, fused_recurrent_gated_delta_rule, fused_recurrent_gdn
from .generalized_delta_rule import (
chunk_dplr_delta_rule,
chunk_iplr_delta_rule,
fused_recurrent_dplr_delta_rule,
fused_recurrent_iplr_delta_rule,
)
from .gla import chunk_gla, fused_chunk_gla, fused_recurrent_gla
from .gsa import chunk_gsa, fused_recurrent_gsa
from .hgrn import fused_recurrent_hgrn
from .kda import chunk_kda, fused_recurrent_kda
from .lightning_attn import chunk_lightning_attn, fused_recurrent_lightning_attn
from .linear_attn import chunk_linear_attn, fused_chunk_linear_attn, fused_recurrent_linear_attn
from .log_linear_attn import chunk_log_linear_attn
from .mesa_net import chunk_mesa_net
from .nsa import parallel_nsa
from .parallax import parallel_parallax
from .path_attn import parallel_path_attn
from .retention import chunk_retention, fused_chunk_retention, fused_recurrent_retention, parallel_retention
from .rwkv6 import chunk_rwkv6, fused_recurrent_rwkv6
from .rwkv7 import chunk_rwkv7, fused_recurrent_rwkv7
from .simple_gla import chunk_simple_gla, fused_chunk_simple_gla, fused_recurrent_simple_gla, parallel_simple_gla
from .wall_attn import parallel_wall_attn, parallel_wall_attn_decode

__all__ = [
'chunk_abc',
'chunk_comba',
'chunk_delta_rule',
'chunk_dplr_delta_rule',
'chunk_gated_delta_rule',
'chunk_gdn',
'chunk_gla',
'chunk_gsa',
'chunk_iplr_delta_rule',
'chunk_kda',
'chunk_lightning_attn',
'chunk_linear_attn',
'chunk_log_linear_attn',
'chunk_mesa_net',
'chunk_retention',
'chunk_rwkv6',
'chunk_rwkv7',
'chunk_simple_gla',
'fused_attnres',
'fused_chunk_based',
'fused_chunk_delta_rule',
'fused_chunk_gla',
'fused_chunk_linear_attn',
'fused_chunk_retention',
'fused_chunk_simple_gla',
'fused_recurrent_comba',
'fused_recurrent_delta_rule',
'fused_recurrent_dplr_delta_rule',
'fused_recurrent_gated_delta_rule',
'fused_recurrent_gdn',
'fused_recurrent_gla',
'fused_recurrent_gsa',
'fused_recurrent_hgrn',
'fused_recurrent_iplr_delta_rule',
'fused_recurrent_kda',
'fused_recurrent_lightning_attn',
'fused_recurrent_linear_attn',
'fused_recurrent_retention',
'fused_recurrent_rwkv6',
'fused_recurrent_rwkv7',
'fused_recurrent_simple_gla',
'parallel_attn',
'parallel_based',
'parallel_forgetting_attn',
'parallel_nsa',
'parallel_parallax',
'parallel_path_attn',
'parallel_retention',
'parallel_simple_gla',
'parallel_wall_attn',
'parallel_wall_attn_decode',
]
__all__: list[str] = []
Loading