diff --git a/python/sglang/multimodal_gen/configs/sample/minimax_h3.py b/python/sglang/multimodal_gen/configs/sample/minimax_h3.py index c78216d5b11b..d5d79aba56a0 100644 --- a/python/sglang/multimodal_gen/configs/sample/minimax_h3.py +++ b/python/sglang/multimodal_gen/configs/sample/minimax_h3.py @@ -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: diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py index b558be3e4a21..de01bc49283b 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/request_timing.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/request_timing.py new file mode 100644 index 000000000000..6404a1c75076 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/request_timing.py @@ -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=(",", ":"), + ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py index 71492523dcd7..9c3cba5c0d29 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/rollout_api.py @@ -2,11 +2,12 @@ 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 @@ -14,8 +15,16 @@ 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 @@ -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, @@ -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, @@ -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 @@ -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: @@ -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: @@ -330,9 +355,10 @@ 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( @@ -340,11 +366,36 @@ async def rollout_generate(request: RolloutRequest): ) 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, ) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/utils.py index 495fadf88d0b..8fa5c5f3e6c4 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/utils.py @@ -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()}) @@ -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__"): @@ -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) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py index db1d29a2d5d2..31b3f73bc40a 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/minimax_h3.py @@ -59,6 +59,7 @@ ) from sglang.multimodal_gen.runtime.layers.linear import ( ColumnParallelLinear, + LinearBase, MergedColumnParallelLinear, RowParallelLinear, ) @@ -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 ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py index cea45be71c0c..9d7357dd017a 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -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 diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py index 5f6620b8b3f0..f009e10d1ba0 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/lora/pipeline.py @@ -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 @@ -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}]" diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/denoise_loop.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/denoise_loop.py index 77f389b08ffe..62599da199bc 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/denoise_loop.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/denoise_loop.py @@ -462,6 +462,7 @@ def minimax_h3_denoise_loop( attn_metadata: AttentionMetadata | None = None, on_step: Callable[[int, torch.Tensor, torch.Tensor], None] | None = None, step_profiler: Callable[[int], AbstractContextManager] | None = None, + rollout_ctx=None, ) -> tuple[torch.Tensor, torch.Tensor]: """Run the full denoise loop; returns final (video_rows, audio_rows). @@ -557,11 +558,22 @@ def minimax_h3_denoise_loop( audio_one_minus_sigma_ratios = 1.0 - audio_sigma_ratios video_denoised_scratch = torch.empty_like(video_rows[video_target_slice]) audio_denoised_scratch = torch.empty_like(audio_rows[audio_target_slice]) + if rollout_ctx is not None: + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.minimax_h3_rollout import ( + minimax_h3_rollout_update_video_target, + ) + + rollout_ctx.batch._h3_rollout_sigma_max = float(max(sigmas_video)) + video_target = video_rows[video_target_slice] + rollout_ctx.collector.record_initial(video_target.detach().clone()) for step in range(num_steps): step_cm = step_profiler(step) if step_profiler is not None else nullcontext() with step_cm: s_v = sigmas_video[step] s_a = sigmas_audio[step] + s_n_v = sigmas_video[step + 1] + if rollout_ctx is not None: + rollout_ctx.batch._rollout_loop_step_index = step fk = positive.forward_kwargs( video_rows=video_rows, @@ -592,15 +604,38 @@ def minimax_h3_denoise_loop( mv_audio_t = v_audio[audio_target_slice].float() video_target = video_rows[video_target_slice] - _minimax_h3_update_target_rows_( - video_target, - mv_video_t, - sigma_t=video_sigma_t[step], - sigma_curr=s_v, - sigma_ratio=video_sigma_ratios[step], - one_minus_sigma_ratio=video_one_minus_sigma_ratios[step], - denoised_scratch=video_denoised_scratch, - ) + if rollout_ctx is not None: + ( + updated, + log_sum, + log_count, + rollout_ctx.noise_buffer, + ) = minimax_h3_rollout_update_video_target( + video_target, + mv_video_t, + sigma_curr=s_v, + sigma_next=s_n_v, + batch=rollout_ctx.batch, + generator=rollout_ctx.generator, + loop_step_index=step, + noise_buffer=rollout_ctx.noise_buffer, + ) + video_target.copy_(updated) + rollout_ctx.collector.record_step( + video_target.detach().clone(), + log_sum, + log_count, + ) + else: + _minimax_h3_update_target_rows_( + video_target, + mv_video_t, + sigma_t=video_sigma_t[step], + sigma_curr=s_v, + sigma_ratio=video_sigma_ratios[step], + one_minus_sigma_ratio=video_one_minus_sigma_ratios[step], + denoised_scratch=video_denoised_scratch, + ) audio_target = audio_rows[audio_target_slice] _minimax_h3_update_target_rows_( diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/minimax_h3_rollout.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/minimax_h3_rollout.py new file mode 100644 index 000000000000..2bccab0e0ca4 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/minimax_h3_rollout.py @@ -0,0 +1,242 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Rollout (RL) helpers for MiniMax H3 video-only denoising.""" + +from __future__ import annotations + +import math +from dataclasses import dataclass, field +from typing import Any + +import torch + +from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import ( + RolloutDenoisingEnv, + RolloutDitTrajectory, + RolloutTrajectoryData, +) + +_LOG_SQRT_2PI = math.log(math.sqrt(2 * math.pi)) + + +def _effective_sde_type(batch, loop_step_index: int) -> str: + sde_type = str(getattr(batch, "rollout_sde_type", "sde")) + if sde_type == "ode": + return "ode" + sde_step_indices = getattr(batch, "rollout_sde_step_indices", None) + if sde_step_indices is not None and loop_step_index not in sde_step_indices: + return "ode" + return sde_type + + +def minimax_h3_rollout_update_video_target( + video_target: torch.Tensor, + velocity: torch.Tensor, + *, + sigma_curr: float, + sigma_next: float, + batch, + generator: torch.Generator, + loop_step_index: int, + noise_buffer: torch.Tensor | None, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: + """Stochastic or deterministic update for video target rows during rollout. + + H3 velocity ``v`` relates to flow-matching model output as ``model_output = -v``. + Returns ``(updated_target, log_prob_sum[B=1], log_prob_count[B=1], noise_buffer)``. + """ + sde_type = _effective_sde_type(batch, loop_step_index) + noise_level = float(getattr(batch, "rollout_noise_level", 0.0)) + log_prob_no_const = bool(getattr(batch, "rollout_log_prob_no_const", False)) + + if sde_type == "ode" or noise_level == 0.0: + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.denoise_loop import ( + _minimax_h3_update_target_rows_, + ) + + out = video_target.clone() + sigma_t = out.new_tensor(sigma_curr) + ratio = out.new_tensor(sigma_next / sigma_curr if sigma_curr != 0.0 else 1.0) + one_minus_ratio = 1.0 - ratio + scratch = torch.empty_like(out) + _minimax_h3_update_target_rows_( + out, + velocity.float(), + sigma_t=sigma_t, + sigma_curr=sigma_curr, + sigma_ratio=ratio, + one_minus_sigma_ratio=one_minus_ratio, + denoised_scratch=scratch, + ) + elem_count = float(out.numel()) + log_prob_sum = torch.zeros(1, device=out.device, dtype=torch.float32) + log_prob_count = out.new_tensor([elem_count]) + return out, log_prob_sum, log_prob_count, noise_buffer + + model_output = (-velocity).float() + sample = video_target.float() + current_sigma = sample.new_tensor(float(sigma_curr)) + next_sigma = sample.new_tensor(float(sigma_next)) + + if ( + noise_buffer is None + or noise_buffer.shape != sample.shape + or noise_buffer.device != sample.device + ): + noise_buffer = torch.empty_like(sample) + variance_noise = torch.randn( + sample.shape, + generator=generator, + device=sample.device, + dtype=torch.float32, + ) + noise_buffer.copy_(variance_noise) + + if sde_type == "cps": + std_dev_t = next_sigma * math.sin(noise_level * math.pi / 2) + noise_std_dev = std_dev_t + pred_original = sample - current_sigma * model_output + noise_estimate = sample + model_output * (1.0 - current_sigma) + prev_mean = pred_original * (1.0 - next_sigma) + noise_estimate * torch.sqrt( + torch.clamp(next_sigma**2 - std_dev_t**2, min=1e-12) + ) + prev_sample = prev_mean + variance_noise * noise_std_dev + log_prob_no_const_val = -((variance_noise * noise_std_dev) ** 2) + elif sde_type == "sde": + dt = next_sigma - current_sigma + sigma_max = float(getattr(batch, "_h3_rollout_sigma_max", 1.0)) + std_dev_t = ( + torch.sqrt( + current_sigma + / ( + 1.0 + - torch.where( + torch.isclose(current_sigma, current_sigma.new_tensor(1.0)), + current_sigma.new_tensor(sigma_max), + current_sigma, + ) + ) + ) + * noise_level + ) + noise_std_dev = std_dev_t * torch.sqrt(-1.0 * dt) + prev_mean = ( + sample * (1.0 + std_dev_t**2 / (2.0 * current_sigma) * dt) + + model_output + * (1.0 + std_dev_t**2 * (1.0 - current_sigma) / (2.0 * current_sigma)) + * dt + ) + prev_sample = prev_mean + variance_noise * noise_std_dev + log_prob_no_const_val = -((variance_noise * noise_std_dev) ** 2) + else: + raise ValueError(f"Unsupported rollout_sde_type for H3: {sde_type!r}") + + if log_prob_no_const or sde_type == "ode": + log_prob_sum = log_prob_no_const_val.sum().unsqueeze(0) + else: + log_prob_sum = ( + ( + log_prob_no_const_val / (2.0 * (noise_std_dev**2)) + - torch.log(noise_std_dev) + - _LOG_SQRT_2PI + ) + .sum() + .unsqueeze(0) + ) + log_prob_count = sample.new_tensor([float(sample.numel())]) + return ( + prev_sample.to(dtype=video_target.dtype), + log_prob_sum, + log_prob_count, + noise_buffer, + ) + + +@dataclass +class MiniMaxH3RolloutCtx: + """Per-denoise-loop rollout state for H3.""" + + batch: Any + generator: torch.Generator + sigmas_video: list[float] + collector: MiniMaxH3RolloutCollector + noise_buffer: torch.Tensor | None = None + denoising_env_kwargs: dict[str, Any] = field(default_factory=dict) + + +@dataclass +class MiniMaxH3RolloutCollector: + """Accumulates video-target trajectory and per-step log probs.""" + + sigmas_video: list[float] + return_step_indices: set[int] | None = None + latent_steps: list[torch.Tensor] = field(default_factory=list) + latent_step_indices: list[int] = field(default_factory=list) + log_prob_sums: list[torch.Tensor] = field(default_factory=list) + log_prob_counts: list[torch.Tensor] = field(default_factory=list) + pos_cond_kwargs: dict[str, Any] = field(default_factory=dict) + _next_latent_index: int = 0 + + def _record_latent(self, video_target: torch.Tensor) -> None: + index = self._next_latent_index + self._next_latent_index += 1 + if ( + self.return_step_indices is not None + and index not in self.return_step_indices + ): + return + self.latent_steps.append(video_target.detach().cpu().clone()) + self.latent_step_indices.append(index) + + def record_initial(self, video_target: torch.Tensor) -> None: + self._record_latent(video_target) + + def record_step( + self, + video_target: torch.Tensor, + log_prob_sum: torch.Tensor, + log_prob_count: torch.Tensor, + ) -> None: + self._record_latent(video_target) + self.log_prob_sums.append(log_prob_sum.detach().cpu()) + self.log_prob_counts.append(log_prob_count.detach().cpu()) + + def build_trajectory_data(self) -> RolloutTrajectoryData: + # latents: [B=1, T+1, num_video_target_rows, width] + stacked = torch.stack(self.latent_steps, dim=0).unsqueeze(0) + divisor = 1000.0 + step_sigmas = torch.tensor( + [float(s) for s in self.sigmas_video], + dtype=torch.float32, + ) + timesteps = step_sigmas * divisor + sigmas = torch.tensor( + [float(s) for s in self.sigmas_video], + dtype=torch.float32, + ) + log_probs = None + if self.log_prob_sums: + sums = torch.stack(self.log_prob_sums, dim=0) + counts = torch.stack(self.log_prob_counts, dim=0) + per_step = sums.squeeze(-1) / counts.squeeze(-1).clamp(min=1.0) + log_probs = per_step.unsqueeze(0) + return RolloutTrajectoryData( + rollout_log_probs=log_probs, + denoising_env=RolloutDenoisingEnv( + pos_cond_kwargs=self.pos_cond_kwargs, + ), + dit_trajectory=RolloutDitTrajectory( + latents=stacked, + timesteps=timesteps, + sigmas=sigmas, + latent_step_indices=torch.tensor( + self.latent_step_indices, dtype=torch.long + ), + ), + ) + + +__all__ = [ + "MiniMaxH3RolloutCollector", + "MiniMaxH3RolloutCtx", + "minimax_h3_rollout_update_video_target", +] diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py index ab65f65c54fa..71998ec3b0c7 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/stages/denoising.py @@ -780,6 +780,54 @@ def _run_full_loop(self, batch: Req, server_args: ServerArgs) -> None: device=device, ) initial_video, initial_audio = _expand_initial_rows(ctx, positive) + rollout_ctx = None + if getattr(batch, "rollout", False): + from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.minimax_h3.minimax_h3_rollout import ( + MiniMaxH3RolloutCollector, + MiniMaxH3RolloutCtx, + ) + from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import ( + RolloutTrajectoryData, + ) + + task = str(getattr(batch.sampling_params, "task", "") or "t2va").lower() + if task not in ("t2va",): + raise ValueError( + f"MiniMax H3 rollout currently supports task=t2va only, got {task!r}" + ) + generator = torch.Generator(device=device) + seed = getattr(batch.sampling_params, "seed", 0) + if isinstance(seed, list): + seed = seed[0] + generator.manual_seed(int(seed)) + return_step_indices = getattr( + batch, "rollout_return_step_indices", None + ) + collector = MiniMaxH3RolloutCollector( + sigmas_video=sigmas_video, + return_step_indices=( + set(return_step_indices) + if return_step_indices is not None + else None + ), + ) + packed_cpu = { + k: (v.detach().cpu() if isinstance(v, torch.Tensor) else v) + for k, v in packed.items() + } + collector.pos_cond_kwargs = { + "encoder_hidden_states": emb["hidden_states"].detach().cpu(), + "h3_packed_layout": packed_cpu, + "h3_token_tags": tags.detach().cpu(), + "h3_video_target_start": positive.video_target_start, + } + rollout_ctx = MiniMaxH3RolloutCtx( + batch=batch, + generator=generator, + sigmas_video=sigmas_video, + collector=collector, + ) + batch.rollout_trajectory_data = RolloutTrajectoryData() with ( maybe_nvtx_range("denoising_loop", self.current_use_nvtx), self.progress_bar( @@ -818,7 +866,12 @@ def on_step(_step, _video_rows, _audio_rows): self._profile_denoising_step, batch=batch, ), + rollout_ctx=rollout_ctx, ) + if rollout_ctx is not None: + batch.rollout_trajectory_data = ( + rollout_ctx.collector.build_trajectory_data() + ) finally: self._finish_active_component_use() _publish_full_loop_outputs( diff --git a/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py b/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py index 1882776d573d..c48a4fcb3bdd 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py +++ b/python/sglang/multimodal_gen/runtime/post_training/rl_dataclasses.py @@ -51,9 +51,12 @@ class RolloutDitTrajectory: # [B, T+1, ...]: per-step noisy latents x_{t_0..t_{T-1}} followed by the # final denoised latent x_{t_T} (last scheduler.step output). latents: torch.Tensor | None = None - timesteps: torch.Tensor | None = None # [T] + timesteps: torch.Tensor | None = None # [T+1], includes the terminal timestep # [T+1] scheduler.sigmas snapshot (post-shift, includes terminal 0). sigmas: torch.Tensor | None = None + # [K] original step index of each kept latent (0..T); None means the full + # 0..T trajectory. Set whenever rollout_return_step_indices filtered it. + latent_step_indices: torch.Tensor | None = None @dataclass diff --git a/python/sglang/multimodal_gen/runtime/post_training/rollout_denoising_mixin.py b/python/sglang/multimodal_gen/runtime/post_training/rollout_denoising_mixin.py index 674fb1a87933..fde2d199a5e8 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/rollout_denoising_mixin.py +++ b/python/sglang/multimodal_gen/runtime/post_training/rollout_denoising_mixin.py @@ -138,6 +138,7 @@ def _maybe_init_denoising_env_collection( batch._rollout_denoising_env_state = { "env": env, "step_latents": [], + "step_latent_indices": [], "step_timesteps": [], "pos_cond_kwargs_src": pos_src, "neg_cond_kwargs_src": neg_src, @@ -156,12 +157,13 @@ def _maybe_append_dit_trajectory_step( if state is None: return + state["step_timesteps"].append(timestep_value.detach().cpu()) return_step_indices = getattr(batch, "rollout_return_step_indices", None) if return_step_indices is not None and step_index not in return_step_indices: return state["step_latents"].append(latents.detach()) - state["step_timesteps"].append(timestep_value.detach().cpu()) + state["step_latent_indices"].append(step_index) def _maybe_finalize_denoising_env_collection(self, batch, pipeline_config) -> None: state = getattr(batch, "_rollout_denoising_env_state", None) @@ -186,6 +188,9 @@ def _maybe_finalize_denoising_env_collection(self, batch, pipeline_config) -> No latents=step_latents_tensor.cpu(), timesteps=torch.stack(step_timesteps, dim=0).cpu(), sigmas=batch.scheduler.sigmas.detach().cpu().clone(), + latent_step_indices=torch.tensor( + state["step_latent_indices"], dtype=torch.long + ), ) if env is not None and batch.rollout_return_denoising_env: diff --git a/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py b/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py index 75359a82cefe..bf92ab67d73a 100644 --- a/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py +++ b/python/sglang/multimodal_gen/runtime/post_training/weights_updater.py @@ -65,6 +65,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.lora.pipeline import ( LoRAPipeline, stack_or_compose_fused_lora, + swap_peft_swiglu_fc1_lora_b, ) from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import maybe_download_model from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger @@ -263,12 +264,18 @@ def _resolve_lora_ipc_layer_dict_key( if map_name is None: return None, layer_prefix, None - mapped_name, merge_index = map_name(f"{layer_prefix}.weight") - mapped = _strip_param_weight_suffix(mapped_name) - if mapped != layer_prefix: - layer = layer_dict.get(mapped) - if layer is not None: - return layer, mapped, merge_index + # miles-h3 disk mapping is LoRA-suffix-only; main also maps ".weight". + for suffix in (".weight", ".lora_A"): + mapped_name, merge_index = map_name(f"{layer_prefix}{suffix}") + mapped = ( + mapped_name[: -len(suffix)] + if mapped_name.endswith(suffix) + else _strip_param_weight_suffix(mapped_name) + ) + if mapped != layer_prefix: + layer = layer_dict.get(mapped) + if layer is not None: + return layer, mapped, merge_index return None, layer_prefix, None @@ -707,7 +714,7 @@ def _update_lora_from_tensor( plain_pairs: list[tuple[torch.Tensor, torch.Tensor, Any]] = [] fused_sections: dict[Any, dict[int, tuple[torch.Tensor, torch.Tensor]]] = {} for layer_name, (lora_a, lora_b) in pairs.items(): - layer, _resolved_key, merge_index = _resolve_lora_ipc_layer_dict_key( + layer, resolved_key, merge_index = _resolve_lora_ipc_layer_dict_key( layer_name, layer_dict, dit_module ) if layer is None: @@ -719,6 +726,15 @@ def _update_lora_from_tensor( unknown_layers.append(layer_name) skipped += 1 continue + # Same PEFT → native rewrite as disk load_lora_adapter: fused QKV + # uses merge_index; H3 gated FFN swaps [up, gate] → [gate, up]. + # Native names already in lora_layers skip both transforms. + lora_a = swap_peft_swiglu_fc1_lora_b( + f"{layer_name}.lora_A", f"{resolved_key}.lora_A", lora_a + ) + lora_b = swap_peft_swiglu_fc1_lora_b( + f"{layer_name}.lora_B", f"{resolved_key}.lora_B", lora_b + ) inferred_rank = int(lora_a.shape[-2]) if lora_rank is not None and lora_rank != inferred_rank: logger.warning( diff --git a/python/sglang/multimodal_gen/test/unit/test_fused_lora_compose.py b/python/sglang/multimodal_gen/test/unit/test_fused_lora_compose.py index cdb221e31693..7dc45e100701 100644 --- a/python/sglang/multimodal_gen/test/unit/test_fused_lora_compose.py +++ b/python/sglang/multimodal_gen/test/unit/test_fused_lora_compose.py @@ -8,6 +8,7 @@ import torch import torch.nn.functional as F +from sglang.multimodal_gen.configs.models.dits.minimax_h3 import MiniMaxH3DiTArchConfig from sglang.multimodal_gen.runtime.layers.lora.linear import ( MergedColumnParallelLinearWithLoRA, ) @@ -15,6 +16,7 @@ LoRAPipeline, _store_fused_lora_groups, stack_or_compose_fused_lora, + swap_peft_swiglu_fc1_lora_b, ) from sglang.multimodal_gen.runtime.post_training.weights_updater import ( _resolve_lora_ipc_layer_dict_key, @@ -397,3 +399,44 @@ def test_apply_composed_adapter_end_to_end(): + _reference_delta(x, a_list, b_list, adapter_alpha) * strength ) torch.testing.assert_close(out, expected, rtol=1e-5, atol=1e-5) + + +def test_h3_ipc_reuses_disk_mapping_and_ffn_swap(): + module = torch.nn.Module() + module.param_names_mapping = MiniMaxH3DiTArchConfig().param_names_mapping + qkv, fc1 = object(), object() + layer_dict = { + "blocks.0.attn.qkv_proj": qkv, + "blocks.0.mlp.fc1": fc1, + "token_refiner.blocks.1.attn.qkv_proj": qkv, + } + cases = [ + ("transformer_blocks.0.attn.to_k", qkv, "blocks.0.attn.qkv_proj", 1), + ("transformer_blocks.0.ff.net.0.proj", fc1, "blocks.0.mlp.fc1", None), + ( + "token_refiner.refiner_blocks.1.attn.to_q", + qkv, + "token_refiner.blocks.1.attn.qkv_proj", + 0, + ), + ] + for src, layer, key, merge_index in cases: + assert _resolve_lora_ipc_layer_dict_key(src, layer_dict, module) == ( + layer, + key, + merge_index, + ) + + peft_b = torch.arange(8, dtype=torch.float32).reshape(4, 2) + swapped = swap_peft_swiglu_fc1_lora_b( + "transformer_blocks.0.ff.net.0.proj.lora_B", + "blocks.0.mlp.fc1.lora_B", + peft_b, + ) + torch.testing.assert_close(swapped, torch.cat([peft_b[2:], peft_b[:2]])) + assert ( + swap_peft_swiglu_fc1_lora_b( + "blocks.0.mlp.fc1.lora_B", "blocks.0.mlp.fc1.lora_B", swapped + ) + is swapped + ) diff --git a/python/sglang/multimodal_gen/test/unit/test_msgpack_splice.py b/python/sglang/multimodal_gen/test/unit/test_msgpack_splice.py new file mode 100644 index 000000000000..faa6dd0487f4 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_msgpack_splice.py @@ -0,0 +1,67 @@ +"""Unit tests for msgpack_encode_spliced. + +Mental model (what we exercise): + + payload = {meta..., "data": , nested: [...]} + | + msgpack_encode_spliced(payload) + | + [small-parts bytes][big bytes BY REFERENCE][small-parts bytes] + | + b"".join(parts) == msgspec.msgpack.encode(payload) (byte-identical) + +Covered: +1. Joined parts are byte-identical to a plain msgspec encode, across nested + dicts/lists, >15-element containers (map16/array16 headers), and all + scalar leaves the rollout payload uses. +2. Large bytes land in the parts list by reference (zero copy), small bytes + stay inline. +""" + +import unittest + +import msgspec + +from sglang.multimodal_gen.runtime.entrypoints.post_training.utils import ( + msgpack_encode_spliced, +) + + +def _sample_payload(big: bytes) -> list: + wide_map = {f"k{i}": i for i in range(20)} + return [ + { + "request_id": "req-0", + "seed": 7, + "scale": 1.5, + "flag": True, + "missing": None, + "dit_trajectory": { + "timesteps": {"__tensor__": True, "data": b"tiny", "shape": [4]}, + "latents": {"__tensor__": True, "data": big, "shape": [2, 3]}, + }, + "wide": wide_map, + "long_list": list(range(30)), + } + ] + + +class TestMsgpackEncodeSpliced(unittest.TestCase): + def test_byte_identical_to_msgspec(self): + big = bytes(range(256)) * 8192 + payload = _sample_payload(big) + parts = msgpack_encode_spliced(payload, threshold=1 << 10) + self.assertEqual(b"".join(parts), msgspec.msgpack.encode(payload)) + + def test_large_bytes_spliced_by_reference(self): + big = b"\x00" * (1 << 20) + payload = _sample_payload(big) + parts = msgpack_encode_spliced(payload, threshold=1 << 10) + self.assertTrue(any(part is big for part in parts)) + decoded = msgspec.msgpack.decode(b"".join(parts)) + self.assertEqual(decoded[0]["dit_trajectory"]["latents"]["data"], big) + self.assertEqual(decoded[0]["dit_trajectory"]["timesteps"]["data"], b"tiny") + + +if __name__ == "__main__": + unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_rollout_api.py b/python/sglang/multimodal_gen/test/unit/test_rollout_api.py index 64b0298c0f12..e4f84d27f03a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_rollout_api.py +++ b/python/sglang/multimodal_gen/test/unit/test_rollout_api.py @@ -1,10 +1,24 @@ -"""Unit tests for the rollout generate API (serialization, io_struct, rollout_api).""" +"""Unit tests for the rollout generate API (serialization, io_struct, rollout_api). +Request timing model (TestRequestStamps / TestStagesHeader): + + srv_recv .. forward .. build .. dump .. msgpack (absolute wall-clock marks) + -> all marks ride the x-sgld-timing response header + stages_header: engine per-stage milliseconds -> x-sgld-stages header +""" + +import json +import time import types import unittest import torch +from sglang.multimodal_gen.runtime.entrypoints.post_training.request_timing import ( + MARKS, + RequestStamps, + stages_header, +) from sglang.multimodal_gen.runtime.entrypoints.post_training.utils import ( _maybe_deserialize, _maybe_serialize, @@ -416,5 +430,66 @@ def test_sampling_params_exposes_filters_via_req_getattr(self): self.assertEqual(req.rollout_return_step_indices, [1, 3]) +class TestRequestStamps(unittest.TestCase): + """The client joins these marks with its own, so what matters is that the + header decodes, keeps wall-clock order, and stays readable when a request + returned early with only some marks taken.""" + + def _decode(self, stamps: RequestStamps) -> dict: + header = stamps.to_header() + self.assertEqual(header, header.encode("ascii").decode()) + return json.loads(header) + + def test_marks_decode_in_wall_clock_order(self): + before = time.time() + stamps = RequestStamps("req-1") + for name in MARKS: + stamps.mark(name) + decoded = self._decode(stamps) + self.assertEqual(decoded.pop("request_id"), "req-1") + self.assertEqual(set(decoded), set(MARKS)) + values = [decoded[name] for name in MARKS] + self.assertEqual(values, sorted(values)) + self.assertGreaterEqual(values[0], round(before, 6) - 1e-6) + + def test_span_pairs_start_and_end_marks(self): + stamps = RequestStamps() + with stamps.span("forward"): + pass + with stamps.span("dump"): + pass + self.assertEqual( + set(json.loads(stamps.to_header())), + {"forward_start", "forward_end", "dump_end"}, + ) + + def test_partial_marks_only_report_what_was_taken(self): + stamps = RequestStamps("req-3") + stamps.mark("srv_recv") + self.assertEqual(set(self._decode(stamps)), {"srv_recv", "request_id"}) + + def test_unknown_mark_is_rejected(self): + with self.assertRaises(AssertionError): + RequestStamps().mark("no_such_mark") + + +class TestStagesHeader(unittest.TestCase): + """The engine's own stage breakdown is the only view inside the forward the + client waits on, so it has to survive the trip and stay absent when unset.""" + + def test_stages_are_reported_in_milliseconds(self): + metrics = types.SimpleNamespace( + stages={"decoding": 8123.4567, "text_encoding": 91.2} + ) + self.assertEqual( + json.loads(stages_header(metrics)), + {"decoding": 8123.457, "text_encoding": 91.2}, + ) + + def test_no_metrics_or_no_stages_yields_no_header(self): + self.assertEqual(stages_header(None), "") + self.assertEqual(stages_header(types.SimpleNamespace(stages={})), "") + + if __name__ == "__main__": unittest.main()