model: support JetVLM (#13289)

This commit is contained in:
Zijian Zhang
2025-11-18 12:02:03 +08:00
committed by GitHub
parent 9188feccca
commit aa8ecbda7a
9 changed files with 210 additions and 1 deletions
+2
View File
@@ -7,6 +7,7 @@ from sglang.srt.configs.exaone import ExaoneConfig
from sglang.srt.configs.falcon_h1 import FalconH1Config
from sglang.srt.configs.janus_pro import MultiModalityConfig
from sglang.srt.configs.jet_nemotron import JetNemotronConfig
from sglang.srt.configs.jet_vlm import JetVLMConfig
from sglang.srt.configs.kimi_linear import KimiLinearConfig
from sglang.srt.configs.kimi_vl import KimiVLConfig
from sglang.srt.configs.kimi_vl_moonvit import MoonViTConfig
@@ -40,4 +41,5 @@ __all__ = [
"FalconH1Config",
"NemotronHConfig",
"JetNemotronConfig",
"JetVLMConfig",
]
+53
View File
@@ -0,0 +1,53 @@
from typing import Any
from transformers.configuration_utils import PretrainedConfig
from transformers.models.siglip import SiglipVisionConfig
from sglang.srt.configs.jet_nemotron import JetNemotronConfig
from sglang.srt.configs.mamba_utils import Mamba2CacheParams
class JetVLMConfig(PretrainedConfig):
model_type = "jet_vlm"
sub_configs = {
"text_config": JetNemotronConfig,
"vision_config": SiglipVisionConfig,
}
_auto_class = "AutoConfig"
def __init__(
self,
*,
text_config: dict[str, Any] | None = None,
vision_config: dict[str, Any] | None = None,
image_token_id: int | None = None,
video_token_id: int | None = None,
**kwargs,
):
self.text_config = (
JetNemotronConfig(**text_config)
if text_config is not None
else JetNemotronConfig()
)
self.vision_config = (
SiglipVisionConfig(**vision_config)
if vision_config is not None
else SiglipVisionConfig()
)
self.image_token_id = image_token_id if image_token_id is not None else -1
self.video_token_id = video_token_id if video_token_id is not None else -1
super().__init__(**kwargs)
@property
def full_attention_layer_ids(self) -> list[int]:
return self.text_config.full_attention_layer_ids
@property
def linear_layer_ids(self) -> list[int]:
return self.text_config.linear_layer_ids
@property
def mamba2_cache_params(self) -> Mamba2CacheParams:
return self.text_config.mamba2_cache_params
@@ -955,6 +955,7 @@ multimodal_model_archs = [
"NVILAForConditionalGeneration",
"NVILALiteForConditionalGeneration",
"DeepseekOCRForCausalLM",
"JetVLMForConditionalGeneration",
]