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
11 changes: 10 additions & 1 deletion fla/layers/gated_deltanet.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,10 @@ class GatedDeltaNet(nn.Module):
The index of the layer. Default: None.
norm_eps (float, Optional):
The epsilon value for the normalization layer. Default: 1e-5.
output_gate_activation (str, Optional):
Activation applied by the output gate when `use_gate=True`.
Supported values are `swish`, `silu`, and `sigmoid`.
`swish` and `silu` are equivalent aliases. Default: `swish`.
"""

def __init__(
Expand All @@ -100,6 +104,7 @@ def __init__(
conv_bias: bool = False,
layer_idx: int = None,
norm_eps: float = 1e-5,
output_gate_activation: str = 'swish',
**kwargs,
) -> GatedDeltaNet:
super().__init__()
Expand Down Expand Up @@ -141,6 +146,10 @@ def __init__(
f"Resulting head_v_dim would be {head_dim * expand_v}, which is invalid for FusedRMSNormGated.",
)
assert mode in ['chunk', 'fused_recurrent'], f"Not supported mode `{mode}`."
if output_gate_activation not in ('swish', 'silu', 'sigmoid'):
raise ValueError(
f"output_gate_activation must be one of 'swish', 'silu', 'sigmoid', got {output_gate_activation!r}")
self.output_gate_activation = output_gate_activation

self.q_proj = nn.Linear(hidden_size, self.key_dim, bias=False)
self.k_proj = nn.Linear(hidden_size, self.key_dim, bias=False)
Expand Down Expand Up @@ -194,7 +203,7 @@ def __init__(
)
if use_gate:
self.g_proj = nn.Linear(hidden_size, self.value_dim, bias=False)
self.o_norm = FusedRMSNormGated(self.head_v_dim, eps=norm_eps)
self.o_norm = FusedRMSNormGated(self.head_v_dim, eps=norm_eps, activation=self.output_gate_activation)
else:
self.o_norm = RMSNorm(self.head_v_dim, eps=norm_eps, dtype=torch.float32)
self.o_proj = nn.Linear(self.value_dim, hidden_size, bias=False)
Expand Down
5 changes: 5 additions & 0 deletions fla/models/gated_deltanet/configuration_gated_deltanet.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ def __init__(
use_l2warp: bool = False,
vocab_size: int = 32000,
attnres_block_size: int | None = None,
output_gate_activation: str = "swish",
**kwargs,
):
self.attn_mode = attn_mode
Expand Down Expand Up @@ -78,6 +79,10 @@ def __init__(
self.vocab_size = vocab_size
self.allow_neg_eigval = allow_neg_eigval
self.attnres_block_size = attnres_block_size
self.output_gate_activation = output_gate_activation
if self.output_gate_activation not in ("swish", "silu", "sigmoid"):
raise ValueError(
f"output_gate_activation must be one of 'swish', 'silu', 'sigmoid', got {self.output_gate_activation!r}")

if fuse_cross_entropy and fuse_linear_cross_entropy:
raise ValueError(
Expand Down
3 changes: 2 additions & 1 deletion fla/models/gated_deltanet/modeling_gated_deltanet.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,8 +73,9 @@ def __init__(self, config: GatedDeltaNetConfig, layer_idx: int):
use_short_conv=config.use_short_conv,
allow_neg_eigval=config.allow_neg_eigval,
conv_size=config.conv_size,
norm_eps=config.norm_eps,
layer_idx=layer_idx,
norm_eps=config.norm_eps,
output_gate_activation=config.output_gate_activation,
)
self.mlp_norm = (RMSNorm if config.fuse_norm else nn.RMSNorm)(config.hidden_size, eps=config.norm_eps)
self.mlp = GatedDeltaNetMLP(
Expand Down
167 changes: 167 additions & 0 deletions tests/layers/test_gated_deltanet_output_gate.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,167 @@
# 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 pytest
import torch

from fla.layers.gated_deltanet import GatedDeltaNet
from fla.models.gated_deltanet.configuration_gated_deltanet import GatedDeltaNetConfig
from fla.models.gated_deltanet.modeling_gated_deltanet import GatedDeltaNetBlock
from fla.utils import device


def _tiny_config(**overrides):
base = dict(
hidden_size=128,
head_dim=32,
num_heads=2,
num_v_heads=2,
expand_v=1,
intermediate_size=256,
hidden_ratio=2,
max_position_embeddings=512,
num_hidden_layers=2,
vocab_size=512,
)
base.update(overrides)
return GatedDeltaNetConfig(**base)


def test_gated_deltanet_default_output_gate_activation_is_swish():
layer = GatedDeltaNet(hidden_size=128, head_dim=32, num_heads=2, expand_v=1)
assert layer.output_gate_activation == "swish"
assert layer.o_norm.activation == "swish"
cfg = _tiny_config()
assert cfg.output_gate_activation == "swish"
block = GatedDeltaNetBlock(cfg, layer_idx=0)
assert block.attn.output_gate_activation == "swish"
assert block.attn.o_norm.activation == "swish"


@pytest.mark.parametrize("gate", ["sigmoid", "silu"])
def test_gated_deltanet_output_gate_activation_wiring(gate):
layer = GatedDeltaNet(hidden_size=128, head_dim=32, num_heads=2, expand_v=1, output_gate_activation=gate)
assert layer.output_gate_activation == gate
assert layer.o_norm.activation == gate
cfg = _tiny_config(output_gate_activation=gate)
block = GatedDeltaNetBlock(cfg, layer_idx=0)
assert block.attn.output_gate_activation == gate
assert block.attn.o_norm.activation == gate
assert cfg.to_dict()["output_gate_activation"] == gate


def test_gated_deltanet_config_serialization():
cfg = _tiny_config(output_gate_activation="sigmoid")
d = cfg.to_dict()
assert d["output_gate_activation"] == "sigmoid"
cfg2 = GatedDeltaNetConfig.from_dict(d)
assert cfg2.output_gate_activation == "sigmoid"
block = GatedDeltaNetBlock(cfg2, layer_idx=0)
assert block.attn.o_norm.activation == "sigmoid"
# old checkpoint without field defaults to swish
cfg3 = _tiny_config(output_gate_activation="sigmoid")
d3 = cfg3.to_dict()
d3.pop("output_gate_activation")
restored = GatedDeltaNetConfig.from_dict(d3)
assert restored.output_gate_activation == "swish"
block = GatedDeltaNetBlock(restored, layer_idx=0)
assert block.attn.output_gate_activation == "swish"
assert block.attn.o_norm.activation == "swish"


def test_gated_deltanet_state_dict_compatibility():
tiny_kwargs = dict(hidden_size=128, head_dim=32, num_heads=2, expand_v=1)
base = GatedDeltaNet(**tiny_kwargs)
swish = GatedDeltaNet(**tiny_kwargs, output_gate_activation="swish")
sigmoid = GatedDeltaNet(**tiny_kwargs, output_gate_activation="sigmoid")
# activation choice adds no parameter/state-dict key
assert "output_gate_activation" not in base.state_dict()
assert set(base.state_dict().keys()) == set(swish.state_dict().keys()) == set(sigmoid.state_dict().keys())
assert sum(p.numel() for p in base.parameters()) == sum(
p.numel() for p in swish.parameters()
) == sum(p.numel() for p in sigmoid.parameters())
# strict load between variants
sd = base.state_dict()
swish.load_state_dict(sd, strict=True)
sigmoid.load_state_dict(sd, strict=True)


def test_gated_deltanet_sigmoid_forward_backward():
torch.manual_seed(42)

layer = GatedDeltaNet(
hidden_size=128,
head_dim=32,
num_heads=2,
expand_v=1,
output_gate_activation="sigmoid",
).to(device).train()

x = torch.randn(2, 16, 128, device=device, requires_grad=True)
y, _, _ = layer(x)

assert torch.isfinite(y).all()

y.sum().backward()

assert x.grad is not None
assert torch.isfinite(x.grad).all()


def test_gated_deltanet_validation_semantics():
with pytest.raises(ValueError):
GatedDeltaNetConfig(output_gate_activation="relu")
with pytest.raises(ValueError):
GatedDeltaNet(hidden_size=128, head_dim=32, num_heads=2, expand_v=1, output_gate_activation="relu")


def test_gated_deltanet_config_positional_compatibility():
attn_spec = {"layers": [0], "num_heads": 2}
cfg_pos = GatedDeltaNetConfig(
"chunk",
128,
1.0,
True,
True,
False,
4,
32,
2,
2,
512,
2,
256,
"swish",
2,
1e-5,
attn_spec,
False,
)
cfg_kw = GatedDeltaNetConfig(
attn_mode="chunk",
hidden_size=128,
expand_v=1.0,
use_gate=True,
use_short_conv=True,
allow_neg_eigval=False,
conv_size=4,
head_dim=32,
num_heads=2,
num_v_heads=2,
max_position_embeddings=512,
hidden_ratio=2,
intermediate_size=256,
hidden_act="swish",
num_hidden_layers=2,
norm_eps=1e-5,
attn=attn_spec,
use_cache=False,
)
assert cfg_pos.norm_eps == cfg_kw.norm_eps == 1e-5
assert cfg_pos.attn == cfg_kw.attn
assert cfg_pos.use_cache == cfg_kw.use_cache is False
assert cfg_pos.output_gate_activation == cfg_kw.output_gate_activation == "swish"
Loading