from __future__ import annotations import math from dataclasses import dataclass from typing import Any import torch import torch.nn as nn @dataclass(frozen=True, slots=True) class SequenceAdapterConfig: input_features: int = 4096 output_features: int = 5120 width: int = 512 num_heads: int = 6 head_dim: int = 64 expand_v: float = 2.0 mode: str = "chunk" use_short_conv: bool = True max_log_gain: float = 2.0 num_lanes: int = 4 output_head_reference_width: int = 512 def _gated_delta_net(config: SequenceAdapterConfig) -> nn.Module: try: from fla.layers import GatedDeltaNet except ImportError as error: raise RuntimeError( "This model requires fla-core and flash-linear-attention 0.5.1. " "Install the sequence dependencies documented in README.md." ) from error return GatedDeltaNet( hidden_size=config.width, expand_v=config.expand_v, head_dim=config.head_dim, num_heads=config.num_heads, mode=config.mode, use_gate=True, use_short_conv=config.use_short_conv, ) class _SequenceLane(nn.Module): def __init__(self, config: SequenceAdapterConfig) -> None: super().__init__() self.down = nn.Linear(config.input_features, config.width) self.norm = nn.RMSNorm(config.width) self.mixer = _gated_delta_net(config) def forward( self, source: torch.Tensor, *, attention_mask: torch.Tensor | None, cu_seqlens: torch.Tensor | None, ) -> torch.Tensor: mixed = self.mixer( self.norm(self.down(source)), attention_mask=attention_mask, cu_seqlens=cu_seqlens, ) if isinstance(mixed, tuple): mixed = mixed[0] if not isinstance(mixed, torch.Tensor): raise TypeError("Gated DeltaNet did not return hidden states") return mixed class GatedDeltaResidualAdapter(nn.Module): """Frozen affine map plus four causal Gated DeltaNet residual lanes.""" def __init__(self, config: SequenceAdapterConfig) -> None: super().__init__() self.config = config self.register_buffer( "affine_weight", torch.empty(config.input_features, config.output_features), ) self.register_buffer("affine_bias", torch.empty(config.output_features)) self.lanes = nn.ModuleList([_SequenceLane(config) for _ in range(config.num_lanes)]) feature_width = config.width * config.num_lanes self.gain = nn.Linear(feature_width, 1) self.up = nn.Linear(feature_width, config.output_features) def output_head_features(self, features: torch.Tensor) -> torch.Tensor: scale = math.sqrt( self.config.output_head_reference_width / features.shape[-1] ) return features * scale def forward( self, source: torch.Tensor, *, attention_mask: torch.Tensor | None = None, cu_seqlens: torch.Tensor | None = None, **_: Any, ) -> torch.Tensor: features = torch.cat( [ lane( source, attention_mask=attention_mask, cu_seqlens=cu_seqlens, ) for lane in self.lanes ], dim=-1, ) head_features = self.output_head_features(features) affine = source @ self.affine_weight + self.affine_bias log_gain = self.gain(head_features).float().clamp( min=-self.config.max_log_gain, max=self.config.max_log_gain, ) return affine * log_gain.exp().to(affine.dtype) + self.up(head_features) __all__ = ["GatedDeltaResidualAdapter", "SequenceAdapterConfig"]