Skip to content
Open
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
8 changes: 8 additions & 0 deletions fla/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,8 @@
RodimusAttention,
RWKV6Attention,
RWKV7Attention,
SSEGLA,
SSEGDN,
)
from fla.models import (
ABCForCausalLM,
Expand Down Expand Up @@ -86,6 +88,8 @@
RWKV6Model,
RWKV7ForCausalLM,
RWKV7Model,
SSEForCausalLM,
SSEModel,
TransformerForCausalLM,
TransformerModel,
)
Expand Down Expand Up @@ -169,6 +173,10 @@
"RodimusAttention",
"RodimusForCausalLM",
"RodimusModel",
"SSEGLA",
"SSEGDN",
"SSEForCausalLM",
"SSEModel",
"TransformerForCausalLM",
"TransformerModel",
]
Expand Down
3 changes: 3 additions & 0 deletions fla/layers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@
from .rodimus import RodimusAttention, SlidingWindowSharedKeyAttention
from .rwkv6 import RWKV6Attention
from .rwkv7 import RWKV7Attention
from .sse import SSEGLA, SSEGDN

__all__ = [
'ABCAttention',
Expand Down Expand Up @@ -72,4 +73,6 @@
'ReBasedLinearAttention',
'RodimusAttention',
'SlidingWindowSharedKeyAttention',
'SSEGLA',
'SSEGDN',
]
983 changes: 983 additions & 0 deletions fla/layers/sse.py

Large diffs are not rendered by default.

4 changes: 4 additions & 0 deletions fla/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
from fla.models.rwkv6 import RWKV6Config, RWKV6ForCausalLM, RWKV6Model
from fla.models.rwkv7 import RWKV7Config, RWKV7ForCausalLM, RWKV7Model
from fla.models.samba import SambaConfig, SambaForCausalLM, SambaModel
from fla.models.sse import SSEConfig, SSEForCausalLM, SSEModel
from fla.models.transformer import TransformerConfig, TransformerForCausalLM, TransformerModel

__all__ = [
Expand Down Expand Up @@ -132,6 +133,9 @@
'SambaConfig',
'SambaForCausalLM',
'SambaModel',
'SSEConfig',
'SSEForCausalLM',
'SSEModel',
'TransformerConfig',
'TransformerForCausalLM',
'TransformerModel',
Expand Down
17 changes: 17 additions & 0 deletions fla/models/sse/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors

from transformers import AutoConfig, AutoModel, AutoModelForCausalLM

from fla.models.sse.configuration_sse import SSEConfig
from fla.models.sse.modeling_sse import SSEForCausalLM, SSEModel

AutoConfig.register(SSEConfig.model_type, SSEConfig, exist_ok=True)
AutoModel.register(SSEConfig, SSEModel, exist_ok=True)
AutoModelForCausalLM.register(SSEConfig, SSEForCausalLM, exist_ok=True)

__all__ = ['SSEConfig', 'SSEForCausalLM', 'SSEModel']
119 changes: 119 additions & 0 deletions fla/models/sse/configuration_sse.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,119 @@
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
# https://github.com/fla-org/flash-linear-attention/graphs/contributors

import warnings

from transformers.configuration_utils import PretrainedConfig


class SSEConfig(PretrainedConfig):
model_type = 'sse'
keys_to_ignore_at_inference = ['past_key_values']

def __init__(
self,
attn_mode: str = "chunk",
hidden_size: int = 2048,
expand_v: float = 1.0,
use_output_gate: bool = True,
use_short_conv: bool = False,
allow_neg_eigval: bool = False,
conv_size: int = 4,
head_dim: int = 256,
num_heads: int = 6,
num_v_heads: int | None = None,
num_sparse_partition: int = 4,
num_writer: int = 2,
num_reader: int = 2,
linear_attn_type: str = "gla",
sse_implementation: str = "varlen",
aux_loss_coef: float = 0.01,
max_position_embeddings: int = 2048,
hidden_ratio: int | None = 4,
intermediate_size: int | None = None,
hidden_act: str = "swish",
num_hidden_layers: int = 24,
norm_eps: float = 1e-6,
attn: dict | None = None,
use_cache: bool = True,
pad_token_id: int | None = None,
bos_token_id: int = 1,
eos_token_id: int = 2,
tie_word_embeddings: bool = False,
initializer_range: float = 0.02,
fuse_norm: bool = True,
fuse_swiglu: bool = True,
fuse_cross_entropy: bool = True,
fuse_linear_cross_entropy: bool = False,
use_l2warp: bool = False,
vocab_size: int = 32000,
**kwargs,
):
self.attn_mode = attn_mode
self.hidden_size = hidden_size
self.expand_v = expand_v
self.use_output_gate = use_output_gate
self.use_short_conv = use_short_conv
self.conv_size = conv_size
self.head_dim = head_dim
self.num_heads = num_heads
self.num_v_heads = num_v_heads
self.num_sparse_partition = num_sparse_partition
self.num_writer = num_writer
self.num_reader = num_reader
self.linear_attn_type = linear_attn_type
self.sse_implementation = sse_implementation
self.aux_loss_coef = aux_loss_coef
self.max_position_embeddings = max_position_embeddings

self.hidden_ratio = hidden_ratio
self.intermediate_size = intermediate_size
self.hidden_act = hidden_act
self.num_hidden_layers = num_hidden_layers
self.norm_eps = norm_eps
self.attn = attn
self.use_cache = use_cache
self.initializer_range = initializer_range

self.fuse_norm = fuse_norm
self.fuse_swiglu = fuse_swiglu
self.fuse_cross_entropy = fuse_cross_entropy
self.fuse_linear_cross_entropy = fuse_linear_cross_entropy
self.use_l2warp = use_l2warp
self.vocab_size = vocab_size
self.allow_neg_eigval = allow_neg_eigval

if fuse_cross_entropy and fuse_linear_cross_entropy:
raise ValueError(
"`fuse_cross_entropy` and `fuse_linear_cross_entropy` cannot be True at the same time.",
)
if fuse_linear_cross_entropy:
warnings.warn(
"`fuse_linear_cross_entropy` is enabled, which can improves memory efficiency "
"at the potential cost of reduced precision. "
"If you observe issues like loss divergence, consider disabling this setting.",
)

if attn is not None:
if not isinstance(attn, dict):
raise ValueError("attn must be a dictionary")
if 'layers' not in attn:
raise ValueError("Layer indices must be provided to initialize hybrid attention layers")
if 'num_heads' not in attn:
raise ValueError("Number of heads must be provided to initialize hybrid attention layers")
attn['num_kv_heads'] = attn.get('num_kv_heads', attn['num_heads'])
attn['qkv_bias'] = attn.get('qkv_bias', False)
attn['window_size'] = attn.get('window_size', None)
attn['rope_theta'] = attn.get('rope_theta', 10000.)

super().__init__(
pad_token_id=pad_token_id,
bos_token_id=bos_token_id,
eos_token_id=eos_token_id,
tie_word_embeddings=tie_word_embeddings,
**kwargs,
)
Loading
Loading