[diffusion] feat: support saving videos directly on the server to avoid the overhead of tensor transfer (#18253)

This commit is contained in:
wxy
2026-02-07 22:08:42 +08:00
committed by GitHub
parent 31d4cd2ffd
commit 64950d8f97
7 changed files with 188 additions and 110 deletions
@@ -153,6 +153,8 @@ class SamplingParams:
# if True, suppress verbose logging for this request # if True, suppress verbose logging for this request
suppress_logs: bool = False suppress_logs: bool = False
return_file_paths_only: bool = True
def _set_output_file_ext(self): def _set_output_file_ext(self):
# add extension if needed # add extension if needed
if not any( if not any(
@@ -738,6 +740,12 @@ class SamplingParams:
"Default: true. Examples: --adjust-frames, --adjust-frames true, --adjust-frames false." "Default: true. Examples: --adjust-frames, --adjust-frames true, --adjust-frames false."
), ),
) )
parser.add_argument(
"--return-file-paths-only",
action=StoreBoolean,
default=SamplingParams.return_file_paths_only,
help="If set, output file will be saved early to get a performance boost, while output tensors will not be returned.",
)
return parser return parser
@classmethod @classmethod
@@ -807,9 +815,6 @@ class SamplingParams:
n_tokens = -1 n_tokens = -1
return n_tokens return n_tokens
def output_file_path(self):
return os.path.join(self.output_path, self.output_file_name)
@dataclass @dataclass
class CacheParams: class CacheParams:
@@ -14,12 +14,8 @@ import time
from typing import Any, List, Union from typing import Any, List, Union
import numpy as np import numpy as np
import torch
from sglang.multimodal_gen.configs.sample.sampling_params import ( from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
DataType,
SamplingParams,
)
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import ( from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
ListLorasReq, ListLorasReq,
MergeLoraWeightsReq, MergeLoraWeightsReq,
@@ -28,8 +24,8 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
format_lora_message, format_lora_message,
) )
from sglang.multimodal_gen.runtime.entrypoints.utils import ( from sglang.multimodal_gen.runtime.entrypoints.utils import (
post_process_sample,
prepare_request, prepare_request,
save_outputs,
) )
from sglang.multimodal_gen.runtime.launch_server import launch_server from sglang.multimodal_gen.runtime.launch_server import launch_server
from sglang.multimodal_gen.runtime.pipelines_core import Req from sglang.multimodal_gen.runtime.pipelines_core import Req
@@ -229,49 +225,64 @@ class DiffGenerator:
if output_batch.error: if output_batch.error:
raise Exception(f"{output_batch.error}") raise Exception(f"{output_batch.error}")
if output_batch.output is None: if (
output_batch.output is None
and output_batch.output_file_paths is None
):
logger.error( logger.error(
"Received empty output from scheduler for prompt %d", "Received empty output from scheduler for prompt %d",
request_idx + 1, request_idx + 1,
) )
continue continue
audio_sample_rate = output_batch.audio_sample_rate audio_sample_rate = output_batch.audio_sample_rate
for output_idx, sample in enumerate(output_batch.output):
num_outputs = len(output_batch.output)
audio = output_batch.audio
if req.data_type == DataType.VIDEO:
if isinstance(audio, torch.Tensor) and audio.ndim >= 2:
audio = (
audio[output_idx]
if audio.shape[0] > output_idx
else None
)
elif isinstance(audio, np.ndarray) and audio.ndim >= 2:
audio = (
audio[output_idx]
if audio.shape[0] > output_idx
else None
)
if audio is not None and not (
isinstance(sample, (tuple, list)) and len(sample) == 2
):
sample = (sample, audio)
frames = post_process_sample(
sample,
fps=req.fps,
save_output=req.save_output,
# TODO: output file path for req should be determined
save_file_path=req.output_file_path(
num_outputs, output_idx
),
data_type=req.data_type,
audio_sample_rate=audio_sample_rate,
)
if req.save_output and req.return_file_paths_only:
for output_idx, output_path in enumerate(
output_batch.output_file_paths
):
result_item: dict[str, Any] = {
"samples": None,
"frames": None,
"audio": None,
"prompts": req.prompt,
"size": (req.height, req.width, req.num_frames),
"generation_time": timer.duration,
"peak_memory_mb": output_batch.peak_memory_mb,
"timings": (
output_batch.timings.to_dict()
if output_batch.timings
else {}
),
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
"prompt_index": output_idx,
"output_file_path": output_path,
}
results.append(result_item)
continue
samples_out: list[Any] = []
audios_out: list[Any] = []
frames_out: list[Any] = []
save_outputs(
output_batch.output,
req.data_type,
req.fps,
req.save_output,
lambda idx: req.output_file_path(len(output_batch.output), idx),
audio=output_batch.audio,
audio_sample_rate=audio_sample_rate,
samples_out=samples_out,
audios_out=audios_out,
frames_out=frames_out,
)
for output_idx in range(len(samples_out)):
result_item: dict[str, Any] = { result_item: dict[str, Any] = {
"samples": sample, "samples": samples_out[output_idx],
"frames": frames, "frames": frames_out[output_idx],
"audio": audio, "audio": audios_out[output_idx],
"prompts": req.prompt, "prompts": req.prompt,
"size": (req.height, req.width, req.num_frames), "size": (req.height, req.width, req.num_frames),
"generation_time": timer.duration, "generation_time": timer.duration,
@@ -16,8 +16,8 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import (
VertexGenerateReqInput, VertexGenerateReqInput,
) )
from sglang.multimodal_gen.runtime.entrypoints.utils import ( from sglang.multimodal_gen.runtime.entrypoints.utils import (
post_process_sample,
prepare_request, prepare_request,
save_outputs,
) )
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
@@ -114,33 +114,22 @@ async def forward_to_scheduler(req_obj, sp):
"""Forwards request to scheduler and processes the result.""" """Forwards request to scheduler and processes the result."""
try: try:
response = await async_scheduler_client.forward(req_obj) response = await async_scheduler_client.forward(req_obj)
if response.output is None: if response.output is None and response.output_file_paths is None:
raise RuntimeError("Model generation returned no output.") raise RuntimeError("Model generation returned no output.")
output_file_path = sp.output_file_path() if response.output_file_paths:
sample = response.output[0] output_file_path = response.output_file_paths[0]
try: else:
audio = response.audio output_file_path = sp.output_file_path()
except AttributeError: save_outputs(
audio = None [response.output[0]],
if isinstance(audio, torch.Tensor) and audio.ndim >= 2: sp.data_type,
audio = audio[0] sp.fps,
if audio is not None and not ( True,
isinstance(sample, (tuple, list)) and len(sample) == 2 lambda _idx: output_file_path,
): audio=response.audio,
sample = (sample, audio) audio_sample_rate=response.audio_sample_rate,
post_process_sample( )
sample=sample,
data_type=sp.data_type,
fps=sp.fps or 24,
save_output=True,
save_file_path=output_file_path,
audio_sample_rate=(
response.audio_sample_rate
if hasattr(response, "audio_sample_rate")
else None
),
)
if hasattr(response, "model_dump"): if hasattr(response, "model_dump"):
data = response.model_dump() data = response.model_dump()
@@ -7,11 +7,10 @@ import time
from typing import Any, List, Optional, Union from typing import Any, List, Optional, Union
import httpx import httpx
import torch
from fastapi import UploadFile from fastapi import UploadFile
from sglang.multimodal_gen.configs.sample.sampling_params import DataType from sglang.multimodal_gen.configs.sample.sampling_params import DataType
from sglang.multimodal_gen.runtime.entrypoints.utils import post_process_sample from sglang.multimodal_gen.runtime.entrypoints.utils import save_outputs
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.scheduler_client import AsyncSchedulerClient from sglang.multimodal_gen.runtime.scheduler_client import AsyncSchedulerClient
from sglang.multimodal_gen.runtime.utils.logging_utils import ( from sglang.multimodal_gen.runtime.utils.logging_utils import (
@@ -211,51 +210,43 @@ async def process_generation_batch(
with log_generation_timer(logger, batch.prompt): with log_generation_timer(logger, batch.prompt):
result = await scheduler_client.forward([batch]) result = await scheduler_client.forward([batch])
if result.output is None: if result.output is None and result.output_file_paths is None:
error_msg = result.error or "Unknown error" error_msg = result.error or "Unknown error"
raise RuntimeError( raise RuntimeError(
f"Model generation returned no output. Error from scheduler: {error_msg}" f"Model generation returned no output. Error from scheduler: {error_msg}"
) )
save_file_path_list = [] save_file_path_list = []
audio_sample_rate = result.audio_sample_rate # If output_file_paths is provided, use it instead of output.
if batch.data_type == DataType.VIDEO: if result.output_file_paths:
for idx, output in enumerate(result.output): save_file_path_list = result.output_file_paths
save_file_path = str(
os.path.join(batch.output_path, batch.output_file_name)
)
sample = result.output[idx]
audio = result.audio
if isinstance(audio, torch.Tensor) and audio.ndim >= 2:
audio = audio[idx] if audio.shape[0] > idx else None
if audio is not None and not (
isinstance(sample, (tuple, list)) and len(sample) == 2
):
sample = (sample, audio)
post_process_sample(
sample,
batch.data_type,
batch.fps,
batch.save_output,
save_file_path,
audio_sample_rate=audio_sample_rate,
)
save_file_path_list.append(save_file_path)
else: else:
for idx, output in enumerate(result.output): audio_sample_rate = result.audio_sample_rate
save_file_path = str( if batch.data_type == DataType.VIDEO:
os.path.join( save_file_path_list = save_outputs(
batch.output_path, f"sample_{idx}_" + batch.output_file_name result.output,
)
)
post_process_sample(
output,
batch.data_type, batch.data_type,
batch.fps, batch.fps,
batch.save_output, batch.save_output,
save_file_path, lambda _idx: str(
os.path.join(batch.output_path, batch.output_file_name)
),
audio=result.audio,
audio_sample_rate=audio_sample_rate,
)
else:
save_file_path_list = save_outputs(
result.output,
batch.data_type,
batch.fps,
batch.save_output,
lambda idx: str(
os.path.join(
batch.output_path,
f"sample_{idx}_" + batch.output_file_name,
)
),
audio_sample_rate=audio_sample_rate, audio_sample_rate=audio_sample_rate,
) )
save_file_path_list.append(save_file_path)
total_time = time.perf_counter() - total_start_time total_time = time.perf_counter() - total_start_time
log_batch_completion(logger, 1, total_time) log_batch_completion(logger, 1, total_time)
@@ -12,7 +12,7 @@ import os
import shutil import shutil
import subprocess import subprocess
import tempfile import tempfile
from typing import Any, Optional from typing import Any, Callable, Optional, Sequence
import imageio import imageio
import numpy as np import numpy as np
@@ -234,6 +234,72 @@ def prepare_request(
return req return req
def attach_audio_to_video_sample(
sample: Any,
audio: Any,
output_idx: int,
) -> Any:
"""Attach per-sample audio for video outputs when available."""
if audio is None:
return sample
if isinstance(audio, torch.Tensor) and audio.ndim >= 2:
audio = audio[output_idx] if audio.shape[0] > output_idx else None
elif isinstance(audio, np.ndarray) and audio.ndim >= 2:
audio = audio[output_idx] if audio.shape[0] > output_idx else None
if audio is not None and not (
isinstance(sample, (tuple, list)) and len(sample) == 2
):
return (sample, audio)
return sample
def save_outputs(
outputs: Sequence[Any],
data_type: DataType,
fps: int,
save_output: bool,
build_output_path: Callable[[int], str],
*,
audio: Any = None,
audio_sample_rate: Optional[int] = None,
samples_out: Optional[list[Any]] = None,
audios_out: Optional[list[Any]] = None,
frames_out: Optional[list[Any]] = None,
) -> list[str]:
"""Save outputs to files and return the list of file paths."""
output_paths: list[str] = []
for idx, output in enumerate(outputs):
save_file_path = build_output_path(idx)
sample = output
if data_type == DataType.VIDEO:
sample = attach_audio_to_video_sample(sample, audio, idx)
frames = post_process_sample(
sample,
data_type,
fps,
save_output,
save_file_path,
audio_sample_rate=audio_sample_rate,
)
if samples_out is not None:
samples_out.append(sample)
if audios_out is not None:
if data_type == DataType.VIDEO:
audio_item = audio
if isinstance(audio, torch.Tensor) and audio.ndim >= 2:
audio_item = audio[idx] if audio.shape[0] > idx else None
elif isinstance(audio, np.ndarray) and audio.ndim >= 2:
audio_item = audio[idx] if audio.shape[0] > idx else None
audios_out.append(audio_item)
else:
audios_out.append(audio)
if frames_out is not None:
frames_out.append(frames)
output_paths.append(save_file_path)
return output_paths
def post_process_sample( def post_process_sample(
sample: Any, sample: Any,
data_type: DataType, data_type: DataType,
@@ -18,6 +18,7 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_cfg_group, get_cfg_group,
get_tp_group, get_tp_group,
) )
from sglang.multimodal_gen.runtime.entrypoints.utils import save_outputs
from sglang.multimodal_gen.runtime.pipelines_core import ( from sglang.multimodal_gen.runtime.pipelines_core import (
ComposedPipelineBase, ComposedPipelineBase,
LoRAPipeline, LoRAPipeline,
@@ -186,6 +187,21 @@ class GPUWorker:
duration_ms = (time.monotonic() - start_time) * 1000 duration_ms = (time.monotonic() - start_time) * 1000
output_batch.timings.total_duration_ms = duration_ms output_batch.timings.total_duration_ms = duration_ms
# Save output to file and return file path only if requested. Avoid the serialization
# and deserialization overhead between scheduler_client and gpu_worker.
if req.save_output and req.return_file_paths_only:
output_paths = save_outputs(
output_batch.output,
req.data_type,
req.fps,
True,
lambda idx: req.output_file_path(len(output_batch.output), idx),
audio=output_batch.audio,
audio_sample_rate=output_batch.audio_sample_rate,
)
output_batch.output_file_paths = output_paths
output_batch.output = None
# TODO: extract to avoid duplication # TODO: extract to avoid duplication
if req.perf_dump_path is not None or envs.SGLANG_DIFFUSION_STAGE_LOGGING: if req.perf_dump_path is not None or envs.SGLANG_DIFFUSION_STAGE_LOGGING:
# Avoid logging warmup perf records that share the same request_id. # Avoid logging warmup perf records that share the same request_id.
@@ -198,8 +214,7 @@ class GPUWorker:
if output_batch is None: if output_batch is None:
output_batch = OutputBatch() output_batch = OutputBatch()
output_batch.error = f"Error executing request {req.request_id}: {e}" output_batch.error = f"Error executing request {req.request_id}: {e}"
finally: return output_batch
return output_batch
def get_can_stay_resident_components( def get_can_stay_resident_components(
self, remaining_gpu_mem_gb: float self, remaining_gpu_mem_gb: float
@@ -100,7 +100,6 @@ class Req:
raw_audio_latent_shape: tuple[int, ...] | None = None raw_audio_latent_shape: tuple[int, ...] | None = None
# Audio Parameters # Audio Parameters
fps: float = 24.0
generate_audio: bool = True generate_audio: bool = True
raw_latent_shape: torch.Tensor | None = None raw_latent_shape: torch.Tensor | None = None
@@ -295,6 +294,7 @@ class Req:
width: {target_width} width: {target_width}
height: {target_height} height: {target_height}
num_frames: {self.num_frames} num_frames: {self.num_frames}
fps: {self.fps}
prompt: {self.prompt} prompt: {self.prompt}
neg_prompt: {self.negative_prompt} neg_prompt: {self.negative_prompt}
seed: {self.seed} seed: {self.seed}
@@ -325,6 +325,7 @@ class OutputBatch:
trajectory_latents: torch.Tensor | None = None trajectory_latents: torch.Tensor | None = None
trajectory_decoded: list[torch.Tensor] | None = None trajectory_decoded: list[torch.Tensor] | None = None
error: str | None = None error: str | None = None
output_file_paths: list[str] | None = None
# logged timings info, directly from Req.timings # logged timings info, directly from Req.timings
timings: Optional["RequestTimings"] = None timings: Optional["RequestTimings"] = None