Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
70 commits
Select commit Hold shift + click to select a range
bdaf2e2
dsv4.1: extract Top-k kernels and candidate helpers
hnyls2002 Sep 15, 2026
9b17c0d
dsv4.1: skip unwritten Top-k plans in metadata comparisons
hnyls2002 Sep 16, 2026
f51894f
fix dsv4 top-k edge cases and trim tests
BBuf Sep 16, 2026
ed1c649
dsv4.1: extract communication kernels and wrappers
hnyls2002 Sep 15, 2026
5873fba
dsv4.1: extract vocab gather and sharded greedy selection
hnyls2002 Sep 15, 2026
8c02375
fix sharded greedy test suite
hnyls2002 Sep 16, 2026
b297ac2
clarify sharded greedy docstring; split test cases; fix stale usage path
hnyls2002 Sep 16, 2026
d3d9a85
fix communication guards and cover ragged collectives
BBuf Sep 16, 2026
5a8166b
remove standalone nvlink communication test
BBuf Sep 16, 2026
3c5130f
dsv4.1: extract compression and metadata kernels
hnyls2002 Sep 15, 2026
b14c741
dsv4.1: extract KV store and dequantization paths
hnyls2002 Sep 15, 2026
a0ae65a
dsv4.1: preserve 64K prefill planner indices and reject sentinel coll…
hnyls2002 Sep 15, 2026
e36a479
fix c2 padding and trim metadata tests
BBuf Sep 16, 2026
1b94c71
dsv4.1: restore C2 padding test entry point
hnyls2002 Sep 16, 2026
ed29816
remove standalone dsv4 metadata tests
BBuf Sep 16, 2026
1333f2a
dsv4.1: extract candidate indexer library
hnyls2002 Sep 15, 2026
dee8e7d
dsv4.1: extract RoPE and FP4 packing kernels
hnyls2002 Sep 15, 2026
e95bed5
dsv4.1: extract Hopper FP8 matmul kernels and tuning
hnyls2002 Sep 15, 2026
c4bb0d6
dsv4.1: extract mHC computation and compensated projections
hnyls2002 Sep 15, 2026
f41fc30
dsv4.1: extract vision tower and image preprocessing
hnyls2002 Sep 15, 2026
cde9da0
merge main; drop stale topk lineage
hnyls2002 Sep 16, 2026
313f98d
merge dsv4.1-communication
hnyls2002 Sep 16, 2026
0eb0eb8
merge dsv4.1-metadata
hnyls2002 Sep 16, 2026
cc33d85
merge dsv4.1-candidate
hnyls2002 Sep 16, 2026
4cdd8d0
merge dsv4.1-rope-fp4
hnyls2002 Sep 16, 2026
bc686eb
merge dsv4.1-hopper
hnyls2002 Sep 16, 2026
800c056
merge dsv4.1-mhc
hnyls2002 Sep 16, 2026
0a66478
remove standalone sparse indexer test
BBuf Sep 16, 2026
8295eb6
Merge dsv4.1-candidate test cleanup into rope stack
BBuf Sep 16, 2026
294d3e8
remove standalone compressed KV quant test
BBuf Sep 16, 2026
04e98e9
Merge dsv4.1-rope-fp4 test cleanup into hopper stack
BBuf Sep 16, 2026
f4cff50
remove standalone Hopper FP8 test
BBuf Sep 16, 2026
b9a3c5c
Merge dsv4.1-hopper test cleanup into mHC stack
BBuf Sep 16, 2026
e4a3c13
remove standalone compensated mHC test
BBuf Sep 16, 2026
53cb501
Merge dsv4.1-mhc test cleanup into vision stack
BBuf Sep 16, 2026
be34bc2
merge main
hnyls2002 Sep 16, 2026
02076c3
merge metadata
hnyls2002 Sep 16, 2026
f8979c3
merge candidate
hnyls2002 Sep 16, 2026
e40b3f1
merge hopper
hnyls2002 Sep 16, 2026
e3e4652
merge rope-fp4
hnyls2002 Sep 16, 2026
6af53ee
merge mhc
hnyls2002 Sep 16, 2026
56efc4a
drop norm-only c2 entry and aliases; move small metadata into dsv4; n…
hnyls2002 Sep 16, 2026
c888f9e
drop unused fp4 torch reference quantizer
hnyls2002 Sep 16, 2026
f33acae
c2: drop the dead duplicate freqs_cis load
BBuf Sep 16, 2026
922b0d5
mhc: stop allocating a throwaway sqrsum for the residual projection
BBuf Sep 16, 2026
8af4e43
merge c1/c2 wrappers; split small metadata into its homes; move torch…
hnyls2002 Sep 16, 2026
b32b53d
merge metadata
hnyls2002 Sep 16, 2026
68ce340
merge candidate
hnyls2002 Sep 16, 2026
3e0335f
merge rope-fp4
hnyls2002 Sep 16, 2026
81016b0
merge hopper
hnyls2002 Sep 16, 2026
d32bb43
merge mhc
hnyls2002 Sep 16, 2026
fe06b61
merge main
hnyls2002 Sep 16, 2026
90723f1
merge candidate
hnyls2002 Sep 16, 2026
1989687
merge rope-fp4
hnyls2002 Sep 16, 2026
9bdd420
merge hopper
hnyls2002 Sep 16, 2026
97c689d
merge mhc
hnyls2002 Sep 16, 2026
b1be221
merge main
hnyls2002 Sep 16, 2026
ae963bc
merge dsv4.1-rope-fp4
hnyls2002 Sep 16, 2026
895f441
merge dsv4.1-hopper
hnyls2002 Sep 16, 2026
b7de200
merge dsv4.1-mhc
hnyls2002 Sep 16, 2026
aa10ae4
merge main
hnyls2002 Sep 16, 2026
3bef642
merge dsv4.1-hopper
hnyls2002 Sep 16, 2026
4cbfc19
merge dsv4.1-mhc
hnyls2002 Sep 16, 2026
d8a9654
merge main
hnyls2002 Sep 16, 2026
986f641
merge dsv4.1-mhc
hnyls2002 Sep 16, 2026
d9f57d9
merge main
hnyls2002 Sep 17, 2026
a728dc9
vision: drop unused byte loader and redundant max_seqlen; neutral doc…
hnyls2002 Sep 17, 2026
2d125d8
vision: add the DeepSeek-V4.1 image processor and its config/tokenize…
hnyls2002 Sep 17, 2026
93b1eb7
vision: use the shared RMSNorm natively; name patchify helpers and th…
hnyls2002 Sep 17, 2026
86de741
trim comments
hnyls2002 Sep 17, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions python/sglang/srt/configs/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -628,12 +628,17 @@ def __init__(
or hasattr(self.hf_config, "audio_config")
)
)
has_dsv41_vision = (
self.hf_config.model_type == "deepseek_v41"
and self.hf_config.vision_n_layers > 0
)
self.is_multimodal = (
enable_multimodal
and not self.is_lm_only
and (
is_multimodal_model(self.hf_config.architectures)
or has_multimodal_subconfig
or has_dsv41_vision
)
)
self.is_audio_model = enable_multimodal and is_audio_model(
Expand All @@ -652,6 +657,8 @@ def __init__(
self.is_multimodal
and getattr(self.hf_config, "vision_config", None) is not None
)
if self.is_multimodal and has_dsv41_vision:
self.is_image_understandable_model = True

# Models expose audio_config at different nesting levels:
# - top-level audio_config: e.g. Qwen2Audio
Expand Down
151 changes: 151 additions & 0 deletions python/sglang/srt/models/deepseek_v41_vit.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,151 @@
"""DeepSeek-V4.1 vision tower and aligner."""

from functools import lru_cache

import torch
import torch.nn.functional as F
from torch import nn

from sglang.srt.layers.attention.vision import (
VisionAttention,
VisionAttentionMetadata,
prepare_vision_attention_metadata,
)
from sglang.srt.layers.layernorm import RMSNorm


def _rms_norm(dim: int) -> RMSNorm:
# The fused CUDA kernels do not take an fp32 weight with a bf16 input.
return RMSNorm(dim, eps=1e-6, weight_dtype=torch.float32, force_native=True)


@lru_cache(8)
def get_vision_cos_sin(n_h: int, n_w: int, dim: int, theta: float):
inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
hpos = torch.arange(n_h).unsqueeze(1).expand(n_h, n_w)
wpos = torch.arange(n_w).unsqueeze(0).expand(n_h, n_w)
freqs = torch.stack([hpos, wpos], dim=-1).reshape(-1, 2, 1).float() * inv_freq
freqs = freqs.flatten(1)
return freqs.cos().unsqueeze(1), freqs.sin().unsqueeze(1)


def apply_rotary(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
dtype = x.dtype
x1, x2 = x.float().chunk(2, dim=-1)
return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1).to(dtype)


class PatchEmbed(nn.Module):
def __init__(self, args):
super().__init__()
self.proj = nn.Linear(3 * args.vision_patch_size**2, args.vision_dim)

def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.proj(x.flatten(1))


def apply_vision_rotary(q, k, position_embeddings, x_shape):
# The reference pairs the two halves of each head, with FP32 arithmetic.
cos, sin = position_embeddings
return apply_rotary(q, cos, sin), apply_rotary(k, cos, sin)


class Attention(VisionAttention):
def __init__(self, args):
super().__init__(
embed_dim=args.vision_dim,
num_heads=args.vision_n_heads,
projection_size=args.vision_dim,
use_qkv_parallel=True,
use_data_parallel=True,
customized_position_embedding_applier=apply_vision_rotary,
)

def forward(
self,
x: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
metadata: VisionAttentionMetadata,
) -> torch.Tensor:
return (
super()
.forward(
x,
position_embeddings=(cos, sin),
forward_metadata=metadata,
)
.squeeze(0)
)


class MLP(nn.Module):
def __init__(self, args):
super().__init__()
self.w1 = nn.Linear(args.vision_dim, 2 * args.vision_inter_dim, bias=False)
self.w2 = nn.Linear(args.vision_inter_dim, args.vision_dim, bias=False)

def forward(self, x: torch.Tensor) -> torch.Tensor:
gate, up = self.w1(x).chunk(2, dim=-1)
return self.w2(F.silu(gate) * up)


class Block(nn.Module):
def __init__(self, args):
super().__init__()
self.norm1 = _rms_norm(args.vision_dim)
self.attn = Attention(args)
self.norm2 = _rms_norm(args.vision_dim)
self.mlp = MLP(args)

def forward(
self,
x: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
metadata: VisionAttentionMetadata,
) -> torch.Tensor:
x = x + self.attn(self.norm1(x), cos, sin, metadata)
return x + self.mlp(self.norm2(x))


class ViT(nn.Module):
"""DeepSeek ViT: full bidirectional attention over one image with 2D RoPE."""

def __init__(self, args):
super().__init__()
self.rope_dim = args.vision_dim // args.vision_n_heads // 2
self.rope_theta = args.vision_rope_theta
self.patch_embed = PatchEmbed(args)
self.blocks = nn.ModuleList([Block(args) for _ in range(args.vision_n_layers)])
self.norm = _rms_norm(args.vision_dim)

def forward(self, patches: torch.Tensor, n_h: int, n_w: int) -> torch.Tensor:
x = self.patch_embed(patches)
cos, sin = get_vision_cos_sin(n_h, n_w, self.rope_dim, self.rope_theta)
cos, sin = cos.to(x.device), sin.to(x.device)
# Passing the known length avoids device-to-host length discovery per layer.
metadata = prepare_vision_attention_metadata(
torch.tensor([0, x.shape[0]], dtype=torch.int32),
x.device,
max_seqlen=x.shape[0],
)
for block in self.blocks:
x = block(x, cos, sin, metadata)
return self.norm(x)


class Aligner(nn.Module):
def __init__(self, args):
super().__init__()
self.downsample_ratio = args.vision_downsample_ratio
in_dim = args.vision_dim * self.downsample_ratio**2
self.w1 = nn.Linear(in_dim, args.dim)
self.w2 = nn.Linear(args.dim, args.dim)

def forward(self, x: torch.Tensor, n_h: int, n_w: int) -> torch.Tensor:
r = self.downsample_ratio
x = x.view(n_h, n_w, -1).permute(2, 0, 1)
x = F.pad(x, (0, -n_w % r, 0, -n_h % r))
x = F.unfold(x.unsqueeze(0), r, stride=r).squeeze(0).transpose(0, 1)
return self.w2(F.gelu(self.w1(x)))
200 changes: 200 additions & 0 deletions python/sglang/srt/multimodal/deepseek_v41_image_processing.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,200 @@
"""Image preprocessing.

An image becomes a `n_vit_h x n_vit_w` patch grid for the ViT and a `n_llm_h x n_llm_w` token grid
after the 3x3 aligner downsample, which the LLM sees as

[IMAGE_START] + ([IMAGE] * n_llm_w + [IMAGE_NEW_LINE]) * n_llm_h + [IMAGE_END]

Every one of those positions carries `image_token_id` in `input_ids`; only the token type tells them
apart. The IMAGE slots are filled with aligner rows in reading order.
"""

import math

import numpy as np
import torch
import torch.nn.functional as F
from PIL import Image, ImageOps

IMAGE_START, IMAGE, IMAGE_NEW_LINE, IMAGE_END = range(4)


GPU_PLAN_KEY = "dsv41_gpu_plan"


def num_image_tokens(n_llm_h: int, n_llm_w: int) -> int:
return n_llm_h * (n_llm_w + 1) + 2


def llm_grid(best_height: int, best_width: int, patch_size: int, downsample_ratio: int):
"""Token grid the aligner produces from a patch grid of this pixel size."""
return math.ceil((best_height // patch_size) / downsample_ratio), math.ceil(
(best_width // patch_size) / downsample_ratio
)


def solve_resize_ratio(height, width, patch_size, downsample_ratio, max_n_token):
"""Largest aspect-preserving pixel size whose token grid still fits in max_n_token."""
r = height / width
max_w_float = math.sqrt((max_n_token - 2) / r + 0.25) - 0.5
max_h_float = max_w_float * r
cell = patch_size * downsample_ratio
if max_w_float < 1.0: # very tall: collapse to a single column
return (max_n_token - 2) // 2 * cell, cell
if max_h_float < 1.0: # very wide: collapse to a single row
return cell, (max_n_token - 3) * cell
beta = min(
math.floor(max_w_float) * cell / width, math.floor(max_h_float) * cell / height
)
return math.floor(height * beta / patch_size) * patch_size, math.floor(
width * beta / patch_size
) * patch_size


def safe_resize(
height, width, best_height, best_width, patch_size, downsample_ratio, max_n_token
):
"""Shrink the pixel size until the image costs at most max_n_token LLM tokens."""
n_llm_h, n_llm_w = llm_grid(best_height, best_width, patch_size, downsample_ratio)
if num_image_tokens(n_llm_h, n_llm_w) > max_n_token:
best_height, best_width = solve_resize_ratio(
height, width, patch_size, downsample_ratio, max_n_token
)
n_llm_h, n_llm_w = llm_grid(
best_height, best_width, patch_size, downsample_ratio
)
assert num_image_tokens(n_llm_h, n_llm_w) <= max_n_token
return n_llm_h, n_llm_w, best_height, best_width


def plan_image_grid(width: int, height: int, args):
"""Resize plan for an image of the given original size; a pure function of its arguments."""
p = args.vision_patch_size
if (
args.vision_max_wh_ratio is not None
and width > height * args.vision_max_wh_ratio
):
width = height * args.vision_max_wh_ratio
if 0 < width * height < args.vision_min_pixels:
ratio = (args.vision_min_pixels / (width * height)) ** 0.5
width = int(width * ratio)
height = int(height * ratio)
best_width = math.ceil(width / p) * p
best_height = math.ceil(height / p) * p
return safe_resize(
height,
width,
best_height,
best_width,
p,
args.vision_downsample_ratio,
args.vision_max_n_token,
)


def to_rgb(image: Image.Image) -> Image.Image:
"""The same RGB conversion for every preprocessing backend."""
return image.convert("RGB")


def patchify_image(image, args):
p = args.vision_patch_size
image = to_rgb(image)
n_llm_h, n_llm_w, best_height, best_width = plan_image_grid(
image.width, image.height, args
)
n_vit_h, n_vit_w = best_height // p, best_width // p
if (
args.vision_max_wh_ratio is not None
and image.width >= args.vision_max_wh_ratio * image.height
):
image = image.resize((best_width, best_height))
else:
image = ImageOps.pad(image, (best_width, best_height), color=(127, 127, 127))
x = torch.from_numpy(np.asarray(image, dtype=np.float32)).permute(2, 0, 1) / 255
x = ((x - 0.5) / 0.5).to(torch.bfloat16)
patches = (
x.reshape(3, n_vit_h, p, n_vit_w, p)
.permute(1, 3, 0, 2, 4)
.reshape(n_vit_h * n_vit_w, 3, p, p)
)
return patches, n_vit_h, n_vit_w, n_llm_h, n_llm_w


def image_token_types(n_llm_h: int, n_llm_w: int) -> torch.Tensor:
types = [IMAGE_START]
types += ([IMAGE] * n_llm_w + [IMAGE_NEW_LINE]) * n_llm_h
types.append(IMAGE_END)
return torch.tensor(types, dtype=torch.int64)


def prepare_image(image, args):
image = to_rgb(image)
lh, lw, height, width = plan_image_grid(image.width, image.height, args)
stretch = (
args.vision_max_wh_ratio is not None
and image.width >= args.vision_max_wh_ratio * image.height
)
resize_h, resize_w = height, width
if not stretch:
if image.width / image.height > width / height:
resize_h = round(image.height / image.width * width)
elif image.width / image.height < width / height:
resize_w = round(image.width / image.height * height)
plan = {
"height": height,
"width": width,
"resize_h": resize_h,
"resize_w": resize_w,
"top": round((height - resize_h) / 2),
"left": round((width - resize_w) / 2),
"patch_size": args.vision_patch_size,
}
return np.array(image, dtype=np.uint8), plan, lh, lw


def patchify_image_rust(image, args, *, resize_patchify):
pixels, plan, lh, lw = prepare_image(image, args)
bits = resize_patchify(
pixels,
(plan["height"], plan["width"]),
(plan["resize_h"], plan["resize_w"]),
(plan["top"], plan["left"]),
plan["patch_size"],
)
p = plan["patch_size"]
h, w = plan["height"] // p, plan["width"] // p
patches = torch.from_numpy(bits).view(torch.bfloat16).view(h * w, 3, p, p)
return patches, h, w, lh, lw


def prepare_image_gpu(image, args):
pixels, plan, lh, lw = prepare_image(image, args)
return torch.from_numpy(pixels).permute(2, 0, 1).contiguous(), plan, lh, lw


def materialize_image_gpu(pixels: torch.Tensor, plan: dict) -> torch.Tensor:
"""Resize, pad, normalize and patchify on the input tensor's device."""
x = pixels.unsqueeze(0).float()
target = (plan["resize_h"], plan["resize_w"])
# PIL resizes separably, rounding and clamping to uint8 after each pass;
# fusing the two passes into one float resize diverges on high-contrast images.
for size in ((x.shape[-2], target[1]), target):
if x.shape[-2:] != size:
x = (
F.interpolate(
x, size=size, mode="bicubic", align_corners=False, antialias=True
)
.round()
.clamp_(0, 255)
)
top, left = plan["top"], plan["left"]
x = F.pad(
x,
(left, plan["width"] - target[1] - left, top, plan["height"] - target[0] - top),
value=127,
)
x = ((x / 255 - 0.5) / 0.5).to(torch.bfloat16)
p = plan["patch_size"]
h, w = plan["height"] // p, plan["width"] // p
return x.reshape(3, h, p, w, p).permute(1, 3, 0, 2, 4).reshape(h * w, 3, p, p)
Loading
Loading