[diffusion] optimize: stream and parallelize bit-exact video output saves (#34564)
This commit is contained in:
@@ -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,14 +539,6 @@ 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)
|
|
||||||
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 = [
|
command = [
|
||||||
ffmpeg_exe,
|
ffmpeg_exe,
|
||||||
"-y",
|
"-y",
|
||||||
@@ -540,7 +553,7 @@ def _try_save_cuda_video_direct(
|
|||||||
"-r",
|
"-r",
|
||||||
f"{fps:.02f}",
|
f"{fps:.02f}",
|
||||||
"-i",
|
"-i",
|
||||||
f"/proc/self/fd/{buffer.fd}",
|
"pipe:0",
|
||||||
]
|
]
|
||||||
if tmp_wav_path is None:
|
if tmp_wav_path is None:
|
||||||
command += ["-an"]
|
command += ["-an"]
|
||||||
@@ -581,12 +594,49 @@ def _try_save_cuda_video_direct(
|
|||||||
]
|
]
|
||||||
command += ["-v", "warning", save_file_path]
|
command += ["-v", "warning", save_file_path]
|
||||||
|
|
||||||
subprocess.run(
|
with tempfile.TemporaryFile() as stderr_file:
|
||||||
|
process = subprocess.Popen(
|
||||||
command,
|
command,
|
||||||
check=True,
|
stdin=subprocess.PIPE,
|
||||||
pass_fds=(buffer.fd,),
|
|
||||||
stdout=subprocess.DEVNULL,
|
stdout=subprocess.DEVNULL,
|
||||||
stderr=subprocess.PIPE,
|
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:
|
||||||
@@ -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]
|
||||||
|
if not direct_saved:
|
||||||
|
direct_saved = _try_save_cuda_video_direct(
|
||||||
save_file_path=save_file_path,
|
save_file_path=save_file_path,
|
||||||
sample=sample,
|
sample=sample,
|
||||||
fps=fps,
|
fps=fps,
|
||||||
audio_sample_rate=audio_sample_rate,
|
audio_sample_rate=audio_sample_rate,
|
||||||
output_compression=output_compression,
|
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:
|
||||||
|
|||||||
+13
-6
@@ -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]
|
||||||
|
|||||||
Reference in New Issue
Block a user