[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 tempfile
import threading
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from copy import copy
from dataclasses import dataclass, field
@@ -48,6 +49,8 @@ from sglang.srt.observability.trace import TraceReqContext
logger = init_logger(__name__)
_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()
_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)
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(
*,
save_file_path: str,
@@ -472,8 +494,8 @@ def _try_save_cuda_video_direct(
audio_sample_rate: Optional[int],
output_compression: Optional[int],
) -> bool:
"""Save a CUDA RGB video through a registered memfd instead of a Python pipe."""
if not hasattr(os, "memfd_create") or not os.path.isdir("/proc/self/fd"):
"""Stream CUDA RGB chunks to ffmpeg through a registered memfd."""
if not hasattr(os, "memfd_create") or not hasattr(os, "sendfile"):
return False
sample_without_audio, audio = _split_sample_audio(sample)
@@ -492,9 +514,8 @@ def _try_save_cuda_video_direct(
if video.shape[0] != 3:
return False
frames = (video * 255).clamp_(0, 255).to(torch.uint8)
frames = frames.permute(1, 2, 3, 0).contiguous()
num_frames, height, width, _ = frames.shape
_, num_frames, height, width = video.shape
chunk_frames = _cuda_video_conversion_chunk_frames(video)
quality = output_compression / 10 if output_compression is not None else 5
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)
ffmpeg_exe = _resolve_ffmpeg_exe()
shape = tuple(frames.shape)
with _acquire_cuda_video_buffer(shape) as buffer:
assert buffer.tensor is not None
buffer.tensor.copy_(frames, non_blocking=True)
torch.cuda.current_stream(frames.device).synchronize()
del frames
os.lseek(buffer.fd, 0, os.SEEK_SET)
command = [
ffmpeg_exe,
"-y",
"-f",
"rawvideo",
"-vcodec",
"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 = [
ffmpeg_exe,
"-y",
"-f",
"rawvideo",
"-vcodec",
"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,
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]
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
except Exception as e:
logger.warning(
@@ -603,6 +653,86 @@ def _try_save_cuda_video_direct(
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(
*,
save_file_path: str,
@@ -1013,8 +1143,35 @@ def save_outputs(
upscaling_scale: int = 4,
) -> list[str]:
output_paths: list[str] = []
for idx, sample in enumerate(outputs):
save_file_path = build_output_path(idx)
samples = (
[
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 samples_out is not None:
samples_out.append(sample)
@@ -1031,7 +1188,6 @@ def save_outputs(
continue
if data_type == DataType.VIDEO:
sample = attach_audio_to_video_sample(sample, audio, idx)
if (
save_output
and save_file_path
@@ -1040,13 +1196,16 @@ def save_outputs(
and not enable_upscaling
):
os.makedirs(os.path.dirname(save_file_path) or ".", exist_ok=True)
if _try_save_cuda_video_direct(
save_file_path=save_file_path,
sample=sample,
fps=fps,
audio_sample_rate=audio_sample_rate,
output_compression=output_compression,
):
direct_saved = parallel_results is not None and parallel_results[idx]
if not direct_saved:
direct_saved = _try_save_cuda_video_direct(
save_file_path=save_file_path,
sample=sample,
fps=fps,
audio_sample_rate=audio_sample_rate,
output_compression=output_compression,
)
if direct_saved:
if samples_out is not None:
samples_out.append(sample)
if audios_out is not None:
@@ -6,6 +6,7 @@ from __future__ import annotations
import json
import math
import subprocess
from concurrent.futures import ThreadPoolExecutor
from typing import TYPE_CHECKING, Any
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:
expected_size = (int(shape["width"]), int(shape["height"]))
final_media_fields: dict[str, str] = {}
for output_index, output_path in enumerate(output_paths):
media_fields = _probe_minimax_h3_output_fields(
def probe_output(output_path: str) -> dict[str, str]:
return _probe_minimax_h3_output_fields(
output_path,
expected_frame_count=expected_frame_count,
expected_size=expected_size,
)
if output_index == 0:
final_media_fields = media_fields
elif media_fields != final_media_fields:
if len(output_paths) > 1:
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(
"generated MiniMax H3 outputs have inconsistent media metadata: "
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 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]