[Spec] Add LFM2 and LFM2-MoE DSpark speculative decoding support (#31041)
Co-authored-by: Khoa Pham <khoa.pham@radixark.ai>
This commit is contained in:
co-authored by
Khoa Pham
parent
7700602278
commit
046454404a
@@ -21,6 +21,7 @@ def _fused_kv_norm_rope_write_kernel(
|
|||||||
L: tl.constexpr,
|
L: tl.constexpr,
|
||||||
EPS: tl.constexpr,
|
EPS: tl.constexpr,
|
||||||
HAS_COMMIT_LENS: tl.constexpr,
|
HAS_COMMIT_LENS: tl.constexpr,
|
||||||
|
IS_NEOX: tl.constexpr,
|
||||||
):
|
):
|
||||||
t = tl.program_id(0).to(tl.int64)
|
t = tl.program_id(0).to(tl.int64)
|
||||||
l = tl.program_id(1).to(tl.int64)
|
l = tl.program_id(1).to(tl.int64)
|
||||||
@@ -38,10 +39,18 @@ def _fused_kv_norm_rope_write_kernel(
|
|||||||
HALF: tl.constexpr = D // 2
|
HALF: tl.constexpr = D // 2
|
||||||
half_ar = tl.arange(0, HALF)
|
half_ar = tl.arange(0, HALF)
|
||||||
d_ar = tl.arange(0, D)
|
d_ar = tl.arange(0, D)
|
||||||
|
# Rotation pairing: neox pairs dims (i, i + D/2), interleaved (GPT-J style)
|
||||||
|
# pairs (2i, 2i + 1). The cos/sin cache layout is the same for both.
|
||||||
|
if IS_NEOX:
|
||||||
|
off1 = half_ar
|
||||||
|
off2 = HALF + half_ar
|
||||||
|
else:
|
||||||
|
off1 = 2 * half_ar
|
||||||
|
off2 = 2 * half_ar + 1
|
||||||
cos = tl.load(cos_sin_ptr + pos * D + half_ar).to(tl.float32)
|
cos = tl.load(cos_sin_ptr + pos * D + half_ar).to(tl.float32)
|
||||||
sin = tl.load(cos_sin_ptr + pos * D + HALF + half_ar).to(tl.float32)
|
sin = tl.load(cos_sin_ptr + pos * D + HALF + half_ar).to(tl.float32)
|
||||||
knw1 = tl.load(knw_ptr + l * D + half_ar).to(tl.float32)
|
knw1 = tl.load(knw_ptr + l * D + off1).to(tl.float32)
|
||||||
knw2 = tl.load(knw_ptr + l * D + HALF + half_ar).to(tl.float32)
|
knw2 = tl.load(knw_ptr + l * D + off2).to(tl.float32)
|
||||||
|
|
||||||
k_buf = tl.load(meta_ptr + l * 4 + 0).to(tl.pointer_type(tl.bfloat16))
|
k_buf = tl.load(meta_ptr + l * 4 + 0).to(tl.pointer_type(tl.bfloat16))
|
||||||
v_buf = tl.load(meta_ptr + l * 4 + 1).to(tl.pointer_type(tl.bfloat16))
|
v_buf = tl.load(meta_ptr + l * 4 + 1).to(tl.pointer_type(tl.bfloat16))
|
||||||
@@ -53,14 +62,14 @@ def _fused_kv_norm_rope_write_kernel(
|
|||||||
k = tl.load(row + h * D + d_ar).to(tl.float32)
|
k = tl.load(row + h * D + d_ar).to(tl.float32)
|
||||||
ms = tl.sum(k * k, 0) / D
|
ms = tl.sum(k * k, 0) / D
|
||||||
inv = 1.0 / tl.sqrt(ms + EPS)
|
inv = 1.0 / tl.sqrt(ms + EPS)
|
||||||
k1 = tl.load(row + h * D + half_ar).to(tl.float32) * inv * knw1
|
k1 = tl.load(row + h * D + off1).to(tl.float32) * inv * knw1
|
||||||
k2 = tl.load(row + h * D + HALF + half_ar).to(tl.float32) * inv * knw2
|
k2 = tl.load(row + h * D + off2).to(tl.float32) * inv * knw2
|
||||||
k1 = k1.to(tl.bfloat16).to(tl.float32)
|
k1 = k1.to(tl.bfloat16).to(tl.float32)
|
||||||
k2 = k2.to(tl.bfloat16).to(tl.float32)
|
k2 = k2.to(tl.bfloat16).to(tl.float32)
|
||||||
o1 = k1 * cos - k2 * sin
|
o1 = k1 * cos - k2 * sin
|
||||||
o2 = k2 * cos + k1 * sin
|
o2 = k2 * cos + k1 * sin
|
||||||
tl.store(k_buf + loc * ks0 + h * D + half_ar, o1.to(tl.bfloat16))
|
tl.store(k_buf + loc * ks0 + h * D + off1, o1.to(tl.bfloat16))
|
||||||
tl.store(k_buf + loc * ks0 + h * D + HALF + half_ar, o2.to(tl.bfloat16))
|
tl.store(k_buf + loc * ks0 + h * D + off2, o2.to(tl.bfloat16))
|
||||||
|
|
||||||
v = tl.load(row + KV + h * D + d_ar)
|
v = tl.load(row + KV + h * D + d_ar)
|
||||||
tl.store(v_buf + loc * vs0 + h * D + d_ar, v)
|
tl.store(v_buf + loc * vs0 + h * D + d_ar, v)
|
||||||
@@ -79,6 +88,7 @@ def fused_kv_norm_rope_write(
|
|||||||
eps: float,
|
eps: float,
|
||||||
commit_lens: Optional[torch.Tensor] = None,
|
commit_lens: Optional[torch.Tensor] = None,
|
||||||
locs_row_width: Optional[int] = None,
|
locs_row_width: Optional[int] = None,
|
||||||
|
is_neox_style: bool = True,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Write per-layer normed+roped K and raw V rows into the KV pools.
|
"""Write per-layer normed+roped K and raw V rows into the KV pools.
|
||||||
|
|
||||||
@@ -86,6 +96,8 @@ def fused_kv_norm_rope_write(
|
|||||||
flattened [bs, locs_row_width] verify window and only the first
|
flattened [bs, locs_row_width] verify window and only the first
|
||||||
commit_lens[b] columns of each row are written — the in-kernel
|
commit_lens[b] columns of each row are written — the in-kernel
|
||||||
replacement for masking the tail columns to -1 on the host.
|
replacement for masking the tail columns to -1 on the host.
|
||||||
|
|
||||||
|
is_neox_style selects the RoPE rotation pairing (neox or interleaved).
|
||||||
"""
|
"""
|
||||||
T = kv.shape[0]
|
T = kv.shape[0]
|
||||||
if T == 0:
|
if T == 0:
|
||||||
@@ -123,4 +135,5 @@ def fused_kv_norm_rope_write(
|
|||||||
L=num_layers,
|
L=num_layers,
|
||||||
EPS=eps,
|
EPS=eps,
|
||||||
HAS_COMMIT_LENS=has_commit_lens,
|
HAS_COMMIT_LENS=has_commit_lens,
|
||||||
|
IS_NEOX=is_neox_style,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -641,8 +641,6 @@ class DSparkDraftMixin:
|
|||||||
rotary = attn0.rotary_emb
|
rotary = attn0.rotary_emb
|
||||||
if type(rotary).__name__ != "RotaryEmbedding":
|
if type(rotary).__name__ != "RotaryEmbedding":
|
||||||
return None
|
return None
|
||||||
if not getattr(rotary, "is_neox_style", False):
|
|
||||||
return None
|
|
||||||
if getattr(rotary, "rotary_dim", None) != head_dim:
|
if getattr(rotary, "rotary_dim", None) != head_dim:
|
||||||
return None
|
return None
|
||||||
eps = attn0.k_norm.variance_epsilon
|
eps = attn0.k_norm.variance_epsilon
|
||||||
@@ -662,6 +660,8 @@ class DSparkDraftMixin:
|
|||||||
attn.rotary_emb.cos_sin_cache, rotary.cos_sin_cache
|
attn.rotary_emb.cos_sin_cache, rotary.cos_sin_cache
|
||||||
):
|
):
|
||||||
return None
|
return None
|
||||||
|
if attn.rotary_emb.is_neox_style != rotary.is_neox_style:
|
||||||
|
return None
|
||||||
if attn.k_norm.variance_epsilon != eps:
|
if attn.k_norm.variance_epsilon != eps:
|
||||||
return None
|
return None
|
||||||
k_buf = pool.get_key_buffer(attn.attn.layer_id)
|
k_buf = pool.get_key_buffer(attn.attn.layer_id)
|
||||||
@@ -769,6 +769,7 @@ class DSparkDraftMixin:
|
|||||||
eps,
|
eps,
|
||||||
commit_lens=write_commit_lens,
|
commit_lens=write_commit_lens,
|
||||||
locs_row_width=locs_row_width,
|
locs_row_width=locs_row_width,
|
||||||
|
is_neox_style=self.layers[0].self_attn.rotary_emb.is_neox_style,
|
||||||
)
|
)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|||||||
@@ -12,12 +12,15 @@ Uses optimized causal_conv1d kernels from the mamba package for fast inference.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
from typing import Iterable, Optional, Set, Tuple
|
from typing import TYPE_CHECKING, Iterable, List, Optional, Set, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch import nn
|
from torch import nn
|
||||||
|
|
||||||
|
from sglang.kernels.ops.mamba.causal_conv1d_triton import (
|
||||||
|
causal_conv1d_update as causal_conv1d_update_triton,
|
||||||
|
)
|
||||||
from sglang.srt.configs.lfm2 import Lfm2Config
|
from sglang.srt.configs.lfm2 import Lfm2Config
|
||||||
from sglang.srt.distributed import get_pp_group
|
from sglang.srt.distributed import get_pp_group
|
||||||
from sglang.srt.layers.attention.mamba.causal_conv1d import (
|
from sglang.srt.layers.attention.mamba.causal_conv1d import (
|
||||||
@@ -40,6 +43,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
ParallelLMHead,
|
ParallelLMHead,
|
||||||
VocabParallelEmbedding,
|
VocabParallelEmbedding,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.mem_cache.memory_pool import MambaPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||||
from sglang.srt.model_loader.weight_utils import (
|
from sglang.srt.model_loader.weight_utils import (
|
||||||
@@ -49,6 +53,11 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix, make_layers, set_weight_attrs
|
from sglang.srt.utils import add_prefix, make_layers, set_weight_attrs
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.attention.linear.short_conv_backend import (
|
||||||
|
ShortConvMetadata,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -203,6 +212,51 @@ class Lfm2Attention(nn.Module):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def register_shortconv_verify_buffers(conv: nn.Module) -> None:
|
||||||
|
"""Pre-size the tape slot-index buffer used by TARGET_VERIFY."""
|
||||||
|
conv.register_buffer(
|
||||||
|
"_intermediate_state_indices",
|
||||||
|
torch.arange(256, dtype=torch.int32),
|
||||||
|
persistent=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def shortconv_target_verify(
|
||||||
|
conv: nn.Module,
|
||||||
|
Bx: torch.Tensor,
|
||||||
|
meta: "ShortConvMetadata",
|
||||||
|
draft_token_num: int,
|
||||||
|
arch: str,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Depthwise short conv over a TARGET_VERIFY block.
|
||||||
|
|
||||||
|
Runs the triton update over the [bs, block] layout and records per-step conv
|
||||||
|
windows into the speculative tape so the state can roll back to the accept
|
||||||
|
boundary after verification.
|
||||||
|
"""
|
||||||
|
assert isinstance(meta.layer_cache, MambaPool.SpeculativeState), (
|
||||||
|
f"{arch} TARGET_VERIFY needs the speculative conv tape; the mamba pool "
|
||||||
|
"was built without speculative_num_draft_tokens."
|
||||||
|
)
|
||||||
|
bs = meta.cache_indices.shape[0]
|
||||||
|
Bx_reshaped = Bx.view(bs, draft_token_num, -1).transpose(1, 2)
|
||||||
|
if conv._intermediate_state_indices.shape[0] < bs:
|
||||||
|
conv._intermediate_state_indices = torch.arange(
|
||||||
|
bs, dtype=torch.int32, device=Bx.device
|
||||||
|
)
|
||||||
|
conv_out = causal_conv1d_update_triton(
|
||||||
|
Bx_reshaped,
|
||||||
|
meta.layer_cache.conv[0],
|
||||||
|
conv.conv_weight,
|
||||||
|
conv.conv_bias,
|
||||||
|
activation=None,
|
||||||
|
conv_state_indices=meta.cache_indices,
|
||||||
|
intermediate_conv_window=meta.layer_cache.intermediate_conv_window[0],
|
||||||
|
intermediate_state_indices=conv._intermediate_state_indices[:bs],
|
||||||
|
)
|
||||||
|
return conv_out.transpose(1, 2).reshape(bs * draft_token_num, -1)
|
||||||
|
|
||||||
|
|
||||||
class Lfm2ShortConv(nn.Module):
|
class Lfm2ShortConv(nn.Module):
|
||||||
"""
|
"""
|
||||||
Gated short convolution layer using optimized causal_conv1d kernels.
|
Gated short convolution layer using optimized causal_conv1d kernels.
|
||||||
@@ -260,6 +314,8 @@ class Lfm2ShortConv(nn.Module):
|
|||||||
else:
|
else:
|
||||||
self.register_parameter("conv_bias", None)
|
self.register_parameter("conv_bias", None)
|
||||||
|
|
||||||
|
register_shortconv_verify_buffers(self)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
@@ -289,6 +345,10 @@ class Lfm2ShortConv(nn.Module):
|
|||||||
activation=None,
|
activation=None,
|
||||||
conv_state_indices=meta.cache_indices,
|
conv_state_indices=meta.cache_indices,
|
||||||
)
|
)
|
||||||
|
elif forward_batch.forward_mode.is_target_verify():
|
||||||
|
conv_out = shortconv_target_verify(
|
||||||
|
self, Bx, meta, forward_batch.spec_info.draft_token_num, "LFM2"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# Prefill: multiple tokens, use varlen kernel
|
# Prefill: multiple tokens, use varlen kernel
|
||||||
Bx_t = Bx.transpose(0, 1).contiguous()
|
Bx_t = Bx.transpose(0, 1).contiguous()
|
||||||
@@ -374,6 +434,8 @@ class Lfm2DecoderLayer(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.layer_type = config.layer_types[layer_id]
|
self.layer_type = config.layer_types[layer_id]
|
||||||
self.is_attention_layer = self.layer_type == "full_attention"
|
self.is_attention_layer = self.layer_type == "full_attention"
|
||||||
|
# Set by Lfm2Model.set_dflash_layers_to_capture for DFlash aux capture.
|
||||||
|
self._is_layer_to_capture = False
|
||||||
|
|
||||||
self.operator_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
self.operator_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
||||||
self.ffn_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
self.ffn_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
||||||
@@ -412,9 +474,13 @@ class Lfm2DecoderLayer(nn.Module):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
residual: Optional[torch.Tensor],
|
residual: Optional[torch.Tensor],
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
|
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
if not forward_batch.forward_mode.is_idle():
|
if not forward_batch.forward_mode.is_idle():
|
||||||
|
if captured_last_layer_outputs is not None:
|
||||||
|
captured_last_layer_outputs.append(hidden_states)
|
||||||
|
|
||||||
residual = hidden_states
|
residual = hidden_states
|
||||||
normed = self.operator_norm(hidden_states)
|
normed = self.operator_norm(hidden_states)
|
||||||
|
|
||||||
@@ -467,6 +533,15 @@ class Lfm2Model(nn.Module):
|
|||||||
config.num_hidden_layers, get_layer, prefix=add_prefix("layers", prefix)
|
config.num_hidden_layers, get_layer, prefix=add_prefix("layers", prefix)
|
||||||
)
|
)
|
||||||
self.embedding_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
self.embedding_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
||||||
|
self.layers_to_capture: List[int] = []
|
||||||
|
|
||||||
|
def set_dflash_layers_to_capture(self, layers_to_capture: List[int]):
|
||||||
|
self.layers_to_capture = list(layers_to_capture)
|
||||||
|
for layer_id in self.layers_to_capture:
|
||||||
|
# A tap on the final layer (layer_id == len(self.layers)) is
|
||||||
|
# captured after the loop in forward(); only mark real layers.
|
||||||
|
if layer_id < len(self.layers):
|
||||||
|
self.layers[layer_id]._is_layer_to_capture = True
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -480,16 +555,30 @@ class Lfm2Model(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
residual = None
|
residual = None
|
||||||
|
aux_hidden_states: List[torch.Tensor] = []
|
||||||
for i in range(len(self.layers)):
|
for i in range(len(self.layers)):
|
||||||
hidden_states, residual = self.layers[i](
|
layer = self.layers[i]
|
||||||
|
hidden_states, residual = layer(
|
||||||
layer_id=i,
|
layer_id=i,
|
||||||
positions=positions,
|
positions=positions,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
residual=residual,
|
residual=residual,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
|
captured_last_layer_outputs=(
|
||||||
|
aux_hidden_states if layer._is_layer_to_capture else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
return self.embedding_norm(hidden_states)
|
if (
|
||||||
|
not forward_batch.forward_mode.is_idle()
|
||||||
|
and len(self.layers) in self.layers_to_capture
|
||||||
|
):
|
||||||
|
aux_hidden_states.append(hidden_states)
|
||||||
|
|
||||||
|
hidden_states = self.embedding_norm(hidden_states)
|
||||||
|
if not aux_hidden_states:
|
||||||
|
return hidden_states
|
||||||
|
return hidden_states, aux_hidden_states
|
||||||
|
|
||||||
|
|
||||||
class Lfm2BidirectionalModel(Lfm2Model):
|
class Lfm2BidirectionalModel(Lfm2Model):
|
||||||
@@ -620,6 +709,22 @@ class Lfm2ForCausalLM(nn.Module):
|
|||||||
def get_input_embeddings(self) -> nn.Embedding:
|
def get_input_embeddings(self) -> nn.Embedding:
|
||||||
return self.model.embed_tokens
|
return self.model.embed_tokens
|
||||||
|
|
||||||
|
@property
|
||||||
|
def capture_aux_hidden_states(self) -> bool:
|
||||||
|
return bool(self.model.layers_to_capture)
|
||||||
|
|
||||||
|
def set_dflash_layers_to_capture(self, layer_ids: List[int]):
|
||||||
|
if not self.pp_group.is_last_rank:
|
||||||
|
return
|
||||||
|
|
||||||
|
if layer_ids is None:
|
||||||
|
raise ValueError(
|
||||||
|
"DFLASH requires explicit layer_ids for aux hidden capture."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Mark layer L to capture its input, which is the output of layer L-1.
|
||||||
|
self.model.set_dflash_layers_to_capture([val + 1 for val in layer_ids])
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -630,8 +735,14 @@ class Lfm2ForCausalLM(nn.Module):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
||||||
|
aux_hidden_states = None
|
||||||
|
# Capture-enabled idle batches return a bare tensor (layers skip the
|
||||||
|
# append on IDLE), so narrow on the actual return shape.
|
||||||
|
if isinstance(hidden_states, tuple):
|
||||||
|
hidden_states, aux_hidden_states = hidden_states
|
||||||
|
|
||||||
return self.logits_processor(
|
return self.logits_processor(
|
||||||
input_ids, hidden_states, self.lm_head, forward_batch
|
input_ids, hidden_states, self.lm_head, forward_batch, aux_hidden_states
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_weights(
|
def load_weights(
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
"""LFM2-family DSpark draft model.
|
||||||
|
|
||||||
|
Its own arch so the LFM2 draft can later diverge from other DSpark drafts (e.g.
|
||||||
|
ShortConv layers) without touching shared classes. Current checkpoints are
|
||||||
|
attention-only Qwen3-style GQA with interleaved RoPE, so it is a thin
|
||||||
|
DSparkDraftModel subclass; config and weight names come from the checkpoint.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from sglang.srt.models.dspark import DSparkDraftModel
|
||||||
|
|
||||||
|
|
||||||
|
class Lfm2DSparkDraftModel(DSparkDraftModel):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
EntryClass = [Lfm2DSparkDraftModel]
|
||||||
@@ -12,7 +12,7 @@ Key MoE characteristics:
|
|||||||
- Post-hoc normalization of top-k weights
|
- Post-hoc normalization of top-k weights
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import Iterable, Optional, Set, Tuple
|
from typing import Iterable, List, Optional, Set, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
@@ -47,6 +47,10 @@ from sglang.srt.model_loader.weight_utils import (
|
|||||||
default_weight_loader,
|
default_weight_loader,
|
||||||
sharded_weight_loader,
|
sharded_weight_loader,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.models.lfm2 import (
|
||||||
|
register_shortconv_verify_buffers,
|
||||||
|
shortconv_target_verify,
|
||||||
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils import add_prefix, make_layers, set_weight_attrs
|
from sglang.srt.utils import add_prefix, make_layers, set_weight_attrs
|
||||||
|
|
||||||
@@ -320,6 +324,8 @@ class Lfm2MoeShortConv(nn.Module):
|
|||||||
else:
|
else:
|
||||||
self.register_parameter("conv_bias", None)
|
self.register_parameter("conv_bias", None)
|
||||||
|
|
||||||
|
register_shortconv_verify_buffers(self)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
@@ -347,6 +353,10 @@ class Lfm2MoeShortConv(nn.Module):
|
|||||||
activation=None,
|
activation=None,
|
||||||
conv_state_indices=meta.cache_indices,
|
conv_state_indices=meta.cache_indices,
|
||||||
)
|
)
|
||||||
|
elif forward_batch.forward_mode.is_target_verify():
|
||||||
|
conv_out = shortconv_target_verify(
|
||||||
|
self, Bx, meta, forward_batch.spec_info.draft_token_num, "LFM2-MoE"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
Bx_t = Bx.transpose(0, 1).contiguous()
|
Bx_t = Bx.transpose(0, 1).contiguous()
|
||||||
conv_out = causal_conv1d_fn(
|
conv_out = causal_conv1d_fn(
|
||||||
@@ -382,6 +392,8 @@ class Lfm2MoeDecoderLayer(nn.Module):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.layer_type = config.layer_types[layer_id]
|
self.layer_type = config.layer_types[layer_id]
|
||||||
self.is_attention_layer = self.layer_type == "full_attention"
|
self.is_attention_layer = self.layer_type == "full_attention"
|
||||||
|
# Set by Lfm2MoeModel.set_dflash_layers_to_capture for DFlash aux capture.
|
||||||
|
self._is_layer_to_capture = False
|
||||||
|
|
||||||
self.operator_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
self.operator_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
||||||
self.ffn_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
self.ffn_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
||||||
@@ -424,9 +436,13 @@ class Lfm2MoeDecoderLayer(nn.Module):
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
residual: Optional[torch.Tensor],
|
residual: Optional[torch.Tensor],
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
|
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
if not forward_batch.forward_mode.is_idle():
|
if not forward_batch.forward_mode.is_idle():
|
||||||
|
if captured_last_layer_outputs is not None:
|
||||||
|
captured_last_layer_outputs.append(hidden_states)
|
||||||
|
|
||||||
residual = hidden_states
|
residual = hidden_states
|
||||||
normed = self.operator_norm(hidden_states)
|
normed = self.operator_norm(hidden_states)
|
||||||
|
|
||||||
@@ -477,6 +493,15 @@ class Lfm2MoeModel(nn.Module):
|
|||||||
config.num_hidden_layers, get_layer, prefix=f"{prefix}.layers"
|
config.num_hidden_layers, get_layer, prefix=f"{prefix}.layers"
|
||||||
)
|
)
|
||||||
self.embedding_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
self.embedding_norm = RMSNorm(config.hidden_size, eps=config.norm_eps)
|
||||||
|
self.layers_to_capture: List[int] = []
|
||||||
|
|
||||||
|
def set_dflash_layers_to_capture(self, layers_to_capture: List[int]):
|
||||||
|
self.layers_to_capture = list(layers_to_capture)
|
||||||
|
for layer_id in self.layers_to_capture:
|
||||||
|
# A tap on the final layer (layer_id == len(self.layers)) is
|
||||||
|
# captured after the loop in forward(); only mark real layers.
|
||||||
|
if layer_id < len(self.layers):
|
||||||
|
self.layers[layer_id]._is_layer_to_capture = True
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -490,16 +515,30 @@ class Lfm2MoeModel(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
residual = None
|
residual = None
|
||||||
|
aux_hidden_states: List[torch.Tensor] = []
|
||||||
for i in range(len(self.layers)):
|
for i in range(len(self.layers)):
|
||||||
hidden_states, residual = self.layers[i](
|
layer = self.layers[i]
|
||||||
|
hidden_states, residual = layer(
|
||||||
layer_id=i,
|
layer_id=i,
|
||||||
positions=positions,
|
positions=positions,
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
residual=residual,
|
residual=residual,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
|
captured_last_layer_outputs=(
|
||||||
|
aux_hidden_states if layer._is_layer_to_capture else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
return self.embedding_norm(hidden_states)
|
if (
|
||||||
|
not forward_batch.forward_mode.is_idle()
|
||||||
|
and len(self.layers) in self.layers_to_capture
|
||||||
|
):
|
||||||
|
aux_hidden_states.append(hidden_states)
|
||||||
|
|
||||||
|
hidden_states = self.embedding_norm(hidden_states)
|
||||||
|
if not aux_hidden_states:
|
||||||
|
return hidden_states
|
||||||
|
return hidden_states, aux_hidden_states
|
||||||
|
|
||||||
|
|
||||||
class Lfm2MoeForCausalLM(nn.Module):
|
class Lfm2MoeForCausalLM(nn.Module):
|
||||||
@@ -535,6 +574,25 @@ class Lfm2MoeForCausalLM(nn.Module):
|
|||||||
def get_num_kv_cache_layers(self) -> int:
|
def get_num_kv_cache_layers(self) -> int:
|
||||||
return self.num_attention_layers
|
return self.num_attention_layers
|
||||||
|
|
||||||
|
def get_input_embeddings(self) -> nn.Embedding:
|
||||||
|
return self.model.embed_tokens
|
||||||
|
|
||||||
|
@property
|
||||||
|
def capture_aux_hidden_states(self) -> bool:
|
||||||
|
return bool(self.model.layers_to_capture)
|
||||||
|
|
||||||
|
def set_dflash_layers_to_capture(self, layer_ids: List[int]):
|
||||||
|
if not self.pp_group.is_last_rank:
|
||||||
|
return
|
||||||
|
|
||||||
|
if layer_ids is None:
|
||||||
|
raise ValueError(
|
||||||
|
"DFLASH requires explicit layer_ids for aux hidden capture."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Mark layer L to capture its input, which is the output of layer L-1.
|
||||||
|
self.model.set_dflash_layers_to_capture([val + 1 for val in layer_ids])
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -545,8 +603,14 @@ class Lfm2MoeForCausalLM(nn.Module):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
hidden_states = self.model(input_ids, positions, forward_batch, inputs_embeds)
|
hidden_states = self.model(input_ids, positions, forward_batch, inputs_embeds)
|
||||||
|
aux_hidden_states = None
|
||||||
|
# Capture-enabled idle batches return a bare tensor (layers skip the
|
||||||
|
# append on IDLE), so narrow on the actual return shape.
|
||||||
|
if isinstance(hidden_states, tuple):
|
||||||
|
hidden_states, aux_hidden_states = hidden_states
|
||||||
|
|
||||||
return self.logits_processor(
|
return self.logits_processor(
|
||||||
input_ids, hidden_states, self.lm_head, forward_batch
|
input_ids, hidden_states, self.lm_head, forward_batch, aux_hidden_states
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_weights(
|
def load_weights(
|
||||||
|
|||||||
Reference in New Issue
Block a user