model: support LFM2-VL (Liquid Foundation Model 2 Vision-Language) (#21230)
Co-authored-by: Piotr Mazurek <piotr.mazurek@liquid.ai>
This commit is contained in:
co-authored by
Piotr Mazurek
parent
1fb4bf3558
commit
b5e8c4b9e3
@@ -17,6 +17,7 @@ from sglang.srt.configs.kimi_vl import KimiVLConfig
|
||||
from sglang.srt.configs.kimi_vl_moonvit import MoonViTConfig
|
||||
from sglang.srt.configs.lfm2 import Lfm2Config
|
||||
from sglang.srt.configs.lfm2_moe import Lfm2MoeConfig
|
||||
from sglang.srt.configs.lfm2_vl import Lfm2VlConfig
|
||||
from sglang.srt.configs.longcat_flash import LongcatFlashConfig
|
||||
from sglang.srt.configs.nano_nemotron_vl import NemotronH_Nano_VL_V2_Config
|
||||
from sglang.srt.configs.nemotron_h import NemotronHConfig
|
||||
@@ -56,6 +57,7 @@ __all__ = [
|
||||
"GraniteMoeHybridConfig",
|
||||
"Lfm2Config",
|
||||
"Lfm2MoeConfig",
|
||||
"Lfm2VlConfig",
|
||||
"NemotronHConfig",
|
||||
"NemotronH_Nano_VL_V2_Config",
|
||||
"JetNemotronConfig",
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
# Copyright 2026 Liquid AI. All rights reserved.
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""LFM2-VL (Liquid Foundation Model 2 Vision-Language) configuration"""
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from transformers import CONFIG_MAPPING
|
||||
from transformers import Lfm2VlConfig as HFLfm2VlConfig
|
||||
from transformers.utils import logging
|
||||
|
||||
from sglang.srt.configs.mamba_utils import Mamba2CacheParams, Mamba2StateShape
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
class Lfm2VlConfig(HFLfm2VlConfig):
|
||||
"""
|
||||
SGLang configuration for LFM2-VL models.
|
||||
|
||||
Extends HuggingFace's Lfm2VlConfig with hybrid model properties needed by SGLang.
|
||||
LFM2-VL combines:
|
||||
- SigLip2 vision encoder with NaFlex variable-resolution support
|
||||
- LFM2 language model with hybrid attention + short convolution
|
||||
- Multimodal projector with pixel unshuffle downsampling
|
||||
"""
|
||||
|
||||
@property
|
||||
def full_attention_layer_ids(self) -> List[int]:
|
||||
"""Return indices of attention layers for KV cache (from text_config)."""
|
||||
return [
|
||||
i
|
||||
for i, lt in enumerate(self.text_config.layer_types)
|
||||
if lt == "full_attention"
|
||||
]
|
||||
|
||||
@property
|
||||
def linear_layer_ids(self) -> List[int]:
|
||||
"""Return indices of conv layers for conv state cache (from text_config)."""
|
||||
return [
|
||||
i
|
||||
for i, lt in enumerate(self.text_config.layer_types)
|
||||
if lt in ("conv", "short_conv")
|
||||
]
|
||||
|
||||
@property
|
||||
def mamba_chunk_size(self) -> int:
|
||||
"""Return chunk size for Mamba2 backend. LFM2 doesn't use chunking, return 1."""
|
||||
return 1
|
||||
|
||||
@property
|
||||
def mamba2_cache_params(self) -> Optional[Mamba2CacheParams]:
|
||||
"""
|
||||
Get cache params for HybridReqToTokenPool initialization.
|
||||
|
||||
LFM2 uses ShortConv layers with a small fixed-size cache (kernel_size - 1).
|
||||
Unlike full Mamba2 models, LFM2 only uses the conv state, not SSM temporal state.
|
||||
"""
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||
|
||||
conv_layer_ids = self.linear_layer_ids
|
||||
if not conv_layer_ids:
|
||||
return None
|
||||
|
||||
hidden_size = self.text_config.hidden_size
|
||||
# conv_L_cache in config is kernel_size (e.g., 3)
|
||||
conv_kernel = int(self.text_config.conv_L_cache)
|
||||
|
||||
# get_attention_tp_size() requires initialization, default to 1 if not available
|
||||
try:
|
||||
tp_size = get_attention_tp_size()
|
||||
except (AssertionError, RuntimeError):
|
||||
tp_size = 1
|
||||
|
||||
# For ShortConv layers, we use a simplified Mamba2StateShape
|
||||
# LFM2 doesn't use SSM state (state_size=0), only conv state
|
||||
# We pass num_heads=tp_size so divide(tp_size, tp_size)=1 always works.
|
||||
# Since state_size=0, the temporal state shape has zero elements anyway.
|
||||
shape = Mamba2StateShape.create(
|
||||
tp_world_size=tp_size,
|
||||
intermediate_size=hidden_size,
|
||||
n_groups=1, # ShortConv doesn't use grouping
|
||||
num_heads=tp_size, # Ensures divide works; temporal state is empty anyway
|
||||
head_dim=hidden_size, # Conv operates on full hidden dim
|
||||
state_size=0, # No SSM temporal state for ShortConv
|
||||
conv_kernel=conv_kernel,
|
||||
)
|
||||
|
||||
# Uses default mamba2_state_dtype() which reads SGLANG_MAMBA_CONV_DTYPE env var
|
||||
# (defaults to bfloat16). Set SGLANG_MAMBA_CONV_DTYPE=float16 for fp16 inference.
|
||||
return Mamba2CacheParams(
|
||||
shape=shape,
|
||||
layers=conv_layer_ids,
|
||||
)
|
||||
|
||||
|
||||
# Override HuggingFace's Lfm2VlConfig with our extended version
|
||||
# Cannot use .register() because lfm2_vl may already be registered by transformers
|
||||
# Directly modify the internal _extra_content dict instead
|
||||
CONFIG_MAPPING._extra_content["lfm2_vl"] = Lfm2VlConfig
|
||||
@@ -1311,6 +1311,7 @@ multimodal_model_archs = [
|
||||
"LlavaQwenForCausalLM",
|
||||
"LlavaForConditionalGeneration",
|
||||
"LlavaVidForCausalLM",
|
||||
"Lfm2VlForConditionalGeneration",
|
||||
"LightOnOCRForConditionalGeneration",
|
||||
"MiniCPMO",
|
||||
"MiniCPMV",
|
||||
|
||||
@@ -42,6 +42,7 @@ from sglang.srt.configs import (
|
||||
KimiLinearConfig,
|
||||
Lfm2Config,
|
||||
Lfm2MoeConfig,
|
||||
Lfm2VlConfig,
|
||||
NemotronH_Nano_VL_V2_Config,
|
||||
NemotronHConfig,
|
||||
Qwen3_5Config,
|
||||
@@ -1851,7 +1852,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if pattern is not None and "M" not in pattern:
|
||||
return None
|
||||
if isinstance(
|
||||
config, FalconH1Config | NemotronHConfig | Lfm2Config | Lfm2MoeConfig
|
||||
config,
|
||||
FalconH1Config
|
||||
| NemotronHConfig
|
||||
| Lfm2Config
|
||||
| Lfm2MoeConfig
|
||||
| Lfm2VlConfig,
|
||||
):
|
||||
return config
|
||||
if isinstance(config, NemotronH_Nano_VL_V2_Config):
|
||||
|
||||
@@ -424,10 +424,10 @@ class Lfm2Model(nn.Module):
|
||||
input_ids: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
input_embeds: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
hidden_states = (
|
||||
inputs_embeds if inputs_embeds is not None else self.embed_tokens(input_ids)
|
||||
input_embeds if input_embeds is not None else self.embed_tokens(input_ids)
|
||||
)
|
||||
|
||||
residual = None
|
||||
@@ -474,16 +474,19 @@ class Lfm2ForCausalLM(nn.Module):
|
||||
def get_num_kv_cache_layers(self) -> int:
|
||||
return self.num_attention_layers
|
||||
|
||||
def get_input_embeddings(self) -> nn.Embedding:
|
||||
return self.model.embed_tokens
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
input_embeds: Optional[torch.Tensor] = None,
|
||||
**kwargs,
|
||||
):
|
||||
hidden_states = self.model(input_ids, positions, forward_batch, inputs_embeds)
|
||||
hidden_states = self.model(input_ids, positions, forward_batch, input_embeds)
|
||||
return self.logits_processor(
|
||||
input_ids, hidden_states, self.lm_head, forward_batch
|
||||
)
|
||||
|
||||
@@ -0,0 +1,348 @@
|
||||
# Copyright 2026 Liquid AI. All rights reserved.
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Inference-only LFM2-VL model compatible with HuggingFace weights.
|
||||
|
||||
LFM2-VL is a vision-language model that combines:
|
||||
- SigLip2 vision encoder with NaFlex variable-resolution support
|
||||
- LFM2 language model (hybrid attention + short convolution)
|
||||
- Multimodal projector with pixel unshuffle downsampling
|
||||
"""
|
||||
|
||||
import logging
|
||||
from typing import Iterable, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers.activations import ACT2FN
|
||||
|
||||
from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.managers.mm_utils import (
|
||||
MultiModalityDataPaddingPatternMultimodalTokens,
|
||||
general_mm_embed_routine,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
MultimodalDataItem,
|
||||
MultimodalInputs,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.lfm2 import Lfm2ForCausalLM
|
||||
from sglang.srt.models.siglip2 import Siglip2Model
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Lfm2VlMultiModalProjector(nn.Module):
|
||||
"""Multimodal projector with pixel unshuffle downsampling and TP/DP support."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
in_channels = config.vision_config.hidden_size * (config.downsample_factor**2)
|
||||
self.factor = config.downsample_factor
|
||||
self.use_layer_norm = config.projector_use_layernorm
|
||||
self.layer_norm = (
|
||||
nn.LayerNorm(in_channels) if config.projector_use_layernorm else None
|
||||
)
|
||||
|
||||
self.linear_1 = ColumnParallelLinear(
|
||||
in_channels,
|
||||
config.projector_hidden_size,
|
||||
bias=config.projector_bias,
|
||||
quant_config=quant_config,
|
||||
)
|
||||
self.act = ACT2FN[config.projector_hidden_act]
|
||||
self.linear_2 = RowParallelLinear(
|
||||
config.projector_hidden_size,
|
||||
config.text_config.hidden_size,
|
||||
bias=config.projector_bias,
|
||||
quant_config=quant_config,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
vision_features_packed: torch.Tensor,
|
||||
spatial_shapes: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Project packed vision features with pixel unshuffle.
|
||||
|
||||
Args:
|
||||
vision_features_packed: (total_tokens, hidden_size) packed in tile order.
|
||||
spatial_shapes: (num_tiles, 2) on CPU (height, width) per tile.
|
||||
|
||||
Returns:
|
||||
projected_packed: (total_projected_tokens, text_hidden_size)
|
||||
"""
|
||||
factor = self.factor
|
||||
hidden_size = vision_features_packed.shape[-1]
|
||||
|
||||
# Compute tile lengths from spatial shapes
|
||||
lengths = (spatial_shapes[:, 0] * spatial_shapes[:, 1]).tolist()
|
||||
|
||||
# Split packed tensor into per-tile tensors
|
||||
tile_features = torch.split(vision_features_packed, lengths, dim=0)
|
||||
|
||||
# Apply pixel unshuffle to each tile using reshape/permute (GPU operations)
|
||||
unshuffled_parts = []
|
||||
for tile, (h, w) in zip(tile_features, spatial_shapes.tolist()):
|
||||
if h == 0 or w == 0:
|
||||
continue
|
||||
# Reshape: (H*W, C) -> (H, W, C) -> (H/f, f, W/f, f, C)
|
||||
tile_2d = tile.view(h, w, hidden_size)
|
||||
tile_blocks = tile_2d.view(
|
||||
h // factor, factor, w // factor, factor, hidden_size
|
||||
)
|
||||
# Permute: (H/f, f, W/f, f, C) -> (H/f, W/f, f, f, C)
|
||||
tile_permuted = tile_blocks.permute(0, 2, 1, 3, 4)
|
||||
# Reshape: (H/f, W/f, f*f*C)
|
||||
tile_unshuffled = tile_permuted.reshape(
|
||||
(h // factor) * (w // factor), factor * factor * hidden_size
|
||||
)
|
||||
unshuffled_parts.append(tile_unshuffled)
|
||||
|
||||
if unshuffled_parts:
|
||||
unshuffled = torch.cat(unshuffled_parts, dim=0)
|
||||
else:
|
||||
unshuffled = vision_features_packed.new_empty(
|
||||
(0, factor * factor * hidden_size)
|
||||
)
|
||||
|
||||
if self.use_layer_norm:
|
||||
unshuffled = self.layer_norm(unshuffled)
|
||||
hidden_states, _ = self.linear_1(unshuffled)
|
||||
hidden_states = self.act(hidden_states)
|
||||
projected_packed, _ = self.linear_2(hidden_states)
|
||||
return projected_packed
|
||||
|
||||
|
||||
class Lfm2VlForConditionalGeneration(nn.Module):
|
||||
"""LFM2-VL Vision-Language Model."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.quant_config = quant_config
|
||||
|
||||
# Vision tower: Native Siglip2 implementation
|
||||
self.vision_tower = Siglip2Model(
|
||||
config=config.vision_config,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("vision_tower", prefix),
|
||||
)
|
||||
|
||||
# Multimodal projector
|
||||
self.multi_modal_projector = Lfm2VlMultiModalProjector(
|
||||
config,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("multi_modal_projector", prefix),
|
||||
)
|
||||
|
||||
# Language model: reuse SGLang's LFM2 implementation
|
||||
self.language_model = Lfm2ForCausalLM(
|
||||
config.text_config,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("language_model", prefix),
|
||||
)
|
||||
|
||||
self.logits_processor = LogitsProcessor(config.text_config)
|
||||
|
||||
def pad_input_ids(
|
||||
self, input_ids: List[int], mm_inputs: MultimodalInputs
|
||||
) -> List[int]:
|
||||
pattern = MultiModalityDataPaddingPatternMultimodalTokens()
|
||||
result = pattern.pad_input_tokens(input_ids, mm_inputs)
|
||||
return result
|
||||
|
||||
def get_input_embeddings(self) -> nn.Embedding:
|
||||
return self.language_model.model.embed_tokens
|
||||
|
||||
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
||||
"""Process images through vision tower and projector.
|
||||
|
||||
Handles SigLip2's NaFlex variable-resolution output.
|
||||
Pixel values arrive padded from the base processor; we pack them
|
||||
using the attention mask before feeding into the vision tower.
|
||||
"""
|
||||
# Collect data from all items
|
||||
all_pixel_values = []
|
||||
all_attention_masks = []
|
||||
all_spatial_shapes = []
|
||||
|
||||
for item in items:
|
||||
pv = item.feature
|
||||
am = item.pixel_attention_mask
|
||||
ss = item.spatial_shapes
|
||||
|
||||
if isinstance(pv, np.ndarray):
|
||||
pv = torch.from_numpy(pv)
|
||||
if isinstance(am, np.ndarray):
|
||||
am = torch.from_numpy(am)
|
||||
if isinstance(ss, np.ndarray):
|
||||
ss = torch.from_numpy(ss)
|
||||
|
||||
all_pixel_values.append(pv)
|
||||
all_attention_masks.append(am)
|
||||
all_spatial_shapes.append(ss)
|
||||
|
||||
pixel_values = torch.cat(all_pixel_values, dim=0)
|
||||
attention_mask = torch.cat(all_attention_masks, dim=0)
|
||||
spatial_shapes = torch.cat(all_spatial_shapes, dim=0)
|
||||
|
||||
pixel_values = pixel_values.to(
|
||||
device=self.vision_tower.device,
|
||||
dtype=self.vision_tower.dtype,
|
||||
)
|
||||
spatial_shapes_cpu = spatial_shapes.cpu()
|
||||
|
||||
# Pack padded pixel values using attention mask
|
||||
packed_list = []
|
||||
for i in range(pixel_values.shape[0]):
|
||||
mask = attention_mask[i].bool()
|
||||
packed_list.append(pixel_values[i][mask])
|
||||
|
||||
if not packed_list:
|
||||
return torch.tensor(
|
||||
[], device=self.vision_tower.device, dtype=self.vision_tower.dtype
|
||||
)
|
||||
|
||||
pixel_values_packed = torch.cat(packed_list, dim=0)
|
||||
|
||||
# Compute cu_seqlens and max_seqlen for packed attention
|
||||
spatial_shapes_list = spatial_shapes_cpu.tolist()
|
||||
lengths_list = [int(h * w) for h, w in spatial_shapes_list]
|
||||
total_tokens = sum(lengths_list)
|
||||
|
||||
if total_tokens == 0:
|
||||
return torch.tensor(
|
||||
[], device=self.vision_tower.device, dtype=self.vision_tower.dtype
|
||||
)
|
||||
|
||||
lengths = torch.tensor(
|
||||
lengths_list, dtype=torch.int32, device=pixel_values_packed.device
|
||||
)
|
||||
cu_seqlens = torch.zeros(
|
||||
len(lengths_list) + 1,
|
||||
dtype=torch.int32,
|
||||
device=pixel_values_packed.device,
|
||||
)
|
||||
cu_seqlens[1:] = torch.cumsum(lengths, dim=0)
|
||||
max_seqlen = lengths.max()
|
||||
|
||||
# Forward through vision tower
|
||||
vision_outputs = self.vision_tower(
|
||||
pixel_values_packed=pixel_values_packed,
|
||||
spatial_shapes=spatial_shapes_cpu,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
)
|
||||
|
||||
# Get the packed features (remove batch dim if present)
|
||||
if vision_outputs.dim() == 3:
|
||||
vision_features_packed = vision_outputs[0]
|
||||
else:
|
||||
vision_features_packed = vision_outputs
|
||||
|
||||
# Project through multimodal projector
|
||||
projected_packed = self.multi_modal_projector(
|
||||
vision_features_packed=vision_features_packed,
|
||||
spatial_shapes=spatial_shapes_cpu,
|
||||
)
|
||||
|
||||
return projected_packed
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
) -> torch.Tensor:
|
||||
return general_mm_embed_routine(
|
||||
input_ids=input_ids,
|
||||
forward_batch=forward_batch,
|
||||
language_model=self.language_model,
|
||||
multimodal_model=self,
|
||||
positions=positions,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||
"""Load weights from HuggingFace format."""
|
||||
# Collect weights by destination
|
||||
vision_weights = []
|
||||
projector_weights = []
|
||||
lm_weights = []
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
if name.startswith("model.vision_tower."):
|
||||
# model.vision_tower.* → * (strip model.vision_tower. prefix)
|
||||
# siglip2.py expects names like "vision_model.embeddings.patch_embedding.weight"
|
||||
new_name = name.replace("model.vision_tower.", "", 1)
|
||||
vision_weights.append((new_name, loaded_weight))
|
||||
elif name.startswith("model.multi_modal_projector."):
|
||||
# model.multi_modal_projector.* → multi_modal_projector.*
|
||||
new_name = name.replace(
|
||||
"model.multi_modal_projector.", "multi_modal_projector.", 1
|
||||
)
|
||||
projector_weights.append((new_name, loaded_weight))
|
||||
elif name.startswith("model.language_model."):
|
||||
# model.language_model.* → language_model.model.*
|
||||
new_name = name.replace(
|
||||
"model.language_model.", "language_model.model.", 1
|
||||
)
|
||||
lm_weights.append((new_name, loaded_weight))
|
||||
elif name.startswith("lm_head."):
|
||||
# lm_head.* → language_model.lm_head.*
|
||||
new_name = name.replace("lm_head.", "language_model.lm_head.", 1)
|
||||
lm_weights.append((new_name, loaded_weight))
|
||||
else:
|
||||
# Try direct mapping
|
||||
lm_weights.append((name, loaded_weight))
|
||||
|
||||
# Load vision tower weights using its own load_weights method
|
||||
self.vision_tower.load_weights(vision_weights)
|
||||
|
||||
# Load projector weights
|
||||
params_dict = dict(self.named_parameters())
|
||||
for name, loaded_weight in projector_weights:
|
||||
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)
|
||||
|
||||
# Load language model weights via Lfm2ForCausalLM.load_weights
|
||||
# Strip the "language_model." prefix since Lfm2ForCausalLM expects
|
||||
# names like "model.layers.0..." and "lm_head.weight"
|
||||
lm_weights_stripped = []
|
||||
for name, loaded_weight in lm_weights:
|
||||
if name.startswith("language_model."):
|
||||
name = name[len("language_model.") :]
|
||||
lm_weights_stripped.append((name, loaded_weight))
|
||||
self.language_model.load_weights(lm_weights_stripped)
|
||||
|
||||
|
||||
EntryClass = Lfm2VlForConditionalGeneration
|
||||
@@ -0,0 +1,584 @@
|
||||
# Copyright 2026 Liquid AI. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
#
|
||||
# Adapted from vLLM's implementation of Siglip2VisionModel
|
||||
# https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/models/lfm2_siglip2.py
|
||||
#
|
||||
# Siglip2 is a vision encoder that supports variable-resolution images via NaFlex.
|
||||
# Unlike Siglip v1 which uses fixed-size images, Siglip2 handles images of different
|
||||
# sizes by packing them into sequences and using cu_seqlens for attention.
|
||||
|
||||
from collections.abc import Iterable
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from transformers import Siglip2VisionConfig
|
||||
|
||||
from sglang.srt.layers.activation import get_act_fn
|
||||
from sglang.srt.layers.attention.vision import VisionAttention
|
||||
from sglang.srt.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
RowParallelLinear,
|
||||
)
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
|
||||
class Siglip2VisionEmbeddings(nn.Module):
|
||||
"""Siglip2 vision embeddings with NaFlex variable-resolution support."""
|
||||
|
||||
def __init__(self, config: Siglip2VisionConfig):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.embed_dim = config.hidden_size
|
||||
self.patch_size = config.patch_size
|
||||
|
||||
# Siglip2 uses Linear instead of Conv2d for patch embedding
|
||||
self.patch_embedding = nn.Linear(
|
||||
in_features=config.num_channels * self.patch_size * self.patch_size,
|
||||
out_features=self.embed_dim,
|
||||
)
|
||||
self.num_patches = config.num_patches
|
||||
self.position_embedding_size = int(self.num_patches**0.5)
|
||||
self.position_embedding = nn.Embedding(self.num_patches, self.embed_dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pixel_values_packed: torch.FloatTensor,
|
||||
spatial_shapes: torch.LongTensor,
|
||||
) -> torch.Tensor:
|
||||
"""Embed patchified pixel values in packed (unpadded) form.
|
||||
|
||||
Args:
|
||||
pixel_values_packed: (1, total_tokens, patch_dim) or
|
||||
(total_tokens, patch_dim), packed in tile order.
|
||||
spatial_shapes: (num_tiles, 2) on CPU (height, width) per tile.
|
||||
|
||||
Returns:
|
||||
(1, total_tokens, embed_dim) packed embeddings.
|
||||
"""
|
||||
assert spatial_shapes.device.type == "cpu", (
|
||||
"Expected `spatial_shapes` on CPU to avoid device-to-host sync in "
|
||||
"variable-length packing."
|
||||
)
|
||||
|
||||
if pixel_values_packed.dim() == 3:
|
||||
assert pixel_values_packed.shape[0] == 1
|
||||
pixel_values_flat = pixel_values_packed[0]
|
||||
else:
|
||||
pixel_values_flat = pixel_values_packed
|
||||
|
||||
lengths = (spatial_shapes[:, 0] * spatial_shapes[:, 1]).to(dtype=torch.int64)
|
||||
lengths_list = lengths.tolist()
|
||||
total_tokens = int(sum(lengths_list))
|
||||
if total_tokens != pixel_values_flat.shape[0]:
|
||||
raise ValueError(
|
||||
"Packed pixel_values token count does not match spatial_shapes: "
|
||||
f"{pixel_values_flat.shape[0]} vs {total_tokens}."
|
||||
)
|
||||
|
||||
target_dtype = self.patch_embedding.weight.dtype
|
||||
patch_embeds = self.patch_embedding(pixel_values_flat.to(dtype=target_dtype))
|
||||
|
||||
positional_embeddings = self.position_embedding.weight.reshape(
|
||||
self.position_embedding_size, self.position_embedding_size, -1
|
||||
)
|
||||
packed_pos_embeds = self.resize_positional_embeddings_packed(
|
||||
positional_embeddings,
|
||||
spatial_shapes,
|
||||
lengths_list=lengths_list,
|
||||
)
|
||||
|
||||
embeddings = patch_embeds + packed_pos_embeds
|
||||
return embeddings.unsqueeze(0)
|
||||
|
||||
@staticmethod
|
||||
def resize_positional_embeddings_packed(
|
||||
positional_embeddings: torch.Tensor,
|
||||
spatial_shapes: torch.LongTensor,
|
||||
lengths_list: list[int],
|
||||
) -> torch.Tensor:
|
||||
"""Resize positional embeddings per image and return a packed tensor.
|
||||
|
||||
Args:
|
||||
positional_embeddings: (height, width, embed_dim) base grid.
|
||||
spatial_shapes: (batch_size, 2) on CPU, (height, width) per image.
|
||||
lengths_list: flattened token length per image (height * width).
|
||||
|
||||
Returns:
|
||||
(total_tokens, embed_dim) packed positional embeddings.
|
||||
"""
|
||||
assert spatial_shapes.device.type == "cpu"
|
||||
|
||||
embed_dim = positional_embeddings.shape[-1]
|
||||
source_dtype = positional_embeddings.dtype
|
||||
|
||||
total_tokens = int(sum(lengths_list))
|
||||
packed_pos_embeds = torch.empty(
|
||||
(total_tokens, embed_dim),
|
||||
device=positional_embeddings.device,
|
||||
dtype=source_dtype,
|
||||
)
|
||||
|
||||
# (height, width, embed_dim) -> (1, embed_dim, height, width)
|
||||
pos_4d = positional_embeddings.permute(2, 0, 1).unsqueeze(0)
|
||||
|
||||
# Upcast to float32 on CPU because antialias is not supported for
|
||||
# bfloat16/float16 on CPU.
|
||||
if pos_4d.device.type == "cpu":
|
||||
pos_4d = pos_4d.to(torch.float32)
|
||||
|
||||
offset = 0
|
||||
for i, length in enumerate(lengths_list):
|
||||
if length <= 0:
|
||||
continue
|
||||
height, width = spatial_shapes[i].tolist()
|
||||
resized = F.interpolate(
|
||||
pos_4d,
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
align_corners=False,
|
||||
antialias=True,
|
||||
)
|
||||
resized = resized.reshape(embed_dim, height * width).transpose(0, 1)
|
||||
resized = resized.to(source_dtype)
|
||||
packed_pos_embeds[offset : offset + length] = resized
|
||||
offset += length
|
||||
|
||||
return packed_pos_embeds
|
||||
|
||||
|
||||
class Siglip2Attention(nn.Module):
|
||||
"""Multi-headed attention for Siglip2 using optimized VisionAttention backend."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Siglip2VisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.embed_dim = config.hidden_size
|
||||
self.num_heads = config.num_attention_heads
|
||||
self.head_dim = self.embed_dim // self.num_heads
|
||||
|
||||
if self.head_dim * self.num_heads != self.embed_dim:
|
||||
raise ValueError(
|
||||
f"embed_dim must be divisible by num_heads "
|
||||
f"(got `embed_dim`: {self.embed_dim} and `num_heads`:"
|
||||
f" {self.num_heads})."
|
||||
)
|
||||
|
||||
# Use SGLang's optimized VisionAttention with automatic backend selection
|
||||
self.attn = VisionAttention(
|
||||
embed_dim=self.embed_dim,
|
||||
num_heads=self.num_heads,
|
||||
projection_size=self.embed_dim,
|
||||
use_qkv_parallel=True,
|
||||
dropout=config.attention_dropout,
|
||||
flatten_batch=True, # For variable-length sequence support
|
||||
quant_config=quant_config,
|
||||
prefix=prefix,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: int | torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass with variable-length attention.
|
||||
|
||||
Args:
|
||||
hidden_states: (1, total_tokens, embed_dim) packed hidden states
|
||||
cu_seqlens: Cumulative sequence lengths for variable-length attention
|
||||
max_seqlen: Maximum sequence length (unused, VisionAttention computes internally)
|
||||
|
||||
Returns:
|
||||
(1, total_tokens, embed_dim) attention output
|
||||
"""
|
||||
return self.attn(hidden_states, cu_seqlens=cu_seqlens)
|
||||
|
||||
|
||||
class Siglip2MLP(nn.Module):
|
||||
"""MLP for Siglip2 encoder layers."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Siglip2VisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.activation_fn = get_act_fn(config.hidden_act)
|
||||
|
||||
self.fc1 = ColumnParallelLinear(
|
||||
config.hidden_size,
|
||||
config.intermediate_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("fc1", prefix),
|
||||
)
|
||||
self.fc2 = RowParallelLinear(
|
||||
config.intermediate_size,
|
||||
config.hidden_size,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("fc2", prefix),
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states, _ = self.fc1(hidden_states)
|
||||
hidden_states = self.activation_fn(hidden_states)
|
||||
hidden_states, _ = self.fc2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Siglip2EncoderLayer(nn.Module):
|
||||
"""Single encoder layer for Siglip2."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Siglip2VisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.embed_dim = config.hidden_size
|
||||
self.layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)
|
||||
self.self_attn = Siglip2Attention(
|
||||
config,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("self_attn", prefix),
|
||||
)
|
||||
self.layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps)
|
||||
self.mlp = Siglip2MLP(
|
||||
config,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("mlp", prefix),
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: int | torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass for encoder layer.
|
||||
|
||||
Args:
|
||||
hidden_states: Input tensor of shape (batch, seq_len, embed_dim).
|
||||
cu_seqlens: Cumulative sequence lengths tensor.
|
||||
max_seqlen: Maximum sequence length.
|
||||
"""
|
||||
residual = hidden_states
|
||||
|
||||
hidden_states = self.layer_norm1(hidden_states)
|
||||
hidden_states = self.self_attn(
|
||||
hidden_states=hidden_states,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
)
|
||||
hidden_states = residual + hidden_states
|
||||
|
||||
residual = hidden_states
|
||||
hidden_states = self.layer_norm2(hidden_states)
|
||||
hidden_states = self.mlp(hidden_states)
|
||||
hidden_states = residual + hidden_states
|
||||
return hidden_states
|
||||
|
||||
|
||||
class Siglip2Encoder(nn.Module):
|
||||
"""Transformer encoder for Siglip2."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Siglip2VisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
|
||||
if num_hidden_layers_override is None:
|
||||
num_hidden_layers = config.num_hidden_layers
|
||||
else:
|
||||
num_hidden_layers = num_hidden_layers_override
|
||||
|
||||
self.layers = nn.ModuleList(
|
||||
[
|
||||
Siglip2EncoderLayer(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix(f"layers.{idx}", prefix),
|
||||
)
|
||||
for idx in range(num_hidden_layers)
|
||||
]
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
inputs_embeds: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: int | torch.Tensor,
|
||||
return_all_hidden_states: bool = False,
|
||||
) -> torch.Tensor | list[torch.Tensor]:
|
||||
hidden_states_pool = [inputs_embeds]
|
||||
hidden_states = inputs_embeds
|
||||
|
||||
for encoder_layer in self.layers:
|
||||
hidden_states = encoder_layer(
|
||||
hidden_states,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
)
|
||||
if return_all_hidden_states:
|
||||
hidden_states_pool.append(hidden_states)
|
||||
if return_all_hidden_states:
|
||||
return hidden_states_pool
|
||||
return hidden_states
|
||||
|
||||
|
||||
def resolve_visual_encoder_outputs(
|
||||
encoder_outputs: torch.Tensor | list[torch.Tensor],
|
||||
post_layer_norm: Optional[nn.LayerNorm],
|
||||
select_layers: Optional[list[int]] = None,
|
||||
max_possible_layers: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Resolve outputs from visual encoder based on select_layers."""
|
||||
if select_layers is None:
|
||||
if isinstance(encoder_outputs, list):
|
||||
encoder_outputs = encoder_outputs[-1]
|
||||
if post_layer_norm is not None:
|
||||
encoder_outputs = post_layer_norm(encoder_outputs)
|
||||
return encoder_outputs
|
||||
|
||||
if max_possible_layers is None:
|
||||
raise ValueError(
|
||||
"`max_possible_layers` must be provided alongside `select_layers`"
|
||||
)
|
||||
|
||||
if not isinstance(encoder_outputs, list):
|
||||
raise ValueError(
|
||||
"Expected encoder_outputs to be a list when select_layers is provided"
|
||||
)
|
||||
|
||||
# Get the hidden states corresponding to the layer indices
|
||||
num_loaded_layers = len(encoder_outputs) - 1
|
||||
offset = max_possible_layers - num_loaded_layers
|
||||
hs_pool = [
|
||||
(
|
||||
encoder_outputs[layer_idx]
|
||||
if layer_idx >= 0
|
||||
else encoder_outputs[layer_idx + offset]
|
||||
)
|
||||
for layer_idx in select_layers
|
||||
]
|
||||
|
||||
uses_last_layer = select_layers[-1] in (max_possible_layers - 1, -1)
|
||||
if post_layer_norm is not None and uses_last_layer:
|
||||
hs_pool[-1] = post_layer_norm(hs_pool[-1])
|
||||
|
||||
return torch.cat(hs_pool, dim=-1)
|
||||
|
||||
|
||||
class Siglip2VisionTransformer(nn.Module):
|
||||
"""Siglip2 Vision Transformer with NaFlex variable-resolution support."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Siglip2VisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
require_post_norm: Optional[bool] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
embed_dim = config.hidden_size
|
||||
self.config = config
|
||||
self.embeddings = Siglip2VisionEmbeddings(config)
|
||||
self.encoder = Siglip2Encoder(
|
||||
config,
|
||||
quant_config=quant_config,
|
||||
num_hidden_layers_override=num_hidden_layers_override,
|
||||
prefix=add_prefix("encoder", prefix),
|
||||
)
|
||||
num_hidden_layers = config.num_hidden_layers
|
||||
if len(self.encoder.layers) > config.num_hidden_layers:
|
||||
raise ValueError(
|
||||
f"The original encoder only has {num_hidden_layers} "
|
||||
f"layers, but you requested {len(self.encoder.layers)} layers."
|
||||
)
|
||||
|
||||
if require_post_norm is None:
|
||||
require_post_norm = len(self.encoder.layers) == num_hidden_layers
|
||||
|
||||
if require_post_norm:
|
||||
self.post_layernorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
|
||||
else:
|
||||
self.post_layernorm = None
|
||||
|
||||
@property
|
||||
def dtype(self) -> torch.dtype:
|
||||
return self.embeddings.patch_embedding.weight.dtype
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return self.embeddings.patch_embedding.weight.device
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pixel_values_packed: torch.FloatTensor,
|
||||
spatial_shapes: torch.LongTensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: torch.Tensor,
|
||||
select_layers: Optional[list[int]] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass through the vision transformer.
|
||||
|
||||
Args:
|
||||
pixel_values_packed: Packed pixel values
|
||||
spatial_shapes: (batch_size, 2) tensor with (height, width) per image
|
||||
cu_seqlens: Cumulative sequence lengths
|
||||
max_seqlen: Maximum sequence length
|
||||
select_layers: Optional layer indices to select hidden states from
|
||||
|
||||
Returns:
|
||||
Vision features tensor
|
||||
"""
|
||||
hidden_states = self.embeddings(pixel_values_packed, spatial_shapes)
|
||||
|
||||
encoder_outputs = self.encoder(
|
||||
inputs_embeds=hidden_states,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
return_all_hidden_states=select_layers is not None,
|
||||
)
|
||||
|
||||
encoder_outputs = resolve_visual_encoder_outputs(
|
||||
encoder_outputs,
|
||||
self.post_layernorm,
|
||||
select_layers=select_layers,
|
||||
max_possible_layers=self.config.num_hidden_layers,
|
||||
)
|
||||
|
||||
return encoder_outputs
|
||||
|
||||
|
||||
class Siglip2Model(nn.Module):
|
||||
"""Siglip2 Vision Model for use in vision-language models."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Siglip2VisionConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
num_hidden_layers_override: Optional[int] = None,
|
||||
require_post_norm: Optional[bool] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.vision_model = Siglip2VisionTransformer(
|
||||
config,
|
||||
quant_config=quant_config,
|
||||
num_hidden_layers_override=num_hidden_layers_override,
|
||||
require_post_norm=require_post_norm,
|
||||
prefix=add_prefix("vision_model", prefix),
|
||||
)
|
||||
|
||||
@property
|
||||
def dtype(self) -> torch.dtype:
|
||||
return self.vision_model.dtype
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return self.vision_model.device
|
||||
|
||||
def forward(
|
||||
self,
|
||||
pixel_values_packed: torch.FloatTensor,
|
||||
spatial_shapes: torch.LongTensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: torch.Tensor,
|
||||
select_layers: Optional[list[int]] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Forward pass through the vision model."""
|
||||
return self.vision_model(
|
||||
pixel_values_packed=pixel_values_packed,
|
||||
spatial_shapes=spatial_shapes,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
select_layers=select_layers,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
# VisionAttention uses attn.qkv_proj for fused Q/K/V
|
||||
("attn.qkv_proj", "q_proj", "q"),
|
||||
("attn.qkv_proj", "k_proj", "k"),
|
||||
("attn.qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
# VisionAttention uses attn.proj instead of out_proj
|
||||
params_rename_mapping = {
|
||||
"out_proj": "attn.proj",
|
||||
}
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: set[str] = set()
|
||||
layer_count = len(self.vision_model.encoder.layers)
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
# post_layernorm is optional in Siglip2Model
|
||||
if (
|
||||
name.startswith("vision_model.post_layernorm")
|
||||
and self.vision_model.post_layernorm is None
|
||||
):
|
||||
continue
|
||||
|
||||
# omit layers when num_hidden_layers_override is set
|
||||
if name.startswith("vision_model.encoder.layers"):
|
||||
layer_idx = int(name.split(".")[3])
|
||||
if layer_idx >= layer_count:
|
||||
continue
|
||||
|
||||
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 not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
break
|
||||
else:
|
||||
# Apply rename mappings (e.g., out_proj -> attn.proj)
|
||||
for old_name, new_name in params_rename_mapping.items():
|
||||
if old_name in name:
|
||||
name = name.replace(old_name, new_name)
|
||||
break
|
||||
|
||||
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
|
||||
@@ -0,0 +1,84 @@
|
||||
# Copyright 2026 Liquid AI. All rights reserved.
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Multimodal processor for LFM2-VL models with SigLip2 NaFlex support."""
|
||||
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
from sglang.srt.managers.schedule_batch import Modality
|
||||
from sglang.srt.models.lfm2_vl import Lfm2VlForConditionalGeneration
|
||||
from sglang.srt.multimodal.processors.base_processor import (
|
||||
BaseMultimodalProcessor as SGLangBaseProcessor,
|
||||
)
|
||||
from sglang.srt.multimodal.processors.base_processor import (
|
||||
MultimodalSpecialTokens,
|
||||
)
|
||||
|
||||
|
||||
class Lfm2VlImageProcessor(SGLangBaseProcessor):
|
||||
"""Multimodal processor for LFM2-VL vision-language models.
|
||||
|
||||
Uses the base class load_mm_data + process_and_combine_mm_data flow.
|
||||
The HF processor handles NaFlex variable-resolution tiling internally.
|
||||
"""
|
||||
|
||||
models = [Lfm2VlForConditionalGeneration]
|
||||
|
||||
def __init__(self, hf_config, server_args, _processor, *args, **kwargs):
|
||||
super().__init__(hf_config, server_args, _processor, *args, **kwargs)
|
||||
|
||||
self.IMAGE_TOKEN_ID = hf_config.image_token_id
|
||||
self.IMAGE_TOKEN = "<image>"
|
||||
|
||||
self.mm_tokens = MultimodalSpecialTokens(
|
||||
image_token=self.IMAGE_TOKEN,
|
||||
image_token_id=hf_config.image_token_id,
|
||||
).build(_processor)
|
||||
|
||||
# Register NaFlex-specific HF processor outputs so
|
||||
# collect_mm_items_from_processor_output picks them up
|
||||
self.ATTR_NAME_TO_MODALITY["pixel_attention_mask"] = Modality.IMAGE
|
||||
self.ATTR_NAME_TO_MODALITY["spatial_shapes"] = Modality.IMAGE
|
||||
|
||||
async def process_mm_data_async(
|
||||
self,
|
||||
image_data: List[Union[str, bytes]],
|
||||
audio_data,
|
||||
input_text: str,
|
||||
request_obj,
|
||||
**kwargs,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
if not image_data:
|
||||
input_ids = self._tokenizer(
|
||||
input_text, return_tensors="pt", add_special_tokens=False
|
||||
).input_ids
|
||||
return {
|
||||
"input_ids": input_ids.squeeze(0).tolist(),
|
||||
"mm_items": [],
|
||||
"im_token_id": self.IMAGE_TOKEN_ID,
|
||||
}
|
||||
|
||||
base_output = self.load_mm_data(
|
||||
prompt=input_text,
|
||||
image_data=image_data,
|
||||
multimodal_tokens=self.mm_tokens,
|
||||
)
|
||||
|
||||
mm_items, input_ids, ret = self.process_and_combine_mm_data(
|
||||
base_output, self.mm_tokens
|
||||
)
|
||||
|
||||
return {
|
||||
"input_ids": input_ids.tolist(),
|
||||
"mm_items": mm_items,
|
||||
"im_token_id": self.IMAGE_TOKEN_ID,
|
||||
}
|
||||
Reference in New Issue
Block a user