| |
| |
| |
| |
| |
| |
| |
| |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from typing_extensions import Self, override |
|
|
| from coreai_models.models.base import BaseForCausalLM |
| from coreai_models.primitives._ops import mutable_slice_update |
| from coreai_models.primitives.macos.cache import KVCache |
| from coreai_models.primitives.macos.rms_norm import RMSNorm |
| from coreai_models.primitives.macos.rope import initialize_rope |
| from coreai_models.primitives.macos.sdpa import SDPA |
|
|
| USE_FUSED_KV = True |
|
|
|
|
| def lfm2_resolve_rope_theta(config) -> float: |
| """rope theta lives at config.rope_theta (lfm2) or config.rope_parameters (lfm2_moe).""" |
| rope_params = getattr(config, "rope_parameters", None) |
| if isinstance(rope_params, dict) and "rope_theta" in rope_params: |
| return float(rope_params["rope_theta"]) |
| return float(getattr(config, "rope_theta", 1000000.0)) |
|
|
|
|
| def lfm2_mlp_intermediate_size(config, intermediate_size: int) -> int: |
| """LFM2 dense checkpoints store the pre-adjust ff dim; apply the 2/3 SwiGLU adjust.""" |
| if getattr(config, "block_auto_adjust_ff_dim", False): |
| intermediate_size = int(2 * intermediate_size / 3) |
| multiplier = getattr(config, "block_ffn_dim_multiplier", None) |
| if multiplier is not None: |
| intermediate_size = int(multiplier * intermediate_size) |
| multiple_of = config.block_multiple_of |
| intermediate_size = multiple_of * ( |
| (intermediate_size + multiple_of - 1) // multiple_of |
| ) |
| return intermediate_size |
|
|
|
|
| class ConvState: |
| """Per-conv-layer rolling window of the last (L_cache - 1) gated inputs. |
| |
| Layout: (n_conv_layers, 1, conv_dim, L_cache - 1). Zeros mean "sequence start", |
| which reproduces the reference implementation's causal left-padding, so the same |
| graph serves prefill and decode. |
| """ |
|
|
| def __init__(self: Self, states: torch.Tensor) -> None: |
| self._states = states |
|
|
| def fetch(self: Self, conv_idx: int) -> torch.Tensor: |
| torch._check_is_size(conv_idx) |
| torch._check(conv_idx < self._states.size(0)) |
| return self._states.narrow(0, conv_idx, 1).squeeze(0) |
|
|
| def update(self: Self, conv_idx: int, new_state: torch.Tensor) -> None: |
| cache = self._states |
| torch._check_is_size(conv_idx) |
| torch._check(conv_idx < cache.size(0)) |
| begin_layer = torch.tensor((conv_idx,), dtype=torch.int32) |
| end_layer = torch.tensor((conv_idx + 1,), dtype=torch.int32) |
| zeros = [torch.tensor((0,), dtype=torch.int32) for _ in range(cache.dim() - 1)] |
| ends = [ |
| torch.tensor((cache.size(i),), dtype=torch.int32) for i in range(1, cache.dim()) |
| ] |
| mutable_slice_update( |
| x=cache, |
| update=new_state.unsqueeze(0), |
| begin=torch.concatenate([begin_layer, *zeros]), |
| end=torch.cat([end_layer, *ends]), |
| ) |
|
|
|
|
| class _ConvWeightHolder(nn.Module): |
| """Holds the depthwise kernel under the HF key `<layer>.conv.conv.weight` without |
| being an nn.Conv1d (keeps the tiny (D,1,L) kernel away from Linear-targeted quant).""" |
|
|
| def __init__(self, conv_dim: int, l_cache: int) -> None: |
| super().__init__() |
| self.weight = nn.Parameter(torch.empty(conv_dim, 1, l_cache)) |
|
|
|
|
| class ShortConv(nn.Module): |
| """LFM2 gated short conv: BCx = in_proj(x); Bx = B*x; y = C * causal_dwconv(Bx).""" |
|
|
| def __init__(self, config, conv_idx: int) -> None: |
| super().__init__() |
| self.conv_idx = conv_idx |
| dim = getattr(config, "conv_dim", config.hidden_size) |
| self.dim = dim |
| self.l_cache = config.conv_L_cache |
| bias = getattr(config, "conv_bias", False) |
| assert not bias, "conv_bias=True not wired (both LFM2.5 checkpoints use False)" |
| self.in_proj = nn.Linear(config.hidden_size, 3 * dim, bias=False) |
| self.out_proj = nn.Linear(dim, config.hidden_size, bias=False) |
| self.conv = _ConvWeightHolder(dim, self.l_cache) |
|
|
| def forward(self, x: torch.Tensor, conv_state: ConvState | None = None) -> torch.Tensor: |
| dim, l_cache = self.dim, self.l_cache |
| query_len = x.shape[1] |
| torch._check_is_size(query_len) |
|
|
| bcx = self.in_proj(x).transpose(1, 2) |
| b = bcx.narrow(1, 0, dim) |
| c = bcx.narrow(1, dim, dim) |
| xg = bcx.narrow(1, 2 * dim, dim) |
| bx = b * xg |
|
|
| if conv_state is not None: |
| past = conv_state.fetch(self.conv_idx) |
| full = torch.cat([past, bx], dim=-1) |
| conv_state.update(self.conv_idx, full.narrow(-1, query_len, l_cache - 1)) |
| else: |
| full = F.pad(bx, (l_cache - 1, 0)) |
|
|
| |
| |
| |
| |
| conv_out = F.conv1d(full, self.conv.weight, bias=None, groups=dim) |
|
|
| y = c * conv_out |
| return self.out_proj(y.transpose(1, 2)) |
|
|
|
|
| class Lfm2Attention(nn.Module): |
| """GQA attention with per-head q/k RMSNorm and RoPE; KV cache indexed by the |
| layer's ordinal among attention layers (attn_idx), not the global layer index.""" |
|
|
| def __init__(self, config, attn_idx: int) -> None: |
| super().__init__() |
| self.attn_idx = attn_idx |
|
|
| dim = config.hidden_size |
| self.n_heads = n_heads = config.num_attention_heads |
| self.n_kv_heads = n_kv_heads = config.num_key_value_heads |
| head_dim = getattr(config, "head_dim", None) |
| self.head_dim = head_dim = head_dim if head_dim else dim // n_heads |
|
|
| self.qkv_proj = nn.Linear( |
| dim, (n_heads + 2 * n_kv_heads) * head_dim, bias=False |
| ) |
| self.out_proj = nn.Linear(n_heads * head_dim, dim, bias=False) |
|
|
| eps = getattr(config, "norm_eps", 1e-5) |
| if USE_FUSED_KV: |
| self.qk_norm = RMSNorm(head_dim, eps=eps, n_heads=n_heads + n_kv_heads) |
| else: |
| self.q_layernorm = RMSNorm(head_dim, eps=eps) |
| self.k_layernorm = RMSNorm(head_dim, eps=eps) |
|
|
| self.sdpa = SDPA(is_causal=True, scale=head_dim**-0.5) |
| self.rope = initialize_rope(base=lfm2_resolve_rope_theta(config)) |
|
|
| def forward( |
| self, |
| x: torch.Tensor, |
| position_ids: torch.IntTensor, |
| cache: KVCache | None = None, |
| ) -> torch.Tensor: |
| batch_size, query_len, _ = x.shape |
| n_heads, n_kv_heads = self.n_heads, self.n_kv_heads |
|
|
| qkv = ( |
| self.qkv_proj(x) |
| .reshape(batch_size, query_len, n_heads + 2 * n_kv_heads, self.head_dim) |
| .permute(0, 2, 1, 3) |
| ) |
|
|
| if USE_FUSED_KV: |
| query_key = qkv.narrow(1, 0, n_heads + n_kv_heads) |
| else: |
| query = qkv.narrow(1, 0, n_heads) |
| key = qkv.narrow(1, n_heads, n_kv_heads) |
| value = qkv.narrow(1, n_heads + n_kv_heads, n_kv_heads) |
|
|
| if USE_FUSED_KV: |
| query_key = self.qk_norm(query_key) |
| else: |
| query = self.q_layernorm(query) |
| key = self.k_layernorm(key) |
|
|
| seq_len = position_ids.shape[-1] |
| torch._check_is_size(query_len) |
| torch._check_is_size(seq_len) |
| offset = seq_len - query_len |
| torch._check_is_size(offset) |
| rope_positions = position_ids.narrow(-1, offset, query_len) |
|
|
| if USE_FUSED_KV: |
| query_key = self.rope(query_key, position_ids=rope_positions) |
| query = query_key.narrow(1, 0, n_heads) |
| key = query_key.narrow(1, n_heads, n_kv_heads) |
| else: |
| query = self.rope(query, position_ids=rope_positions) |
| key = self.rope(key, position_ids=rope_positions) |
|
|
| if cache is not None: |
| key, value = cache.update_and_fetch( |
| self.attn_idx, offset, key, value, seq_len=seq_len, query_len=query_len |
| ) |
|
|
| output = ( |
| self.sdpa(query=query, key=key, value=value) |
| .permute(0, 2, 1, 3) |
| .reshape(batch_size, query_len, self.n_heads * self.head_dim) |
| ) |
| return self.out_proj(output) |
|
|
|
|
| class Lfm2MLP(nn.Module): |
| """SwiGLU MLP with LFM2's w1/w3/w2 naming (w1=gate, w3=up, w2=down).""" |
|
|
| def __init__(self, config, intermediate_size: int | None = None, auto_adjust: bool = True) -> None: |
| super().__init__() |
| hidden_size = config.hidden_size |
| inter = intermediate_size if intermediate_size else config.intermediate_size |
| if auto_adjust: |
| inter = lfm2_mlp_intermediate_size(config, inter) |
| self.w1 = nn.Linear(hidden_size, inter, bias=False) |
| self.w3 = nn.Linear(hidden_size, inter, bias=False) |
| self.w2 = nn.Linear(inter, hidden_size, bias=False) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return self.w2(F.silu(self.w1(x)) * self.w3(x)) |
|
|
|
|
| class Lfm2DecoderLayer(nn.Module): |
| def __init__(self, config, layer_idx: int, attn_idx: int, conv_idx: int) -> None: |
| super().__init__() |
| self.is_attention_layer = config.layer_types[layer_idx] == "full_attention" |
| if self.is_attention_layer: |
| self.self_attn = Lfm2Attention(config, attn_idx=attn_idx) |
| else: |
| self.conv = ShortConv(config, conv_idx=conv_idx) |
| self.feed_forward = self._build_feed_forward(config, layer_idx) |
| eps = getattr(config, "norm_eps", 1e-5) |
| self.operator_norm = RMSNorm(config.hidden_size, eps=eps) |
| self.ffn_norm = RMSNorm(config.hidden_size, eps=eps) |
|
|
| def _build_feed_forward(self, config, layer_idx: int) -> nn.Module: |
| return Lfm2MLP(config) |
|
|
| def forward( |
| self, |
| x: torch.Tensor, |
| position_ids: torch.IntTensor, |
| cache: KVCache | None = None, |
| conv_state: ConvState | None = None, |
| ) -> torch.Tensor: |
| if self.is_attention_layer: |
| r = self.self_attn(self.operator_norm(x), position_ids, cache) |
| else: |
| r = self.conv(self.operator_norm(x), conv_state) |
| h = x + r |
| return h + self.feed_forward(self.ffn_norm(h)) |
|
|
|
|
| def _layer_ordinals(config) -> list[tuple[int, int]]: |
| """Per global layer: (attn_idx, conv_idx) ordinals (the one not applicable = -1).""" |
| ordinals, attn_i, conv_i = [], 0, 0 |
| for lt in config.layer_types: |
| if lt == "full_attention": |
| ordinals.append((attn_i, -1)) |
| attn_i += 1 |
| else: |
| ordinals.append((-1, conv_i)) |
| conv_i += 1 |
| return ordinals |
|
|
|
|
| def num_attention_layers(config) -> int: |
| return sum(1 for lt in config.layer_types if lt == "full_attention") |
|
|
|
|
| def num_conv_layers(config) -> int: |
| return sum(1 for lt in config.layer_types if lt != "full_attention") |
|
|
|
|
| class Lfm2Model(nn.Module): |
| layer_cls = Lfm2DecoderLayer |
|
|
| def __init__(self, config) -> None: |
| super().__init__() |
| self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) |
| ordinals = _layer_ordinals(config) |
| self.layers = nn.ModuleList( |
| [ |
| self.layer_cls(config, layer_idx, attn_idx=a, conv_idx=c) |
| for layer_idx, (a, c) in enumerate(ordinals) |
| ] |
| ) |
| eps = getattr(config, "norm_eps", 1e-5) |
| self.embedding_norm = RMSNorm(config.hidden_size, eps=eps) |
|
|
| def forward( |
| self, |
| input_ids: torch.Tensor, |
| position_ids: torch.IntTensor, |
| cache: KVCache | None = None, |
| conv_state: ConvState | None = None, |
| ) -> torch.Tensor: |
| h = self.embed_tokens(input_ids) |
| for layer in self.layers: |
| h = layer(h, position_ids, cache, conv_state) |
| return self.embedding_norm(h) |
|
|
|
|
| def _fuse_lfm2_attention_weights(model, state_dict: dict[str, torch.Tensor]) -> None: |
| """Fuse q/k/v_proj -> qkv_proj and q/k_layernorm -> qk_norm, per attention layer.""" |
| for i, layer in enumerate(model.layers): |
| if not getattr(layer, "is_attention_layer", False): |
| continue |
| prefix = f"model.layers.{i}.self_attn" |
| combined = [] |
| for proj in ["q_proj", "k_proj", "v_proj"]: |
| key = f"{prefix}.{proj}.weight" |
| if key in state_dict: |
| combined.append(state_dict.pop(key)) |
| if combined: |
| state_dict[f"{prefix}.qkv_proj.weight"] = torch.concat(combined, axis=0) |
|
|
| if USE_FUSED_KV: |
| qn, kn = f"{prefix}.q_layernorm.weight", f"{prefix}.k_layernorm.weight" |
| if qn in state_dict and kn in state_dict: |
| attn = layer.self_attn |
| qw = state_dict.pop(qn).unsqueeze(0).unsqueeze(0) |
| kw = state_dict.pop(kn).unsqueeze(0).unsqueeze(0) |
| fused = torch.cat( |
| [ |
| qw.expand(attn.n_heads, 1, attn.head_dim), |
| kw.expand(attn.n_kv_heads, 1, attn.head_dim), |
| ], |
| dim=0, |
| ) |
| state_dict[f"{prefix}.qk_norm.weight"] = fused |
|
|
|
|
| def build_lfm2_reference_inputs(config, target_dtype, max_context_length, trace_query_len, trace_offset, trace_kv_seq_len): |
| """Reference inputs + dynamic shapes for the hybrid: KV cache sized to the number of |
| ATTENTION layers only, plus a static-shape conv state for the conv layers.""" |
| batch_size = 1 |
| input_ids = torch.randint( |
| 1, config.vocab_size, (batch_size, trace_query_len), dtype=torch.int32 |
| ) |
| position_ids = ( |
| torch.arange(trace_query_len + trace_offset, dtype=torch.int32) |
| .unsqueeze(0) |
| .expand(batch_size, trace_query_len + trace_offset) |
| ) |
| n_attn = num_attention_layers(config) |
| head_dim = getattr(config, "head_dim", None) or ( |
| config.hidden_size // config.num_attention_heads |
| ) |
| k_cache = torch.zeros( |
| n_attn, 1, config.num_key_value_heads, trace_kv_seq_len, head_dim, dtype=target_dtype |
| ) |
| v_cache = torch.zeros_like(k_cache) |
| conv_dim = getattr(config, "conv_dim", config.hidden_size) |
| conv_state = torch.zeros( |
| num_conv_layers(config), 1, conv_dim, config.conv_L_cache - 1, dtype=target_dtype |
| ) |
|
|
| reference_inputs = { |
| "input_ids": input_ids, |
| "position_ids": position_ids, |
| "k_cache": k_cache, |
| "v_cache": v_cache, |
| "conv_state": conv_state, |
| } |
| dynamic_shapes = { |
| "input_ids": {1: torch.export.Dim("seq_ids", max=max_context_length - 2)}, |
| "position_ids": { |
| 1: torch.export.Dim("seq_pos", min=trace_query_len, max=max_context_length - 1) |
| }, |
| |
| |
| |
| |
| "k_cache": ( |
| { |
| KVCache.seq_len_dim(): torch.export.Dim( |
| "k_seq_len", min=trace_kv_seq_len, max=max_context_length |
| ) |
| } |
| if max_context_length > trace_kv_seq_len |
| else None |
| ), |
| "v_cache": ( |
| { |
| KVCache.seq_len_dim(): torch.export.Dim( |
| "v_seq_len", min=trace_kv_seq_len, max=max_context_length |
| ) |
| } |
| if max_context_length > trace_kv_seq_len |
| else None |
| ), |
| "conv_state": None, |
| } |
| return reference_inputs, dynamic_shapes |
|
|
|
|
| class Lfm2ForCausalLM(BaseForCausalLM): |
| _HF_MODEL_CLASS = None |
|
|
| @override |
| def _init_model(self, config) -> None: |
| self.model = Lfm2Model(config) |
| self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) |
| if getattr(config, "tie_embedding", False) or getattr( |
| config, "tie_word_embeddings", False |
| ): |
| self.lm_head.weight = self.model.embed_tokens.weight |
|
|
| @BaseForCausalLM.cast_logits_bfloat16_to_float16 |
| def forward( |
| self, |
| input_ids: torch.Tensor, |
| position_ids: torch.IntTensor, |
| k_cache: torch.Tensor, |
| v_cache: torch.Tensor, |
| conv_state: torch.Tensor, |
| ) -> torch.Tensor: |
| cache = KVCache(k_cache, v_cache) |
| conv = ConvState(conv_state) |
| out = self.model(input_ids, position_ids, cache, conv) |
| return self.lm_head(out) |
|
|
| @override |
| def _mutate_state_dict(self: Self, state_dict: dict[str, torch.Tensor]) -> None: |
| _fuse_lfm2_attention_weights(self.model, state_dict) |
|
|
| def load_state_dict(self, state_dict, strict: bool = True, assign: bool = False): |
| result = super().load_state_dict(state_dict, strict=strict, assign=assign) |
| if getattr(self.config, "tie_embedding", False) or getattr( |
| self.config, "tie_word_embeddings", False |
| ): |
| self.lm_head.weight = self.model.embed_tokens.weight |
| return result |
|
|
| |
| @staticmethod |
| def state_names() -> tuple[str, ...]: |
| return ("k_cache", "v_cache", "conv_state") |
|
|
| @classmethod |
| def build_reference_inputs(cls, config, target_dtype, max_context_length, trace_query_len, trace_offset, trace_kv_seq_len): |
| return build_lfm2_reference_inputs( |
| config, target_dtype, max_context_length, trace_query_len, trace_offset, trace_kv_seq_len |
| ) |
|
|