[Diffusion] msgpack raw-bytes transport (drop base64/JSON) (#31565)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Rockdu
2026-07-20 11:32:02 -07:00
committed by GitHub
co-authored by Claude Opus 4.8
parent 5ab3d90b81
commit 4e8eb1457b
3 changed files with 42 additions and 33 deletions
@@ -4,9 +4,9 @@ from __future__ import annotations
from typing import Any from typing import Any
import msgspec
import torch import torch
from fastapi import APIRouter, HTTPException from fastapi import APIRouter, HTTPException, Response
from fastapi.responses import ORJSONResponse
from sglang.multimodal_gen.configs.sample.sampling_params import generate_request_id 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.openai.utils import build_sampling_params
@@ -298,7 +298,16 @@ def _build_sampling_kwargs(request: RolloutRequest) -> dict:
return {k: v for k, v in sampling_kwargs.items() if v is not None} return {k: v for k, v in sampling_kwargs.items() if v is not None}
@router.post("/generate", response_model=list[RolloutResponse]) @router.post(
"/generate",
response_class=Response,
responses={
200: {
"model": list[RolloutResponse],
"content": {"application/msgpack": {}},
}
},
)
async def rollout_generate(request: RolloutRequest): async def rollout_generate(request: RolloutRequest):
request_id = generate_request_id() request_id = generate_request_id()
server_args = get_global_server_args() server_args = get_global_server_args()
@@ -326,4 +335,8 @@ async def rollout_generate(request: RolloutRequest):
rollout_responses = _build_response( rollout_responses = _build_response(
request_id, request.prompt, request.seed, request.rollout, output_batch request_id, request.prompt, request.seed, request.rollout, output_batch
) )
return ORJSONResponse(content=[r.model_dump() for r in rollout_responses]) payload = [r.model_dump() for r in rollout_responses]
return Response(
content=msgspec.msgpack.encode(payload),
media_type="application/msgpack",
)
@@ -2,7 +2,6 @@
from __future__ import annotations from __future__ import annotations
import base64
from typing import Any from typing import Any
import numpy as np import numpy as np
@@ -10,22 +9,19 @@ import torch
from safetensors.torch import load, save from safetensors.torch import load, save
def tensor_to_base64(t: torch.Tensor) -> str: def tensor_to_bytes(t: torch.Tensor) -> bytes:
t = t.detach().contiguous().cpu() return save({"t": t.detach().contiguous().cpu()})
raw = save({"t": t})
return base64.b64encode(raw).decode("ascii")
def base64_to_tensor(s: str) -> torch.Tensor: def bytes_to_tensor(b: bytes) -> torch.Tensor:
raw = base64.b64decode(s) return load(b)["t"]
return load(raw)["t"]
def _maybe_serialize(obj: Any) -> Any: def _maybe_serialize(obj: Any) -> Any:
if isinstance(obj, torch.Tensor): if isinstance(obj, torch.Tensor):
return { return {
"__tensor__": True, "__tensor__": True,
"data": tensor_to_base64(obj), "data": tensor_to_bytes(obj),
"shape": list(obj.shape), "shape": list(obj.shape),
"dtype": str(obj.dtype), "dtype": str(obj.dtype),
} }
@@ -41,7 +37,7 @@ def _maybe_serialize(obj: Any) -> Any:
def _maybe_deserialize(obj: Any) -> Any: def _maybe_deserialize(obj: Any) -> Any:
if isinstance(obj, dict): if isinstance(obj, dict):
if obj.get("__tensor__"): if obj.get("__tensor__"):
return base64_to_tensor(obj["data"]) return bytes_to_tensor(obj["data"])
return {k: _maybe_deserialize(v) for k, v in obj.items()} return {k: _maybe_deserialize(v) for k, v in obj.items()}
if isinstance(obj, (list, tuple)): if isinstance(obj, (list, tuple)):
return [_maybe_deserialize(v) for v in obj] return [_maybe_deserialize(v) for v in obj]
@@ -8,8 +8,8 @@ import torch
from sglang.multimodal_gen.runtime.entrypoints.post_training.utils import ( from sglang.multimodal_gen.runtime.entrypoints.post_training.utils import (
_maybe_deserialize, _maybe_deserialize,
_maybe_serialize, _maybe_serialize,
base64_to_tensor, bytes_to_tensor,
tensor_to_base64, tensor_to_bytes,
) )
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import ( from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import (
@@ -20,12 +20,12 @@ from sglang.multimodal_gen.runtime.post_training.rl_dataclasses import (
) )
class TestTensorToBase64Roundtrip(unittest.TestCase): class TestTensorToBytesRoundtrip(unittest.TestCase):
def _roundtrip(self, t: torch.Tensor): def _roundtrip(self, t: torch.Tensor):
encoded = tensor_to_base64(t) encoded = tensor_to_bytes(t)
self.assertIsInstance(encoded, str) self.assertIsInstance(encoded, bytes)
decoded = base64_to_tensor(encoded) decoded = bytes_to_tensor(encoded)
self.assertTrue( self.assertTrue(
torch.equal(t, decoded), f"Mismatch for shape={t.shape} dtype={t.dtype}" torch.equal(t, decoded), f"Mismatch for shape={t.shape} dtype={t.dtype}"
) )
@@ -55,21 +55,21 @@ class TestTensorToBase64Roundtrip(unittest.TestCase):
if not torch.cuda.is_available(): if not torch.cuda.is_available():
self.skipTest("CUDA not available") self.skipTest("CUDA not available")
t = torch.randn(4, device="cuda") t = torch.randn(4, device="cuda")
encoded = tensor_to_base64(t) encoded = tensor_to_bytes(t)
decoded = base64_to_tensor(encoded) decoded = bytes_to_tensor(encoded)
self.assertTrue(torch.equal(t.cpu(), decoded)) self.assertTrue(torch.equal(t.cpu(), decoded))
def test_non_contiguous(self): def test_non_contiguous(self):
t = torch.randn(4, 6)[:, ::2] t = torch.randn(4, 6)[:, ::2]
self.assertFalse(t.is_contiguous()) self.assertFalse(t.is_contiguous())
self._roundtrip(t.contiguous()) self._roundtrip(t.contiguous())
decoded = base64_to_tensor(tensor_to_base64(t)) decoded = bytes_to_tensor(tensor_to_bytes(t))
self.assertTrue(torch.equal(t.contiguous(), decoded)) self.assertTrue(torch.equal(t.contiguous(), decoded))
def test_grad_tensor_detaches(self): def test_grad_tensor_detaches(self):
t = torch.randn(3, requires_grad=True) t = torch.randn(3, requires_grad=True)
encoded = tensor_to_base64(t) encoded = tensor_to_bytes(t)
decoded = base64_to_tensor(encoded) decoded = bytes_to_tensor(encoded)
self.assertFalse(decoded.requires_grad) self.assertFalse(decoded.requires_grad)
self.assertTrue(torch.equal(t.detach(), decoded)) self.assertTrue(torch.equal(t.detach(), decoded))
@@ -82,7 +82,7 @@ class TestMaybeSerialize(unittest.TestCase):
self.assertTrue(result["__tensor__"]) self.assertTrue(result["__tensor__"])
self.assertEqual(result["shape"], [2, 3]) self.assertEqual(result["shape"], [2, 3])
self.assertEqual(result["dtype"], "torch.float32") self.assertEqual(result["dtype"], "torch.float32")
decoded = base64_to_tensor(result["data"]) decoded = bytes_to_tensor(result["data"])
self.assertTrue(torch.equal(t, decoded)) self.assertTrue(torch.equal(t, decoded))
def test_dict_with_tensors(self): def test_dict_with_tensors(self):
@@ -243,7 +243,7 @@ class TestBuildResponse(unittest.TestCase):
self.assertEqual(resp.seed, 42) self.assertEqual(resp.seed, 42)
self.assertIsNotNone(resp.generated_output) self.assertIsNotNone(resp.generated_output)
self.assertIsNotNone(resp.rollout_log_probs) self.assertIsNotNone(resp.rollout_log_probs)
lp = base64_to_tensor(resp.rollout_log_probs["data"]) lp = bytes_to_tensor(resp.rollout_log_probs["data"])
self.assertEqual(lp.shape, ()) self.assertEqual(lp.shape, ())
self.assertAlmostEqual(resp.inference_time_s, 2.5) self.assertAlmostEqual(resp.inference_time_s, 2.5)
@@ -308,12 +308,12 @@ class TestBuildResponse(unittest.TestCase):
batch.metrics = self._make_metrics(1.0) batch.metrics = self._make_metrics(1.0)
resps = _build_response("rb", "p", 0, True, batch) resps = _build_response("rb", "p", 0, True, batch)
self.assertEqual(len(resps), B) self.assertEqual(len(resps), B)
lp0 = base64_to_tensor(resps[0].rollout_log_probs["data"]) lp0 = bytes_to_tensor(resps[0].rollout_log_probs["data"])
lp1 = base64_to_tensor(resps[1].rollout_log_probs["data"]) lp1 = bytes_to_tensor(resps[1].rollout_log_probs["data"])
self.assertEqual(lp0.shape, (T,)) self.assertEqual(lp0.shape, (T,))
self.assertEqual(lp1.shape, (T,)) self.assertEqual(lp1.shape, (T,))
g0 = base64_to_tensor(resps[0].generated_output["data"]) g0 = bytes_to_tensor(resps[0].generated_output["data"])
g1 = base64_to_tensor(resps[1].generated_output["data"]) g1 = bytes_to_tensor(resps[1].generated_output["data"])
self.assertEqual(g0.shape, (1, 8, 8)) self.assertEqual(g0.shape, (1, 8, 8))
self.assertEqual(g1.shape, (1, 8, 8)) self.assertEqual(g1.shape, (1, 8, 8))
self.assertFalse(torch.equal(g0, g1)) self.assertFalse(torch.equal(g0, g1))
@@ -335,8 +335,8 @@ class TestBuildResponse(unittest.TestCase):
self.assertEqual(len(resps), B) self.assertEqual(len(resps), B)
self.assertIsNotNone(resps[0].dit_trajectory) self.assertIsNotNone(resps[0].dit_trajectory)
self.assertIsNotNone(resps[1].dit_trajectory) self.assertIsNotNone(resps[1].dit_trajectory)
ts0 = base64_to_tensor(resps[0].dit_trajectory["timesteps"]["data"]) ts0 = bytes_to_tensor(resps[0].dit_trajectory["timesteps"]["data"])
ts1 = base64_to_tensor(resps[1].dit_trajectory["timesteps"]["data"]) ts1 = bytes_to_tensor(resps[1].dit_trajectory["timesteps"]["data"])
self.assertEqual(ts0.shape, (T,)) self.assertEqual(ts0.shape, (T,))
self.assertTrue(torch.equal(ts0, ts1)) self.assertTrue(torch.equal(ts0, ts1))
self.assertEqual( self.assertEqual(