[diffusion] optimize: stream and parallelize bit-exact video output saves (#34564)

This commit is contained in:
Mick
2026-08-12 20:03:37 +08:00
committed by GitHub
parent 9701cc138c
commit dc5f6c4883
3 changed files with 299 additions and 89 deletions
@@ -16,6 +16,7 @@ import shutil
import subprocess import subprocess
import tempfile import tempfile
import threading import threading
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager from contextlib import contextmanager
from copy import copy from copy import copy
from dataclasses import dataclass, field from dataclasses import dataclass, field
@@ -48,6 +49,8 @@ from sglang.srt.observability.trace import TraceReqContext
logger = init_logger(__name__) logger = init_logger(__name__)
_MAX_CACHED_CUDA_VIDEO_BUFFER_BYTES = 1024 * 1024 * 1024 _MAX_CACHED_CUDA_VIDEO_BUFFER_BYTES = 1024 * 1024 * 1024
_MAX_CUDA_VIDEO_CONVERSION_CHUNK_BYTES = 128 * 1024 * 1024
_MAX_PARALLEL_CUDA_VIDEO_SAVES = 2
_cuda_video_buffer_cache_lock = threading.Lock() _cuda_video_buffer_cache_lock = threading.Lock()
_cached_cuda_video_buffer: "_CudaMemfdVideoBuffer | None" = None _cached_cuda_video_buffer: "_CudaMemfdVideoBuffer | None" = None
@@ -464,6 +467,25 @@ def _x264_auto_thread_count(height: int) -> int:
return min(cpu_limit, row_limit, 128) return min(cpu_limit, row_limit, 128)
def _cuda_video_conversion_chunk_frames(video: torch.Tensor) -> int:
_, num_frames, height, width = video.shape
temporary_bytes_per_frame = 3 * height * width * (video.element_size() + 2)
return min(
num_frames,
max(1, _MAX_CUDA_VIDEO_CONVERSION_CHUNK_BYTES // temporary_bytes_per_frame),
)
def _sendfile_all(output_fd: int, input_fd: int, count: int) -> None:
offset = 0
while count:
sent = os.sendfile(output_fd, input_fd, offset, count)
if sent <= 0:
raise RuntimeError("sendfile made no progress")
offset += sent
count -= sent
def _try_save_cuda_video_direct( def _try_save_cuda_video_direct(
*, *,
save_file_path: str, save_file_path: str,
@@ -472,8 +494,8 @@ def _try_save_cuda_video_direct(
audio_sample_rate: Optional[int], audio_sample_rate: Optional[int],
output_compression: Optional[int], output_compression: Optional[int],
) -> bool: ) -> bool:
"""Save a CUDA RGB video through a registered memfd instead of a Python pipe.""" """Stream CUDA RGB chunks to ffmpeg through a registered memfd."""
if not hasattr(os, "memfd_create") or not os.path.isdir("/proc/self/fd"): if not hasattr(os, "memfd_create") or not hasattr(os, "sendfile"):
return False return False
sample_without_audio, audio = _split_sample_audio(sample) sample_without_audio, audio = _split_sample_audio(sample)
@@ -492,9 +514,8 @@ def _try_save_cuda_video_direct(
if video.shape[0] != 3: if video.shape[0] != 3:
return False return False
frames = (video * 255).clamp_(0, 255).to(torch.uint8) _, num_frames, height, width = video.shape
frames = frames.permute(1, 2, 3, 0).contiguous() chunk_frames = _cuda_video_conversion_chunk_frames(video)
num_frames, height, width, _ = frames.shape
quality = output_compression / 10 if output_compression is not None else 5 quality = output_compression / 10 if output_compression is not None else 5
if not 1 <= quality <= 10: if not 1 <= quality <= 10:
@@ -518,76 +539,105 @@ def _try_save_cuda_video_direct(
scipy_wavfile.write(tmp_wav_path, selected_sr, audio_np) scipy_wavfile.write(tmp_wav_path, selected_sr, audio_np)
ffmpeg_exe = _resolve_ffmpeg_exe() ffmpeg_exe = _resolve_ffmpeg_exe()
shape = tuple(frames.shape) command = [
with _acquire_cuda_video_buffer(shape) as buffer: ffmpeg_exe,
assert buffer.tensor is not None "-y",
buffer.tensor.copy_(frames, non_blocking=True) "-f",
torch.cuda.current_stream(frames.device).synchronize() "rawvideo",
del frames "-vcodec",
os.lseek(buffer.fd, 0, os.SEEK_SET) "rawvideo",
"-s",
f"{width}x{height}",
"-pix_fmt",
"rgb24",
"-r",
f"{fps:.02f}",
"-i",
"pipe:0",
]
if tmp_wav_path is None:
command += ["-an"]
else:
command += ["-i", tmp_wav_path]
command += [
"-vcodec",
"libx264",
"-pix_fmt",
"yuv420p",
"-crf",
str(crf),
]
command = [ macro_block_size = 16
ffmpeg_exe, if width % macro_block_size or height % macro_block_size:
"-y", output_width = (
"-f", width
"rawvideo", if width % macro_block_size == 0
"-vcodec", else width + macro_block_size - width % macro_block_size
"rawvideo",
"-s",
f"{width}x{height}",
"-pix_fmt",
"rgb24",
"-r",
f"{fps:.02f}",
"-i",
f"/proc/self/fd/{buffer.fd}",
]
if tmp_wav_path is None:
command += ["-an"]
else:
command += ["-i", tmp_wav_path]
command += [
"-vcodec",
"libx264",
"-pix_fmt",
"yuv420p",
"-crf",
str(crf),
]
macro_block_size = 16
if width % macro_block_size or height % macro_block_size:
output_width = (
width
if width % macro_block_size == 0
else width + macro_block_size - width % macro_block_size
)
output_height = (
height
if height % macro_block_size == 0
else height + macro_block_size - height % macro_block_size
)
command += ["-vf", f"scale={output_width}:{output_height}"]
command += ["-threads", str(_x264_auto_thread_count(height))]
if tmp_wav_path is not None:
command += [
"-acodec",
"aac",
"-map",
"0:v:0",
"-map",
"1:a:0",
]
command += ["-v", "warning", save_file_path]
subprocess.run(
command,
check=True,
pass_fds=(buffer.fd,),
stdout=subprocess.DEVNULL,
stderr=subprocess.PIPE,
) )
output_height = (
height
if height % macro_block_size == 0
else height + macro_block_size - height % macro_block_size
)
command += ["-vf", f"scale={output_width}:{output_height}"]
command += ["-threads", str(_x264_auto_thread_count(height))]
if tmp_wav_path is not None:
command += [
"-acodec",
"aac",
"-map",
"0:v:0",
"-map",
"1:a:0",
]
command += ["-v", "warning", save_file_path]
with tempfile.TemporaryFile() as stderr_file:
process = subprocess.Popen(
command,
stdin=subprocess.PIPE,
stdout=subprocess.DEVNULL,
stderr=stderr_file,
)
try:
if process.stdin is None:
raise RuntimeError("ffmpeg stdin pipe was not created")
with _acquire_cuda_video_buffer(
(chunk_frames, height, width, 3)
) as buffer:
assert buffer.tensor is not None
for start in range(0, num_frames, chunk_frames):
end = min(start + chunk_frames, num_frames)
frames = (
(video[:, start:end] * 255).clamp_(0, 255).to(torch.uint8)
)
frames = frames.permute(1, 2, 3, 0).contiguous()
buffer.tensor[: end - start].copy_(frames, non_blocking=True)
torch.cuda.current_stream(video.device).synchronize()
del frames
_sendfile_all(
process.stdin.fileno(),
buffer.fd,
(end - start) * height * width * 3,
)
process.stdin.close()
process.stdin = None
returncode = process.wait()
finally:
if process.stdin is not None:
process.stdin.close()
if process.poll() is None:
process.kill()
process.wait()
if returncode:
stderr_file.seek(0)
raise subprocess.CalledProcessError(
returncode,
command,
stderr=stderr_file.read(),
)
return True return True
except Exception as e: except Exception as e:
logger.warning( logger.warning(
@@ -603,6 +653,86 @@ def _try_save_cuda_video_direct(
pass pass
def _try_save_cuda_videos_direct(
samples: Sequence[Any],
save_file_paths: Sequence[str],
*,
fps: int,
audio_sample_rate: Optional[int],
output_compression: Optional[int],
) -> list[bool] | None:
"""Save independent CUDA videos concurrently when memory permits."""
if len(samples) < 2 or len(samples) != len(save_file_paths):
return None
videos = []
for sample, save_file_path in zip(samples, save_file_paths):
video, _ = _split_sample_audio(sample)
if not (
isinstance(video, torch.Tensor)
and video.device.type == "cuda"
and video.dim() in (3, 4)
and int(video.shape[0]) == 3
and os.path.splitext(save_file_path)[1].lower() == ".mp4"
):
return None
if video.dim() == 3:
video = video.unsqueeze(1)
if videos and video.device != videos[0].device:
return None
videos.append(video)
# Each direct save creates one multiply result and two uint8 layouts for a
# temporal chunk. Keep two such chunks below 25% of currently free device
# memory so postprocessing cannot turn a tight inference into an OOM.
temporary_bytes = sorted(
(
_cuda_video_conversion_chunk_frames(video)
* 3
* int(video.shape[-2])
* int(video.shape[-1])
* (int(video.element_size()) + 2)
for video in videos
),
reverse=True,
)[:_MAX_PARALLEL_CUDA_VIDEO_SAVES]
try:
free_bytes, _ = torch.cuda.mem_get_info(videos[0].device)
except (RuntimeError, TypeError):
return None
if sum(temporary_bytes) > int(free_bytes) // 4:
return None
try:
available_cpus = len(os.sched_getaffinity(0))
except (AttributeError, OSError):
available_cpus = os.cpu_count() or 1
encoder_threads = sorted(
(_x264_auto_thread_count(int(video.shape[-2])) for video in videos),
reverse=True,
)[:_MAX_PARALLEL_CUDA_VIDEO_SAVES]
if sum(encoder_threads) > available_cpus:
return None
def save_one(idx: int) -> bool:
return _try_save_cuda_video_direct(
save_file_path=save_file_paths[idx],
sample=samples[idx],
fps=fps,
audio_sample_rate=audio_sample_rate,
output_compression=output_compression,
)
try:
with ThreadPoolExecutor(max_workers=_MAX_PARALLEL_CUDA_VIDEO_SAVES) as pool:
return list(pool.map(save_one, range(len(samples))))
except Exception as exc:
logger.warning(
"Parallel CUDA video save failed; falling back to serial output: %s",
str(exc),
)
return None
def _mux_audio_np_into_mp4( def _mux_audio_np_into_mp4(
*, *,
save_file_path: str, save_file_path: str,
@@ -1013,8 +1143,35 @@ def save_outputs(
upscaling_scale: int = 4, upscaling_scale: int = 4,
) -> list[str]: ) -> list[str]:
output_paths: list[str] = [] output_paths: list[str] = []
for idx, sample in enumerate(outputs): samples = (
save_file_path = build_output_path(idx) [
attach_audio_to_video_sample(sample, audio, idx)
for idx, sample in enumerate(outputs)
]
if data_type == DataType.VIDEO
else outputs
)
save_file_paths = [build_output_path(idx) for idx in range(len(outputs))]
parallel_results = None
if (
data_type == DataType.VIDEO
and len(outputs) > 1
and save_output
and frames_out is None
and not enable_frame_interpolation
and not enable_upscaling
):
for path in save_file_paths:
os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
parallel_results = _try_save_cuda_videos_direct(
samples,
save_file_paths,
fps=fps,
audio_sample_rate=audio_sample_rate,
output_compression=output_compression,
)
for idx, (sample, save_file_path) in enumerate(zip(samples, save_file_paths)):
if data_type == DataType.ACTION: if data_type == DataType.ACTION:
if samples_out is not None: if samples_out is not None:
samples_out.append(sample) samples_out.append(sample)
@@ -1031,7 +1188,6 @@ def save_outputs(
continue continue
if data_type == DataType.VIDEO: if data_type == DataType.VIDEO:
sample = attach_audio_to_video_sample(sample, audio, idx)
if ( if (
save_output save_output
and save_file_path and save_file_path
@@ -1040,13 +1196,16 @@ def save_outputs(
and not enable_upscaling and not enable_upscaling
): ):
os.makedirs(os.path.dirname(save_file_path) or ".", exist_ok=True) os.makedirs(os.path.dirname(save_file_path) or ".", exist_ok=True)
if _try_save_cuda_video_direct( direct_saved = parallel_results is not None and parallel_results[idx]
save_file_path=save_file_path, if not direct_saved:
sample=sample, direct_saved = _try_save_cuda_video_direct(
fps=fps, save_file_path=save_file_path,
audio_sample_rate=audio_sample_rate, sample=sample,
output_compression=output_compression, fps=fps,
): audio_sample_rate=audio_sample_rate,
output_compression=output_compression,
)
if direct_saved:
if samples_out is not None: if samples_out is not None:
samples_out.append(sample) samples_out.append(sample)
if audios_out is not None: if audios_out is not None:
@@ -6,6 +6,7 @@ from __future__ import annotations
import json import json
import math import math
import subprocess import subprocess
from concurrent.futures import ThreadPoolExecutor
from typing import TYPE_CHECKING, Any from typing import TYPE_CHECKING, Any
from sglang.multimodal_gen.configs.pipeline_configs.minimax_h3 import ( from sglang.multimodal_gen.configs.pipeline_configs.minimax_h3 import (
@@ -329,16 +330,22 @@ class MiniMaxH3VideoModelAdapter:
if shape.get("width") is not None and shape.get("height") is not None: if shape.get("width") is not None and shape.get("height") is not None:
expected_size = (int(shape["width"]), int(shape["height"])) expected_size = (int(shape["width"]), int(shape["height"]))
final_media_fields: dict[str, str] = {} def probe_output(output_path: str) -> dict[str, str]:
for output_index, output_path in enumerate(output_paths): return _probe_minimax_h3_output_fields(
media_fields = _probe_minimax_h3_output_fields(
output_path, output_path,
expected_frame_count=expected_frame_count, expected_frame_count=expected_frame_count,
expected_size=expected_size, expected_size=expected_size,
) )
if output_index == 0:
final_media_fields = media_fields if len(output_paths) > 1:
elif media_fields != final_media_fields: with ThreadPoolExecutor(max_workers=min(4, len(output_paths))) as pool:
media_fields_by_output = list(pool.map(probe_output, output_paths))
else:
media_fields_by_output = [probe_output(output_paths[0])]
final_media_fields = media_fields_by_output[0]
for output_index, media_fields in enumerate(media_fields_by_output[1:], 1):
if media_fields != final_media_fields:
raise RuntimeError( raise RuntimeError(
"generated MiniMax H3 outputs have inconsistent media metadata: " "generated MiniMax H3 outputs have inconsistent media metadata: "
f"output 0={final_media_fields}, output " f"output 0={final_media_fields}, output "
@@ -199,3 +199,47 @@ def test_video_direct_save_short_circuits_materialization(tmp_path, monkeypatch)
assert paths == [str(output_path)] assert paths == [str(output_path)]
assert len(direct_calls) == 1 assert len(direct_calls) == 1
def test_multiple_videos_use_parallel_direct_save_with_serial_fallback(
tmp_path, monkeypatch
):
outputs = [torch.zeros((3, 1, 2, 3)), torch.ones((3, 1, 2, 3))]
direct_calls = []
def parallel_save(samples, paths, **kwargs):
direct_calls.append((samples, paths, kwargs))
return [True, False]
serial_calls = []
monkeypatch.setattr(output_utils, "_try_save_cuda_videos_direct", parallel_save)
monkeypatch.setattr(
output_utils,
"_try_save_cuda_video_direct",
lambda **kwargs: serial_calls.append(kwargs) or True,
)
monkeypatch.setattr(
output_utils,
"post_process_sample",
lambda *_args, **_kwargs: pytest.fail(
"successful parallel direct saves should skip frame materialization"
),
)
paths = output_utils.save_outputs(
outputs,
DataType.VIDEO,
fps=24,
save_output=True,
build_output_path=lambda idx: str(tmp_path / f"sample_{idx}.mp4"),
)
assert paths == [str(tmp_path / "sample_0.mp4"), str(tmp_path / "sample_1.mp4")]
assert len(direct_calls) == 1
samples, save_paths, kwargs = direct_calls[0]
assert all(actual is expected for actual, expected in zip(samples, outputs))
assert save_paths == paths
assert kwargs["fps"] == 24
assert len(serial_calls) == 1
assert serial_calls[0]["save_file_path"] == paths[1]