Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
13 changes: 7 additions & 6 deletions python/sglang/multimodal_gen/configs/sample/minimax_h3.py
Original file line number Diff line number Diff line change
Expand Up @@ -241,14 +241,15 @@ def _validate(self) -> None:
"video/audio denoise loop has no lossless TeaCache contract"
)
if self.rollout:
raise ValueError(
"MiniMax H3 does not support rollout: its coupled video/audio "
"scheduler has no SchedulerRLMixin contract"
)
task = str(self.task or "t2va").lower()
if task not in ("t2va",):
raise ValueError(
f"MiniMax H3 rollout currently supports task=t2va only, got {task!r}"
)
if self.return_trajectory_latents or self.return_trajectory_decoded:
raise ValueError(
"MiniMax H3 does not support trajectory output for its coupled "
"video/audio denoise state"
"MiniMax H3 does not support return_trajectory_latents/decoded; "
"use rollout=True with rollout_return_dit_trajectory instead"
)
seeds = self.seed if isinstance(self.seed, list) else [self.seed]
for seed in seeds:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,8 @@ class RolloutRequest(BaseModel):
fps: Optional[int] = None

rollout: bool = True
# "uint8": quantise the video engine-side; None: ship unchanged
rollout_video_dtype: Optional[str] = None
rollout_sde_type: str = "sde"
rollout_noise_level: float = 0.7
rollout_log_prob_no_const: bool = False
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,83 @@
"""Wall-clock marks for the rollout HTTP path.

The client reconstructs one waterfall per request across three processes, so it
needs absolute marks, not durations. They ride a response header: every mark is
taken before the response headers are sent, and a header leaves the msgpack
body contract untouched for clients that ignore it.
"""

from __future__ import annotations

import json
import time
from contextlib import contextmanager

# Compact JSON rides under each header, e.g.
# x-sgld-timing: {"srv_recv":1787472609.104512,...,"msgpack_end":1787472691.63172,"request_id":"abc123"}
# mark name -> epoch seconds (6 dp), plus this request's id
# x-sgld-stages: {"TextEncodingStage":91.2,"DenoisingStage":448512.301,"DecodingStage":8123.457}
# pipeline stage class -> milliseconds (3 dp)
TIMING_HEADER = "x-sgld-timing"
STAGES_HEADER = "x-sgld-stages"

MARKS = (
"srv_recv",
"forward_start",
"forward_end",
"build_start",
"build_end",
"dump_end",
"msgpack_end",
)


class RequestStamps:
"""Absolute wall-clock marks for one rollout request."""

__slots__ = ("request_id", "_marks")

def __init__(self, request_id: str = "") -> None:
self.request_id = request_id
self._marks: dict[str, float] = {}

def mark(self, name: str) -> None:
assert name in MARKS, f"unknown timing mark {name!r}"
self._marks[name] = time.time()

@contextmanager
def span(self, name: str):
"""Mark ``{name}_start`` on entry and ``{name}_end`` on exit.

dump/msgpack have no start mark: each begins where the previous span
ended, so only the end boundary is recorded.
"""
start = f"{name}_start"
if start in MARKS:
self.mark(start)
try:
yield
finally:
self.mark(f"{name}_end")

def to_header(self) -> str:
payload: dict[str, object] = {
name: round(t, 6) for name, t in self._marks.items()
}
if self.request_id:
payload["request_id"] = self.request_id
return json.dumps(payload, separators=(",", ":"))


def stages_header(metrics) -> str:
"""The per-stage milliseconds the engine already recorded, verbatim.

The client sees the whole forward as one number, so without these it cannot
tell a slow denoise from a slow VAE decode. Only stage totals travel, never
the per-step list, which would grow the header with the step count.
"""
if metrics is None or not metrics.stages:
return ""
return json.dumps(
{name: round(ms, 3) for name, ms in metrics.stages.items()},
separators=(",", ":"),
)
Original file line number Diff line number Diff line change
Expand Up @@ -2,20 +2,29 @@

from __future__ import annotations

import asyncio
from typing import Any

import msgspec
import torch
from fastapi import APIRouter, HTTPException, Response
from fastapi.responses import StreamingResponse

