from __future__ import annotations from typing import Any, Optional from transformers import PretrainedConfig _VALID_LM_HEAD_TYPES = {"replicate", "lane", "vocab_parallel"} _VALID_ROPE_SCALING_KEYS = { "scaling_factor", "low_freq_factor", "high_freq_factor", "original_max_position_embeddings", } def _infer_layer_types( *, num_hidden_layers: int, sliding_window: Optional[int], swa_layers: list[int], use_sliding_window: Optional[bool], ) -> list[str]: if ( sliding_window is None or int(sliding_window) <= 0 or use_sliding_window is False ): return ["full_attention"] * int(num_hidden_layers) sliding_layer_indices = {int(layer) for layer in swa_layers} return [ "sliding_attention" if layer_index in sliding_layer_indices else "full_attention" for layer_index in range(int(num_hidden_layers)) ] def _validate_positive_int(name: str, value: int) -> int: value = int(value) if value <= 0: raise ValueError(f"{name} must be a positive integer, got {value}") return value def _validate_nonnegative_int(name: str, value: int) -> int: value = int(value) if value < 0: raise ValueError(f"{name} must be >= 0, got {value}") return value def _validate_positive_float(name: str, value: float) -> float: value = float(value) if value <= 0.0: raise ValueError(f"{name} must be > 0, got {value}") return value def _validate_probability(name: str, value: float) -> float: value = float(value) if not 0.0 <= value < 1.0: raise ValueError(f"{name} must satisfy 0.0 <= {name} < 1.0, got {value}") return value def _validate_divisible(name: str, numerator: int, denominator: int) -> None: if int(numerator) % int(denominator) != 0: raise ValueError(f"{name} must be divisible by {denominator}, got {numerator}") def _is_power_of_two(value: int) -> bool: return value > 0 and (value & (value - 1)) == 0 def _is_rocm_torch_build() -> bool: try: import torch except ImportError: return False return getattr(torch.version, "hip", None) is not None def _validate_swa_layers(swa_layers: list[int], num_hidden_layers: int) -> None: for layer in swa_layers: if not 0 <= layer < num_hidden_layers: raise ValueError( "swa_layers must satisfy 0 <= layer < num_hidden_layers, " f"got layer {layer} with num_hidden_layers={num_hidden_layers}" ) def _validate_rope_scaling(rope_scaling: Optional[dict[str, Any]]) -> None: if rope_scaling is None: return if not isinstance(rope_scaling, dict): raise ValueError("rope_scaling must be a dictionary") unknown_keys = set(rope_scaling) - _VALID_ROPE_SCALING_KEYS if unknown_keys: raise ValueError(f"Unknown rope_scaling keys: {sorted(unknown_keys)}") scaling_factor = float(rope_scaling.get("scaling_factor", 1.0)) low_freq_factor = float(rope_scaling.get("low_freq_factor", 1.0)) high_freq_factor = float(rope_scaling.get("high_freq_factor", 1.0)) original_max_position_embeddings = int( rope_scaling.get("original_max_position_embeddings", 4096) ) if scaling_factor <= 0: raise ValueError("rope_scaling.scaling_factor must be > 0") if low_freq_factor <= 0 or high_freq_factor <= 0: raise ValueError("rope_scaling frequency factors must be > 0") if high_freq_factor == low_freq_factor: raise ValueError( "rope_scaling.high_freq_factor and low_freq_factor must differ" ) if original_max_position_embeddings <= 0: raise ValueError("rope_scaling.original_max_position_embeddings must be > 0") class LaneformerConfig(PretrainedConfig): model_type = "laneformer" keys_to_ignore_at_inference = ["past_key_values"] def __init__( self, hidden_size: int = 4096, num_hidden_layers: int = 32, num_attention_heads: int = 32, num_key_value_heads: Optional[int] = None, intermediate_size: int = 4 * 4096, max_position_embeddings: int = 131072, eos_token_id: int = 0, pad_token_id: Optional[int] = None, vocab_size: int = 128256, rope_theta: float = 10000.0, norm_eps: float = 1e-5, multiple_of: int = 32, depth_init: bool = True, use_flex_attn: bool = False, attn_mask_type: str = "causal", initializer_range: float = 0.02, attention_dropout: float = 0.0, use_cache: bool = False, # laneformer-specific extras (HF is fine with extra fields): num_lanes: int = 1, broadcast_delay: int = 1, use_attention_comm: bool = True, use_mlp_comm: bool = True, use_early_comm: bool = True, lm_head_type: str = "replicate", pre_norm_lane_agg: bool = False, replicated_rmsn_scale: bool = True, tie_word_embeddings: bool = False, sliding_window: Optional[int] = None, sliding_window_n_layers: int = 0, swa_layers: Optional[list[int]] = None, rope_scaling: Optional[dict[str, Any]] = None, **kwargs, ): if pad_token_id is None: pad_token_id = eos_token_id if num_key_value_heads is None: num_key_value_heads = num_attention_heads hidden_size = _validate_positive_int("hidden_size", hidden_size) num_hidden_layers = _validate_positive_int( "num_hidden_layers", num_hidden_layers ) num_attention_heads = _validate_positive_int( "num_attention_heads", num_attention_heads ) num_key_value_heads = _validate_positive_int( "num_key_value_heads", num_key_value_heads ) intermediate_size = _validate_positive_int( "intermediate_size", intermediate_size ) max_position_embeddings = _validate_positive_int( "max_position_embeddings", max_position_embeddings ) vocab_size = _validate_positive_int("vocab_size", vocab_size) num_lanes = _validate_positive_int("num_lanes", num_lanes) rope_theta = _validate_positive_float("rope_theta", rope_theta) norm_eps = _validate_positive_float("norm_eps", norm_eps) multiple_of = _validate_positive_int("multiple_of", multiple_of) attention_dropout = _validate_probability( "attention_dropout", attention_dropout ) broadcast_delay = _validate_nonnegative_int("broadcast_delay", broadcast_delay) sliding_window_n_layers = _validate_nonnegative_int( "sliding_window_n_layers", sliding_window_n_layers ) if sliding_window_n_layers > num_hidden_layers: raise ValueError( "sliding_window_n_layers must satisfy " "0 <= sliding_window_n_layers <= num_hidden_layers, " f"got {sliding_window_n_layers} with " f"num_hidden_layers={num_hidden_layers}" ) _validate_divisible("hidden_size", hidden_size, num_attention_heads) _validate_divisible( "num_attention_heads", num_attention_heads, num_key_value_heads ) _validate_divisible("hidden_size", hidden_size, num_lanes) _validate_divisible("intermediate_size", intermediate_size, num_lanes) head_dim = hidden_size // num_attention_heads if head_dim % 2 != 0: raise ValueError( "hidden_size // num_attention_heads must be even for RoPE, " f"got {head_dim}" ) _validate_divisible( "num_key_value_heads * head_dim", num_key_value_heads * head_dim, num_lanes, ) if lm_head_type not in _VALID_LM_HEAD_TYPES: raise ValueError( f"lm_head_type must be one of {sorted(_VALID_LM_HEAD_TYPES)}, " f"got {lm_head_type!r}" ) if lm_head_type == "vocab_parallel": _validate_divisible("vocab_size", vocab_size, num_lanes) if sliding_window is not None: sliding_window = _validate_positive_int("sliding_window", sliding_window) if tie_word_embeddings: raise ValueError("tie_word_embeddings is not supported for Laneformer") _validate_rope_scaling(rope_scaling) if swa_layers is None: swa_layers = list(range(sliding_window_n_layers)) else: swa_layers = sorted({int(layer) for layer in swa_layers}) _validate_swa_layers(swa_layers, num_hidden_layers) super().__init__( eos_token_id=eos_token_id, pad_token_id=pad_token_id, vocab_size=vocab_size, tie_word_embeddings=tie_word_embeddings, **kwargs, ) if ( self._attn_implementation == "flex_attention" and _is_rocm_torch_build() and not _is_power_of_two(head_dim) ): self._attn_implementation = "sdpa" self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.num_key_value_heads = num_key_value_heads self.intermediate_size = intermediate_size self.max_position_embeddings = max_position_embeddings self.rope_theta = rope_theta self.norm_eps = norm_eps self.multiple_of = multiple_of self.depth_init = depth_init self.use_flex_attn = use_flex_attn self.attn_mask_type = attn_mask_type self.initializer_range = initializer_range self.attention_dropout = attention_dropout self.use_cache = use_cache self.sliding_window = sliding_window self.rope_scaling = rope_scaling self.swa_layers = swa_layers # legacy field kept for backward compatibility self.sliding_window_n_layers = len(self.swa_layers) if getattr(self, "layer_types", None) is None: self.layer_types = _infer_layer_types( num_hidden_layers=num_hidden_layers, sliding_window=sliding_window, swa_layers=self.swa_layers, use_sliding_window=getattr(self, "use_sliding_window", None), ) # custom fields self.num_lanes = num_lanes self.use_attention_comm = use_attention_comm self.use_mlp_comm = use_mlp_comm # legacy aggregate flag kept for compatibility self.use_comm = bool(self.use_attention_comm and self.use_mlp_comm) self.use_early_comm = use_early_comm self.broadcast_delay = broadcast_delay self.lm_head_type = lm_head_type self.pre_norm_lane_agg = pre_norm_lane_agg self.replicated_rmsn_scale = replicated_rmsn_scale