Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
33 changes: 2 additions & 31 deletions fla/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,40 +5,11 @@
# 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] = []
from fla import modules, ops # noqa: E402


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
__all__ = ["modules", "ops"]
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']
11 changes: 6 additions & 5 deletions fla/modules/conv/causal_conv1d.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@

import torch

from fla.modules.conv.cp import causal_conv1d_cp
from fla.modules.conv.triton import CausalConv1dFunction
from fla.ops.cp import FLACPContext
from fla.utils import input_guard

Expand Down Expand Up @@ -65,11 +67,6 @@ 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:
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"
Expand Down Expand Up @@ -98,6 +95,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 +112,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: 2 additions & 83 deletions fla/ops/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,87 +5,6 @@
# 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
from fla.ops import cp, kda, utils

__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__ = ["cp", "kda", "utils"]
2 changes: 2 additions & 0 deletions fla/ops/backends/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -164,6 +164,8 @@ def dispatch(operation: str):
that passes the verifier for the given function call.
"""
def decorator(func: F) -> F:
return func

if _DISPATCH_DISABLED:
return func
func_name = func.__name__
Expand Down
6 changes: 3 additions & 3 deletions fla/ops/common/gate.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ def fused_beta_sigmoid_bwd_kernel(
@dispatch('common')
def fused_beta_sigmoid_fwd(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
y = torch.empty_like(x, dtype=torch.float32)
n_elements = x.numel()
n_elements = x.shape.numel()
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
fused_beta_sigmoid_fwd_kernel[grid](
x,
Expand All @@ -73,7 +73,7 @@ def fused_beta_sigmoid_fwd(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
@dispatch('common')
def fused_beta_sigmoid_bwd(x: torch.Tensor, dy: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
dx = torch.empty_like(x)
n_elements = x.numel()
n_elements = x.shape.numel()
grid = (triton.cdiv(n_elements, _BETA_SIGMOID_BLOCK_SIZE),)
fused_beta_sigmoid_bwd_kernel[grid](
x,
Expand Down Expand Up @@ -103,7 +103,7 @@ def forward(ctx, x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
def backward(ctx, dy: torch.Tensor):
(x,) = ctx.saved_tensors
dx = fused_beta_sigmoid_bwd(x, dy, ctx.scale)
return dx.type_as(x), None
return (dx.type_as(x),)


def fused_beta_sigmoid(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor:
Expand Down
15 changes: 10 additions & 5 deletions fla/ops/cp/comm.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,11 +34,16 @@ def all_gather_into_tensor(
Returns:
Tuple of (output tensor, handle if async_op else None)
"""
world_size = dist.get_world_size(group=group)
if async_op:
raise NotImplementedError("KDA context parallel currently supports synchronous all-gather only")
gathered = []
dist.all_gather(gathered, inp, group=group, sync_op=True)
gathered_tensor = torch.stack(gathered, dim=0)
if out is None:
out = torch.empty(world_size, *inp.shape, device=inp.device, dtype=inp.dtype)
handle = dist.all_gather_into_tensor(out, inp, group=group, async_op=async_op)
return out, handle
out = gathered_tensor
else:
out.copy_(gathered_tensor)
return out, None


def all_reduce_sum(
Expand All @@ -57,7 +62,7 @@ def all_reduce_sum(
Returns:
Tuple of (reduced tensor, handle if async_op else None)
"""
handle = dist.all_reduce(inp, op=dist.ReduceOp.SUM, group=group, async_op=async_op)
handle = dist.all_reduce(inp, op=dist.ReduceOp.SUM, group=group, sync_op=not async_op)
return inp, handle


Expand Down
Loading
Loading