"""Configuration for the owned dual-attention decoder LM.""" from dataclasses import dataclass from .attention_config import SUPPORTED_SEQUENCE_BOUNDARY_POLICIES SUPPORTED_PE_TYPES = {"sinusoidal", "learned", "relative", "rope", "none"} SUPPORTED_SYMBOL_RETRIEVAL = {"symbolic", "positional", "relative", "relsymbolic"} SUPPORTED_RA_TYPES = {"ra", "rca", "disrca"} SUPPORTED_RA_ACTIVATIONS = {"softmax", "identity", "relu", "tanh", "sigmoid", "gelu"} SUPPORTED_FFN_ACTIVATIONS = {"gelu", "relu", "swiglu", "identity"} SUPPORTED_FFN_HIDDEN_DIM_MODES = {"dff_factor", "swiglu_parameter_matched"} SUPPORTED_INIT_SCHEMES = {"xavier_uniform", "normal_0_02_scaled_projection"} @dataclass(frozen=True) class DatLMConfig: vocab_size: int max_seq_len: int pe_type: str = "rope" hidden_dim: int = 256 n_heads_sa: int = 2 n_heads_ra: int = 2 n_layers: int = 4 dropout: float = 0.0 dff_factor: int = 4 ffn_hidden_dim_mode: str = "dff_factor" ffn_activation: str = "gelu" rope_theta: float = 10000.0 max_rel_pos: int | None = None sequence_boundary_policy: str = "eos_document" segment_boundary_token_id: int | None = None init_range: float = 0.15 init_scheme: str = "xavier_uniform" norm_type: str = "rmsnorm" norm_first: bool = True use_bias_qkv: bool = False use_bias_out: bool = True use_bias_ffn: bool = True tie_lm_head: bool = True symbol_dim: int | None = None n_symbols: int | None = None symbolic_attn_n_heads: int | None = None symbol_retrieval: str = "symbolic" symbolic_use_bias: bool = False shared_symbol_retriever: bool = True share_attn_params: bool = False positional_symbols_sinusoidal: bool = False relative_symbols_rope: bool = False relsymbolic_rel_n_heads: int = 4 relsymbolic_symbolic_attn_n_heads: int = 4 relsymbolic_neighborhood_size: int = 2 relsymbolic_include_self: bool = False relsymbolic_normalize_rels: bool = True relsymbolic_trainable_symbols: bool = True relsymbolic_dropout: float = 0.0 relsymbolic_rel_scale: float | None = None relsymbolic_symbolic_attn_scale: float | None = None relsymbolic_use_bias: bool = False ra_type: str = "ra" ra_n_relations: int | None = None ra_rel_activation: str = "identity" ra_symmetric_rels: bool = False pad_token_id: int = 0 bos_token_id: int = 1 eos_token_id: int = 2 mlm_head_enabled: bool = False def __post_init__(self) -> None: if self.vocab_size <= 0: raise ValueError(f"vocab_size must be positive, got {self.vocab_size}") if self.max_seq_len <= 0: raise ValueError(f"max_seq_len must be positive, got {self.max_seq_len}") if self.pe_type not in SUPPORTED_PE_TYPES: raise ValueError(f"Unsupported pe_type: {self.pe_type}") if self.hidden_dim <= 0: raise ValueError(f"hidden_dim must be positive, got {self.hidden_dim}") if self.n_heads_sa <= 0: raise ValueError(f"n_heads_sa must be positive for DAT, got {self.n_heads_sa}") if self.n_heads_ra <= 0: raise ValueError(f"n_heads_ra must be positive for DAT, got {self.n_heads_ra}") total_heads = self.total_n_heads if self.hidden_dim % total_heads != 0: raise ValueError( f"hidden_dim ({self.hidden_dim}) must be divisible by total DAT heads " f"({total_heads} = {self.n_heads_sa} SA + {self.n_heads_ra} RA)" ) if self.n_layers <= 0: raise ValueError(f"n_layers must be positive, got {self.n_layers}") if not 0.0 <= self.dropout < 1.0: raise ValueError(f"dropout must be in [0.0, 1.0), got {self.dropout}") if self.dff_factor <= 0: raise ValueError(f"dff_factor must be positive, got {self.dff_factor}") if self.ffn_hidden_dim_mode not in SUPPORTED_FFN_HIDDEN_DIM_MODES: raise ValueError(f"Unsupported ffn_hidden_dim_mode for DAT: {self.ffn_hidden_dim_mode}") if self.ffn_activation not in SUPPORTED_FFN_ACTIVATIONS: raise ValueError(f"Unsupported ffn_activation for DAT: {self.ffn_activation}") if self.ffn_hidden_dim_mode == "swiglu_parameter_matched" and self.ffn_activation != "swiglu": raise ValueError( "ffn_hidden_dim_mode='swiglu_parameter_matched' requires " f"ffn_activation='swiglu', got {self.ffn_activation}" ) if self.rope_theta <= 0.0: raise ValueError(f"rope_theta must be positive, got {self.rope_theta}") if self.max_rel_pos is not None and self.max_rel_pos <= 0: raise ValueError(f"max_rel_pos must be positive when provided, got {self.max_rel_pos}") if self.sequence_boundary_policy not in SUPPORTED_SEQUENCE_BOUNDARY_POLICIES: raise ValueError( "Unsupported sequence_boundary_policy for DAT: " f"{self.sequence_boundary_policy}" ) if self.sequence_boundary_policy == "segment_document": if self.segment_boundary_token_id is None: raise ValueError( "segment_boundary_token_id is required when " "sequence_boundary_policy='segment_document'" ) if self.init_range <= 0.0: raise ValueError(f"init_range must be positive, got {self.init_range}") if self.init_scheme not in SUPPORTED_INIT_SCHEMES: raise ValueError(f"Unsupported init_scheme for DAT: {self.init_scheme}") if self.norm_type not in {"layernorm", "rmsnorm"}: raise ValueError( f"norm_type must be 'layernorm' or 'rmsnorm', got {self.norm_type}" ) if self.pe_type == "rope" and self.head_dim % 2 != 0: raise ValueError( "RoPE requires even DAT head_dim, got " f"{self.head_dim} from hidden_dim={self.hidden_dim}, total_heads={total_heads}" ) if self.pe_type == "sinusoidal" and self.hidden_dim % 2 != 0: raise ValueError(f"Sinusoidal encoding requires even hidden_dim, got {self.hidden_dim}") if self.share_attn_params and self.n_heads_sa != self.n_heads_ra: raise ValueError( "share_attn_params=True requires n_heads_sa == n_heads_ra, " f"got {self.n_heads_sa} and {self.n_heads_ra}" ) if self.symbol_dim is not None and self.symbol_dim <= 0: raise ValueError(f"symbol_dim must be positive when provided, got {self.symbol_dim}") if self.symbolic_attn_n_heads is not None and self.symbolic_attn_n_heads <= 0: raise ValueError( "symbolic_attn_n_heads must be positive when provided, " f"got {self.symbolic_attn_n_heads}" ) if self.symbol_retrieval not in SUPPORTED_SYMBOL_RETRIEVAL: raise ValueError(f"Unsupported symbol_retrieval for DAT: {self.symbol_retrieval}") if self.symbol_retrieval == "symbolic": symbolic_heads = self.resolved_symbolic_attn_n_heads if self.hidden_dim % symbolic_heads != 0: raise ValueError( f"hidden_dim ({self.hidden_dim}) must be divisible by symbolic_attn_n_heads " f"({symbolic_heads}) for symbolic retrieval" ) if self.resolved_symbol_dim % symbolic_heads != 0: raise ValueError( f"symbol_dim ({self.resolved_symbol_dim}) must be divisible by symbolic_attn_n_heads " f"({symbolic_heads}) for symbolic retrieval" ) if ( self.symbol_retrieval == "positional" and self.positional_symbols_sinusoidal and self.resolved_symbol_dim % 2 != 0 ): raise ValueError( "Sinusoidal positional symbols require even symbol_dim, " f"got {self.resolved_symbol_dim}" ) if self.relative_symbols_rope: if self.symbol_retrieval != "relative": raise ValueError( "relative_symbols_rope=True requires symbol_retrieval='relative', " f"got {self.symbol_retrieval!r}" ) if self.resolved_symbol_dim % 2 != 0: raise ValueError( "RoPE relative symbols require even symbol_dim, " f"got {self.resolved_symbol_dim}" ) if self.positional_symbols_sinusoidal and self.symbol_retrieval != "positional": raise ValueError( "positional_symbols_sinusoidal=True requires symbol_retrieval='positional', " f"got {self.symbol_retrieval!r}" ) if self.resolved_n_symbols <= 0: raise ValueError(f"resolved_n_symbols must be positive, got {self.resolved_n_symbols}") if self.symbol_retrieval == "relsymbolic": if self.relsymbolic_rel_n_heads <= 0: raise ValueError( f"relsymbolic_rel_n_heads must be positive, got {self.relsymbolic_rel_n_heads}" ) if self.hidden_dim % self.relsymbolic_rel_n_heads != 0: raise ValueError( f"hidden_dim ({self.hidden_dim}) must be divisible by " f"relsymbolic_rel_n_heads ({self.relsymbolic_rel_n_heads})" ) if self.relsymbolic_symbolic_attn_n_heads <= 0: raise ValueError( "relsymbolic_symbolic_attn_n_heads must be positive, " f"got {self.relsymbolic_symbolic_attn_n_heads}" ) if self.hidden_dim % self.relsymbolic_symbolic_attn_n_heads != 0: raise ValueError( f"hidden_dim ({self.hidden_dim}) must be divisible by " f"relsymbolic_symbolic_attn_n_heads ({self.relsymbolic_symbolic_attn_n_heads})" ) if self.resolved_symbol_dim % self.relsymbolic_symbolic_attn_n_heads != 0: raise ValueError( f"symbol_dim ({self.resolved_symbol_dim}) must be divisible by " f"relsymbolic_symbolic_attn_n_heads ({self.relsymbolic_symbolic_attn_n_heads})" ) if self.relsymbolic_neighborhood_size <= 0: raise ValueError( "relsymbolic_neighborhood_size must be positive, " f"got {self.relsymbolic_neighborhood_size}" ) if not 0.0 <= self.relsymbolic_dropout < 1.0: raise ValueError( f"relsymbolic_dropout must be in [0.0, 1.0), got {self.relsymbolic_dropout}" ) if self.relsymbolic_rel_scale is not None and self.relsymbolic_rel_scale <= 0.0: raise ValueError( "relsymbolic_rel_scale must be positive when provided, " f"got {self.relsymbolic_rel_scale}" ) if ( self.relsymbolic_symbolic_attn_scale is not None and self.relsymbolic_symbolic_attn_scale <= 0.0 ): raise ValueError( "relsymbolic_symbolic_attn_scale must be positive when provided, " f"got {self.relsymbolic_symbolic_attn_scale}" ) if self.ra_type not in SUPPORTED_RA_TYPES: raise ValueError(f"Unsupported ra_type for DAT: {self.ra_type}") if self.ra_n_relations is not None and self.ra_n_relations <= 0: raise ValueError( f"ra_n_relations must be positive when provided, got {self.ra_n_relations}" ) if self.ra_type != "ra" and self.ra_n_relations is not None: raise ValueError(f"ra_n_relations applies only to ra_type='ra', got {self.ra_type}") if self.ra_type != "ra" and self.ra_symmetric_rels: raise ValueError(f"ra_symmetric_rels applies only to ra_type='ra', got {self.ra_type}") n_relations = self.resolved_ra_n_relations if self.ra_type == "ra" and (self.head_dim * self.n_heads_ra) % n_relations != 0: raise ValueError( f"head_dim * n_heads_ra ({self.head_dim * self.n_heads_ra}) must be " f"divisible by ra_n_relations ({n_relations})" ) if self.ra_rel_activation not in SUPPORTED_RA_ACTIVATIONS: raise ValueError( f"Unsupported ra_rel_activation for DAT: {self.ra_rel_activation}" ) @property def total_n_heads(self) -> int: return self.n_heads_sa + self.n_heads_ra @property def head_dim(self) -> int: return self.hidden_dim // self.total_n_heads @property def resolved_symbol_dim(self) -> int: return self.hidden_dim if self.symbol_dim is None else self.symbol_dim @property def resolved_symbolic_attn_n_heads(self) -> int: return self.total_n_heads if self.symbolic_attn_n_heads is None else self.symbolic_attn_n_heads @property def resolved_ffn_hidden_dim(self) -> int: if self.ffn_hidden_dim_mode == "swiglu_parameter_matched": return int(8 / 3 * self.hidden_dim) return self.hidden_dim * self.dff_factor @property def resolved_n_symbols(self) -> int: return self.max_seq_len if self.n_symbols is None else self.n_symbols @property def resolved_ra_n_relations(self) -> int: return self.n_heads_ra if self.ra_n_relations is None else self.ra_n_relations