from sglang.multimodal_gen.configs.sample.sampling_params import generate_request_id
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import build_sampling_params
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
RolloutRequest,
RolloutResponse,
)
from sglang.multimodal_gen.runtime.entrypoints.post_training.request_timing import (
STAGES_HEADER,
TIMING_HEADER,
RequestStamps,
stages_header,
)
from sglang.multimodal_gen.runtime.entrypoints.post_training.utils import (
_maybe_serialize,
_quantize_video_uint8,
msgpack_encode_spliced,
)
from sglang.multimodal_gen.runtime.entrypoints.utils import prepare_request
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
Expand Down Expand Up @@ -131,6 +140,7 @@ def _slice_rollout_trajectory_for_sample(
latents=_extract_single_sample_tensor(dit.latents, sample_idx, batch_size),
timesteps=dit.timesteps,
sigmas=dit.sigmas,
latent_step_indices=dit.latent_step_indices,
)
return RolloutTrajectoryData(
rollout_log_probs=log_probs,
Expand Down Expand Up @@ -185,6 +195,11 @@ def _serialize_rollout_trajectory(
),
"timesteps": serialized_dit_timesteps,
"sigmas": serialized_dit_sigmas,
"latent_step_indices": (
_maybe_serialize(dit.latent_step_indices)
if dit.latent_step_indices is not None
else None
),
}
return (
serialized_log_probs,
Expand All @@ -195,7 +210,12 @@ def _serialize_rollout_trajectory(


def _build_response(
request_id: str, prompt: str, seed: int, rollout: bool, result: OutputBatch
request_id: str,
prompt: str,
seed: int,
rollout: bool,
result: OutputBatch,
video_dtype: str | None = None,
) -> list[RolloutResponse]:
"""
rollout: bool - set to False when evaluating the model
Expand Down Expand Up @@ -227,6 +247,8 @@ def _build_response(
for sample_idx in range(batch_size):
out_i = result.output[sample_idx]
if isinstance(out_i, torch.Tensor):
if video_dtype == "uint8":
out_i = _quantize_video_uint8(out_i)
out_i = out_i.contiguous()
serialized_generated_output = _maybe_serialize(out_i)
if not rollout:
Expand Down Expand Up @@ -317,7 +339,10 @@ def _build_sampling_kwargs(request: RolloutRequest) -> dict:
},
)
async def rollout_generate(request: RolloutRequest):
stamps = RequestStamps()
stamps.mark("srv_recv")
request_id = generate_request_id()
stamps.request_id = request_id
server_args = get_global_server_args()
sampling_kwargs = _build_sampling_kwargs(request)
try:
Expand All @@ -330,21 +355,47 @@ async def rollout_generate(request: RolloutRequest):
server_args=server_args, sampling_params=sampling_params
)
try:
output_batch: OutputBatch = await async_scheduler_client.forward(
pipeline_request
)
with stamps.span("forward"):
output_batch: OutputBatch = await async_scheduler_client.forward(
pipeline_request
)
except Exception as exc:
logger.error("Rollout generation failed: %s", exc, exc_info=True)
raise HTTPException(
status_code=500, detail=f"Generation failed: {exc}"
) from exc
if output_batch.error:
raise HTTPException(status_code=500, detail=output_batch.error)
rollout_responses = _build_response(
request_id, request.prompt, request.seed, request.rollout, output_batch
)
payload = [r.model_dump() for r in rollout_responses]
return Response(
content=msgspec.msgpack.encode(payload),

def _serialize_response() -> list[bytes]:
with stamps.span("build"):
rollout_responses = _build_response(
request_id,
request.prompt,
request.seed,
request.rollout,
output_batch,
video_dtype=request.rollout_video_dtype,
)
with stamps.span("dump"):
payload = [r.model_dump() for r in rollout_responses]
with stamps.span("msgpack"):
return msgpack_encode_spliced(payload)

# Off the event loop: building the response copies hundreds of MB
parts = await asyncio.to_thread(_serialize_response)

# Timing rides a header: the last marks postdate the encoded body. The
# explicit content-length keeps uvicorn on identity framing, not chunked.
headers = {
TIMING_HEADER: stamps.to_header(),
"content-length": str(sum(len(p) for p in parts)),
}
stages = stages_header(output_batch.metrics)
if stages:
headers[STAGES_HEADER] = stages
return StreamingResponse(
iter(parts),
media_type="application/msgpack",
headers=headers,
)
Original file line number Diff line number Diff line change
Expand Up @@ -4,10 +4,13 @@

from typing import Any

import msgspec
import numpy as np
import torch
from safetensors.torch import load, save

_SPLICE_MIN_BYTES = 1 << 20


def tensor_to_bytes(t: torch.Tensor) -> bytes:
return save({"t": t.detach().contiguous().cpu()})
Expand All @@ -34,6 +37,48 @@ def _maybe_serialize(obj: Any) -> Any:
return obj


def _container_header(tag_fix: int, tag16: bytes, tag32: bytes, n: int) -> bytes:
if n <= 15:
return bytes([tag_fix | n])
if n <= 0xFFFF:
return tag16 + n.to_bytes(2, "big")
return tag32 + n.to_bytes(4, "big")


def msgpack_encode_spliced(obj: Any, threshold: int = _SPLICE_MIN_BYTES) -> list[bytes]:
"""Encode to msgpack as a list of buffers, splicing large ``bytes`` values.

Container and bin32 headers are hand-written so a large ``bytes`` payload
lands in the output by reference, never copied through the encoder; the
concatenated parts are byte-identical to ``msgspec.msgpack.encode(obj)``.
"""
parts: list[bytes] = []
small = bytearray()

def _emit(value: Any) -> None:
if isinstance(value, bytes) and len(value) >= threshold:
small.extend(b"\xc6" + len(value).to_bytes(4, "big"))
parts.append(bytes(small))
small.clear()
parts.append(value)
elif isinstance(value, dict):
small.extend(_container_header(0x80, b"\xde", b"\xdf", len(value)))
for key, item in value.items():
small.extend(msgspec.msgpack.encode(key))
_emit(item)
elif isinstance(value, (list, tuple)):
small.extend(_container_header(0x90, b"\xdc", b"\xdd", len(value)))
for item in value:
_emit(item)
else:
small.extend(msgspec.msgpack.encode(value))

_emit(obj)
if small:
parts.append(bytes(small))
return parts


def _maybe_deserialize(obj: Any) -> Any:
if isinstance(obj, dict):
if obj.get("__tensor__"):
Expand All @@ -42,3 +87,11 @@ def _maybe_deserialize(obj: Any) -> Any:
if isinstance(obj, (list, tuple)):
return [_maybe_deserialize(v) for v in obj]
return obj


def _quantize_video_uint8(video: torch.Tensor) -> torch.Tensor:
"""Map the decoded [0,1] float video to 0..255 uint8; consumers divide by 255."""
out = video.float()
if float(out.max()) <= 1.0 + 1e-3:
out = out * 255.0
return out.clamp(0, 255).to(torch.uint8)
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@
)
from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
LinearBase,
MergedColumnParallelLinear,
RowParallelLinear,
)
Expand Down Expand Up @@ -370,6 +371,8 @@ def _rotate_half(x: torch.Tensor) -> torch.Tensor:


def _accepts_mxfp8_input(linear: nn.Module) -> bool:
if not isinstance(linear, LinearBase):
return False
return linear.quant_method is not None and linear.quant_method.accepts_mxfp8_input(
linear
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -919,6 +919,7 @@ def _get_added_qkv_projections(
if (
self._unquantized_added_qkv_is_packed
and not qwen_image_added_qkv_active(self)
and isinstance(self.to_added_qkv, MergedColumnParallelLinear)
):
return _split_unquantized_merged_linear(
self.to_added_qkv, encoder_hidden_states
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -53,12 +53,10 @@
logger = init_logger(__name__)


def _swap_peft_swiglu_fc1_lora_b(
def swap_peft_swiglu_fc1_lora_b(
source_name: str, target_name: str, weight: torch.Tensor
) -> torch.Tensor:
# Only the PEFT -> native H3 FFN rewrite: ff.net.0.proj [value; gate]
# onto mlp.fc1 [gate; value]. Native fused mlp.fc1 and other models'
# ff.net.0.proj (e.g. Flux) must not match.
"""Rewrite PEFT H3 FFN lora_B from [value; gate] to native [gate; value]."""
if (
weight.dim() != 2
or ".ff.net.0.proj.lora_B" not in source_name
Expand Down Expand Up @@ -934,7 +932,7 @@ def load_lora_adapter(
else:
continue

weight = _swap_peft_swiglu_fc1_lora_b(name, target_name, weight)
weight = swap_peft_swiglu_fc1_lora_b(name, target_name, weight)
if target_name in self.lora_adapters[lora_nickname]:
raise ValueError(
f"Dit target weight name {target_name} already exists in lora_adapters[{lora_nickname}]"
Expand Down
Loading
Loading