diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index a2839a89d..08322cb8a 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -1706,6 +1706,7 @@ def is_generation_model(model_architectures: List[str], is_embedding: bool = Fal or "XLMRobertaModel" in model_architectures or "XLMRobertaForSequenceClassification" in model_architectures or "Gemma2ForSequenceClassification" in model_architectures + or "Lfm2BidirectionalModel" in model_architectures ): return False else: diff --git a/python/sglang/srt/models/lfm2.py b/python/sglang/srt/models/lfm2.py index d494c6a59..765bc6b39 100644 --- a/python/sglang/srt/models/lfm2.py +++ b/python/sglang/srt/models/lfm2.py @@ -32,8 +32,9 @@ from sglang.srt.layers.linear import ( RowParallelLinear, ) from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.layers.pooler import EmbeddingPoolerOutput, Pooler, PoolingType from sglang.srt.layers.quantization.base_config import QuantizationConfig -from sglang.srt.layers.radix_attention import RadixAttention +from sglang.srt.layers.radix_attention import AttentionType, RadixAttention from sglang.srt.layers.rotary_embedding import get_rope from sglang.srt.layers.vocab_parallel_embedding import ( ParallelLMHead, @@ -111,6 +112,7 @@ class Lfm2Attention(nn.Module): config: Lfm2Config, layer_id: int, quant_config: Optional[QuantizationConfig] = None, + attn_type: AttentionType = AttentionType.DECODER, prefix: str = "", ) -> None: super().__init__() @@ -167,6 +169,7 @@ class Lfm2Attention(nn.Module): scaling=self.scaling, num_kv_heads=self.num_local_kv_heads, layer_id=layer_id, + attn_type=attn_type, prefix=add_prefix("attn", prefix), ) @@ -304,6 +307,59 @@ class Lfm2ShortConv(nn.Module): return output +class Lfm2BidirectionalShortConv(Lfm2ShortConv): + """Same-padding non-causal short convolution for bidirectional LFM2 encoders.""" + + def forward( + self, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + if forward_batch.forward_mode.is_idle(): + return hidden_states + + proj, _ = self.in_proj(hidden_states) + B_gate, C_gate, x = proj.chunk(3, dim=-1) + Bx = B_gate * x + + seq_lens = forward_batch.extend_seq_lens_cpu + if seq_lens is None: + seq_lens = forward_batch.extend_seq_lens.detach().cpu().tolist() + max_len = max(seq_lens) + padded = Bx.new_zeros((len(seq_lens), max_len, Bx.shape[-1])) + + offset = 0 + for batch_idx, seq_len in enumerate(seq_lens): + seq_len = int(seq_len) + end = offset + seq_len + padded[batch_idx, :seq_len] = Bx[offset:end] + offset = end + + conv_out = F.conv1d( + padded.transpose(1, 2), + weight=self.conv_weight.unsqueeze(1), + bias=self.conv_bias, + padding=self.conv_kernel // 2, + groups=self.hidden_size_per_partition, + ) + if conv_out.shape[-1] > max_len: + conv_out = conv_out[..., :max_len] + elif conv_out.shape[-1] < max_len: + conv_out = F.pad(conv_out, (0, max_len - conv_out.shape[-1])) + conv_out = conv_out.transpose(1, 2) + + unpadded = Bx.new_empty(Bx.shape) + offset = 0 + for batch_idx, seq_len in enumerate(seq_lens): + seq_len = int(seq_len) + end = offset + seq_len + unpadded[offset:end] = conv_out[batch_idx, :seq_len] + offset = end + + output, _ = self.out_proj(C_gate * unpadded) + return output + + class Lfm2DecoderLayer(nn.Module): """Decoder layer - either attention or conv based on config.""" @@ -312,6 +368,7 @@ class Lfm2DecoderLayer(nn.Module): config: Lfm2Config, layer_id: int, quant_config: Optional[QuantizationConfig] = None, + bidirectional: bool = False, prefix: str = "", ): super().__init__() @@ -326,10 +383,16 @@ class Lfm2DecoderLayer(nn.Module): config=config, layer_id=layer_id, quant_config=quant_config, + attn_type=( + AttentionType.ENCODER_ONLY + if bidirectional + else AttentionType.DECODER + ), prefix=add_prefix("self_attn", prefix), ) else: - self.conv = Lfm2ShortConv( + conv_cls = Lfm2BidirectionalShortConv if bidirectional else Lfm2ShortConv + self.conv = conv_cls( config=config, layer_idx=layer_id, quant_config=quant_config, @@ -373,6 +436,7 @@ class Lfm2Model(nn.Module): self, config: Lfm2Config, quant_config: Optional[QuantizationConfig] = None, + bidirectional: bool = False, prefix: str = "", ): super().__init__() @@ -395,11 +459,12 @@ class Lfm2Model(nn.Module): config=config, layer_id=idx, quant_config=quant_config, + bidirectional=bidirectional, prefix=prefix, ) self.layers = make_layers( - config.num_hidden_layers, get_layer, prefix=f"{prefix}.layers" + config.num_hidden_layers, get_layer, prefix=add_prefix("layers", prefix) ) self.embedding_norm = RMSNorm(config.hidden_size, eps=config.norm_eps) @@ -427,6 +492,100 @@ class Lfm2Model(nn.Module): return self.embedding_norm(hidden_states) +class Lfm2BidirectionalModel(Lfm2Model): + """Bidirectional LFM2 encoder for dense embedding models.""" + + fall_back_to_pt_during_load = False + + def __init__( + self, + config: Lfm2Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__( + config, + quant_config=quant_config, + bidirectional=True, + prefix=prefix, + ) + self.pooler = Pooler(pooling_type=PoolingType.CLS, normalize=True) + + def get_num_kv_cache_layers(self) -> int: + return self.num_attention_layers + + def get_input_embeddings(self) -> nn.Embedding: + return self.embed_tokens + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: Optional[torch.Tensor] = None, + get_embedding: bool = True, + ) -> EmbeddingPoolerOutput: + assert get_embedding, f"{self.__class__.__name__} is only used for embedding" + hidden_states = super().forward( + input_ids=input_ids, + positions=positions, + forward_batch=forward_batch, + input_embeds=input_embeds, + ) + return self.pooler(hidden_states, forward_batch) + + def load_weights( + self, weights: Iterable[Tuple[str, torch.Tensor]], is_mtp: bool = False + ) -> Set[str]: + stacked_params_mapping = [ + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ] + + params_dict = dict(self.named_parameters()) + loaded_params: Set[str] = set() + + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + + # Handle conv weight/bias naming: HF uses conv.conv, we use conv_weight/conv_bias + if ".conv.conv.weight" in name: + name = name.replace(".conv.conv.weight", ".conv.conv_weight") + loaded_weight = loaded_weight.squeeze(1) # (D, 1, K) -> (D, K) + if ".conv.conv.bias" in name: + name = name.replace(".conv.conv.bias", ".conv.conv_bias") + + # Handle QKV stacking + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + name = name.replace(weight_name, param_name) + if name.endswith(".bias") and name not in params_dict: + break + if name not in params_dict: + break + param = params_dict[name] + weight_loader = getattr(param, "weight_loader") + weight_loader(param, loaded_weight, shard_id) + loaded_params.add(name) + break + else: + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + continue + + param = params_dict[name] + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight) + loaded_params.add(name) + + return loaded_params + + class Lfm2ForCausalLM(nn.Module): """LFM2 for causal language modeling with hybrid attention/conv architecture.""" @@ -538,4 +697,4 @@ class Lfm2ForCausalLM(nn.Module): return loaded_params -EntryClass = [Lfm2ForCausalLM] +EntryClass = [Lfm2ForCausalLM, Lfm2BidirectionalModel]