[diffusion] feat: accelerate multiple-outputs generation (#23759)
This commit is contained in:
@@ -48,6 +48,20 @@ PATCH_SIZE = 16
|
|||||||
PATCH_AREA = PATCH_SIZE * PATCH_SIZE
|
PATCH_AREA = PATCH_SIZE * PATCH_SIZE
|
||||||
|
|
||||||
|
|
||||||
|
def _get_response_output_count(resp_json: Dict[str, Any]) -> int:
|
||||||
|
if isinstance(resp_json.get("num_outputs"), int):
|
||||||
|
return resp_json["num_outputs"]
|
||||||
|
if isinstance(resp_json.get("data"), list):
|
||||||
|
return len(resp_json["data"])
|
||||||
|
if isinstance(resp_json.get("file_paths"), list):
|
||||||
|
return len(resp_json["file_paths"])
|
||||||
|
if isinstance(resp_json.get("urls"), list):
|
||||||
|
return len(resp_json["urls"])
|
||||||
|
if resp_json.get("file_path") or resp_json.get("url"):
|
||||||
|
return 1
|
||||||
|
return 0
|
||||||
|
|
||||||
|
|
||||||
def _compute_scale_factor(req: RequestFuncInput, args) -> Optional[float]:
|
def _compute_scale_factor(req: RequestFuncInput, args) -> Optional[float]:
|
||||||
"""Computes the composite scale factor (area × frames × steps) for a request."""
|
"""Computes the composite scale factor (area × frames × steps) for a request."""
|
||||||
width = req.width or args.width
|
width = req.width or args.width
|
||||||
@@ -142,6 +156,7 @@ async def async_request_image_sglang(
|
|||||||
data.add_field("model", input.model)
|
data.add_field("model", input.model)
|
||||||
data.add_field("prompt", input.prompt)
|
data.add_field("prompt", input.prompt)
|
||||||
data.add_field("response_format", "b64_json")
|
data.add_field("response_format", "b64_json")
|
||||||
|
data.add_field("n", str(input.num_outputs_per_prompt))
|
||||||
|
|
||||||
if input.width and input.height:
|
if input.width and input.height:
|
||||||
data.add_field("size", f"{input.width}x{input.height}")
|
data.add_field("size", f"{input.width}x{input.height}")
|
||||||
@@ -172,6 +187,7 @@ async def async_request_image_sglang(
|
|||||||
resp_json = await response.json()
|
resp_json = await response.json()
|
||||||
output.response_body = resp_json
|
output.response_body = resp_json
|
||||||
output.success = True
|
output.success = True
|
||||||
|
output.output_count = _get_response_output_count(resp_json)
|
||||||
if "peak_memory_mb" in resp_json:
|
if "peak_memory_mb" in resp_json:
|
||||||
output.peak_memory_mb = resp_json["peak_memory_mb"]
|
output.peak_memory_mb = resp_json["peak_memory_mb"]
|
||||||
else:
|
else:
|
||||||
@@ -185,7 +201,7 @@ async def async_request_image_sglang(
|
|||||||
payload = {
|
payload = {
|
||||||
"model": input.model,
|
"model": input.model,
|
||||||
"prompt": input.prompt,
|
"prompt": input.prompt,
|
||||||
"n": 1,
|
"n": input.num_outputs_per_prompt,
|
||||||
"response_format": "b64_json",
|
"response_format": "b64_json",
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -203,6 +219,7 @@ async def async_request_image_sglang(
|
|||||||
resp_json = await response.json()
|
resp_json = await response.json()
|
||||||
output.response_body = resp_json
|
output.response_body = resp_json
|
||||||
output.success = True
|
output.success = True
|
||||||
|
output.output_count = _get_response_output_count(resp_json)
|
||||||
if "peak_memory_mb" in resp_json:
|
if "peak_memory_mb" in resp_json:
|
||||||
output.peak_memory_mb = resp_json["peak_memory_mb"]
|
output.peak_memory_mb = resp_json["peak_memory_mb"]
|
||||||
else:
|
else:
|
||||||
@@ -239,6 +256,7 @@ async def async_request_video_sglang(
|
|||||||
data = aiohttp.FormData()
|
data = aiohttp.FormData()
|
||||||
data.add_field("model", input.model)
|
data.add_field("model", input.model)
|
||||||
data.add_field("prompt", input.prompt)
|
data.add_field("prompt", input.prompt)
|
||||||
|
data.add_field("num_outputs_per_prompt", str(input.num_outputs_per_prompt))
|
||||||
|
|
||||||
if input.width and input.height:
|
if input.width and input.height:
|
||||||
data.add_field("size", f"{input.width}x{input.height}")
|
data.add_field("size", f"{input.width}x{input.height}")
|
||||||
@@ -296,6 +314,7 @@ async def async_request_video_sglang(
|
|||||||
payload: Dict[str, Any] = {
|
payload: Dict[str, Any] = {
|
||||||
"model": input.model,
|
"model": input.model,
|
||||||
"prompt": input.prompt,
|
"prompt": input.prompt,
|
||||||
|
"num_outputs_per_prompt": input.num_outputs_per_prompt,
|
||||||
}
|
}
|
||||||
if input.width and input.height:
|
if input.width and input.height:
|
||||||
payload["size"] = f"{input.width}x{input.height}"
|
payload["size"] = f"{input.width}x{input.height}"
|
||||||
@@ -350,6 +369,7 @@ async def async_request_video_sglang(
|
|||||||
if status == "completed":
|
if status == "completed":
|
||||||
output.success = True
|
output.success = True
|
||||||
output.response_body = status_data
|
output.response_body = status_data
|
||||||
|
output.output_count = _get_response_output_count(status_data)
|
||||||
if "peak_memory_mb" in status_data:
|
if "peak_memory_mb" in status_data:
|
||||||
output.peak_memory_mb = status_data["peak_memory_mb"]
|
output.peak_memory_mb = status_data["peak_memory_mb"]
|
||||||
break
|
break
|
||||||
@@ -394,17 +414,29 @@ def calculate_metrics(
|
|||||||
|
|
||||||
num_success = len(success_outputs)
|
num_success = len(success_outputs)
|
||||||
latencies = [o.latency for o in success_outputs]
|
latencies = [o.latency for o in success_outputs]
|
||||||
peak_memories = [o.peak_memory_mb for o in success_outputs if o.peak_memory_mb > 0]
|
completed_outputs = sum(o.output_count for o in success_outputs)
|
||||||
|
peak_memories = [
|
||||||
|
o.peak_memory_mb
|
||||||
|
for o in success_outputs
|
||||||
|
if o.peak_memory_mb is not None and o.peak_memory_mb > 0
|
||||||
|
]
|
||||||
|
|
||||||
metrics = {
|
metrics = {
|
||||||
"duration": total_duration,
|
"duration": total_duration,
|
||||||
"completed_requests": num_success,
|
"completed_requests": num_success,
|
||||||
|
"completed_outputs": completed_outputs,
|
||||||
"failed_requests": len(error_outputs),
|
"failed_requests": len(error_outputs),
|
||||||
"throughput_qps": num_success / total_duration if total_duration > 0 else 0,
|
"throughput_qps": num_success / total_duration if total_duration > 0 else 0,
|
||||||
|
"output_throughput_ops": (
|
||||||
|
completed_outputs / total_duration if total_duration > 0 else 0
|
||||||
|
),
|
||||||
"latency_mean": np.mean(latencies) if latencies else 0,
|
"latency_mean": np.mean(latencies) if latencies else 0,
|
||||||
"latency_median": np.median(latencies) if latencies else 0,
|
"latency_median": np.median(latencies) if latencies else 0,
|
||||||
"latency_p99": np.percentile(latencies, 99) if latencies else 0,
|
|
||||||
"latency_p50": np.percentile(latencies, 50) if latencies else 0,
|
"latency_p50": np.percentile(latencies, 50) if latencies else 0,
|
||||||
|
"latency_p90": np.percentile(latencies, 90) if latencies else 0,
|
||||||
|
"latency_p95": np.percentile(latencies, 95) if latencies else 0,
|
||||||
|
"latency_p99": np.percentile(latencies, 99) if latencies else 0,
|
||||||
|
"num_outputs_per_prompt": args.num_outputs_per_prompt,
|
||||||
"peak_memory_mb_max": max(peak_memories) if peak_memories else 0,
|
"peak_memory_mb_max": max(peak_memories) if peak_memories else 0,
|
||||||
"peak_memory_mb_mean": np.mean(peak_memories) if peak_memories else 0,
|
"peak_memory_mb_mean": np.mean(peak_memories) if peak_memories else 0,
|
||||||
"peak_memory_mb_median": np.median(peak_memories) if peak_memories else 0,
|
"peak_memory_mb_median": np.median(peak_memories) if peak_memories else 0,
|
||||||
@@ -623,14 +655,21 @@ async def benchmark(args):
|
|||||||
"Successful requests:",
|
"Successful requests:",
|
||||||
f"{metrics['completed_requests']}/{len(requests_list)}",
|
f"{metrics['completed_requests']}/{len(requests_list)}",
|
||||||
)
|
)
|
||||||
|
print_value_formatted("Completed outputs:", metrics["completed_outputs"])
|
||||||
|
print_value_formatted("Outputs per prompt:", metrics["num_outputs_per_prompt"])
|
||||||
|
|
||||||
# Section 3: Performance Metrics
|
# Section 3: Performance Metrics
|
||||||
print_divider(50)
|
print_divider(50)
|
||||||
|
|
||||||
print_value_formatted("Request throughput (req/s):", metrics["throughput_qps"])
|
print_value_formatted("Request throughput (req/s):", metrics["throughput_qps"])
|
||||||
|
print_value_formatted(
|
||||||
|
"Output throughput (outputs/s):", metrics["output_throughput_ops"]
|
||||||
|
)
|
||||||
|
|
||||||
print_value_formatted("Latency Mean (s):", metrics["latency_mean"])
|
print_value_formatted("Latency Mean (s):", metrics["latency_mean"])
|
||||||
print_value_formatted("Latency Median (s):", metrics["latency_median"])
|
print_value_formatted("Latency Median (s):", metrics["latency_median"])
|
||||||
|
print_value_formatted("Latency P90 (s):", metrics["latency_p90"])
|
||||||
|
print_value_formatted("Latency P95 (s):", metrics["latency_p95"])
|
||||||
print_value_formatted("Latency P99 (s):", metrics["latency_p99"])
|
print_value_formatted("Latency P99 (s):", metrics["latency_p99"])
|
||||||
|
|
||||||
if metrics["peak_memory_mb_max"] > 0:
|
if metrics["peak_memory_mb_max"] > 0:
|
||||||
@@ -707,6 +746,12 @@ if __name__ == "__main__":
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--num-prompts", type=int, default=10, help="Number of prompts to benchmark."
|
"--num-prompts", type=int, default=10, help="Number of prompts to benchmark."
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--num-outputs-per-prompt",
|
||||||
|
type=int,
|
||||||
|
default=1,
|
||||||
|
help="Number of generated outputs requested per prompt.",
|
||||||
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--max-concurrency",
|
"--max-concurrency",
|
||||||
type=int,
|
type=int,
|
||||||
@@ -735,7 +780,8 @@ if __name__ == "__main__":
|
|||||||
default=None,
|
default=None,
|
||||||
help=(
|
help=(
|
||||||
"JSON string defining random request profiles. "
|
"JSON string defining random request profiles. "
|
||||||
"Each profile may contain: width, height, num_inference_steps, etc. "
|
"Each profile may contain: width, height, num_inference_steps, "
|
||||||
|
"num_outputs_per_prompt, etc. "
|
||||||
"The 'weight' field controls sampling probability (relative weight). "
|
"The 'weight' field controls sampling probability (relative weight). "
|
||||||
"Example: "
|
"Example: "
|
||||||
'[{"width":512,"height":512,"num_inference_steps":20,"weight":0.15},'
|
'[{"width":512,"height":512,"num_inference_steps":20,"weight":0.15},'
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ class RequestFuncInput:
|
|||||||
prompt: str
|
prompt: str
|
||||||
api_url: str = ""
|
api_url: str = ""
|
||||||
model: str = ""
|
model: str = ""
|
||||||
|
num_outputs_per_prompt: int = 1
|
||||||
width: Optional[int] = None
|
width: Optional[int] = None
|
||||||
height: Optional[int] = None
|
height: Optional[int] = None
|
||||||
num_frames: Optional[int] = None
|
num_frames: Optional[int] = None
|
||||||
@@ -42,6 +43,7 @@ class RequestFuncOutput:
|
|||||||
response_body: Dict[str, Any] = field(default_factory=dict)
|
response_body: Dict[str, Any] = field(default_factory=dict)
|
||||||
peak_memory_mb: float = 0.0
|
peak_memory_mb: float = 0.0
|
||||||
slo_achieved: Optional[bool] = None
|
slo_achieved: Optional[bool] = None
|
||||||
|
output_count: int = 0
|
||||||
|
|
||||||
|
|
||||||
def is_dir_not_empty(path: str) -> bool:
|
def is_dir_not_empty(path: str) -> bool:
|
||||||
@@ -274,6 +276,7 @@ class VBenchDataset(BaseDataset):
|
|||||||
prompt=item.get("prompt", ""),
|
prompt=item.get("prompt", ""),
|
||||||
api_url=self.api_url,
|
api_url=self.api_url,
|
||||||
model=self.model,
|
model=self.model,
|
||||||
|
num_outputs_per_prompt=self.args.num_outputs_per_prompt,
|
||||||
width=self.args.width,
|
width=self.args.width,
|
||||||
height=self.args.height,
|
height=self.args.height,
|
||||||
num_frames=self.args.num_frames,
|
num_frames=self.args.num_frames,
|
||||||
@@ -315,6 +318,9 @@ class RandomDataset(BaseDataset):
|
|||||||
prompt=f"Random prompt {idx} for benchmarking diffusion models",
|
prompt=f"Random prompt {idx} for benchmarking diffusion models",
|
||||||
api_url=self.api_url,
|
api_url=self.api_url,
|
||||||
model=self.model,
|
model=self.model,
|
||||||
|
num_outputs_per_prompt=profile.get(
|
||||||
|
"num_outputs_per_prompt", self.args.num_outputs_per_prompt
|
||||||
|
),
|
||||||
width=profile.get("width", self.args.width),
|
width=profile.get("width", self.args.width),
|
||||||
height=profile.get("height", self.args.height),
|
height=profile.get("height", self.args.height),
|
||||||
num_frames=profile.get("num_frames", self.args.num_frames),
|
num_frames=profile.get("num_frames", self.args.num_frames),
|
||||||
|
|||||||
@@ -128,7 +128,7 @@ class SamplingParams:
|
|||||||
|
|
||||||
# Batch info
|
# Batch info
|
||||||
num_outputs_per_prompt: int = 1
|
num_outputs_per_prompt: int = 1
|
||||||
seed: int = 42
|
seed: int | list[int] = 42
|
||||||
generator_device: str | None = None # None means use the pipeline/model default
|
generator_device: str | None = None # None means use the pipeline/model default
|
||||||
|
|
||||||
# Original dimensions (before VAE scaling)
|
# Original dimensions (before VAE scaling)
|
||||||
@@ -313,6 +313,23 @@ class SamplingParams:
|
|||||||
f"num_outputs_per_prompt must be a positive int, got {self.num_outputs_per_prompt!r}"
|
f"num_outputs_per_prompt must be a positive int, got {self.num_outputs_per_prompt!r}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if isinstance(self.seed, list):
|
||||||
|
if not self.seed:
|
||||||
|
raise ValueError("seed list must not be empty")
|
||||||
|
for seed in self.seed:
|
||||||
|
if isinstance(seed, bool) or not isinstance(seed, int) or seed < 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"seed list must contain non-negative ints, got {self.seed!r}"
|
||||||
|
)
|
||||||
|
elif (
|
||||||
|
isinstance(self.seed, bool)
|
||||||
|
or not isinstance(self.seed, int)
|
||||||
|
or self.seed < 0
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"seed must be a non-negative int or list of ints, " f"got {self.seed!r}"
|
||||||
|
)
|
||||||
|
|
||||||
# Used by seconds() and video writer; fps <= 0 is always invalid.
|
# Used by seconds() and video writer; fps <= 0 is always invalid.
|
||||||
if not isinstance(self.fps, int) or self.fps <= 0:
|
if not isinstance(self.fps, int) or self.fps <= 0:
|
||||||
raise ValueError(f"fps must be a positive int, got {self.fps!r}")
|
raise ValueError(f"fps must be a positive int, got {self.fps!r}")
|
||||||
@@ -728,6 +745,7 @@ class SamplingParams:
|
|||||||
add_argument(
|
add_argument(
|
||||||
"--seed",
|
"--seed",
|
||||||
type=int,
|
type=int,
|
||||||
|
nargs="+",
|
||||||
help="Random seed for generation",
|
help="Random seed for generation",
|
||||||
)
|
)
|
||||||
add_argument(
|
add_argument(
|
||||||
@@ -961,11 +979,14 @@ class SamplingParams:
|
|||||||
sampling_params_fields = {attr.name for attr in dataclasses.fields(cls)}
|
sampling_params_fields = {attr.name for attr in dataclasses.fields(cls)}
|
||||||
args_attrs = set(vars(args).keys())
|
args_attrs = set(vars(args).keys())
|
||||||
attrs = sampling_params_fields & args_attrs
|
attrs = sampling_params_fields & args_attrs
|
||||||
return {
|
cli_args = {
|
||||||
attr: getattr(args, attr)
|
attr: getattr(args, attr)
|
||||||
for attr in attrs
|
for attr in attrs
|
||||||
if hasattr(args, attr) and getattr(args, attr) is not None
|
if hasattr(args, attr) and getattr(args, attr) is not None
|
||||||
}
|
}
|
||||||
|
if isinstance(cli_args.get("seed"), list) and len(cli_args["seed"]) == 1:
|
||||||
|
cli_args["seed"] = cli_args["seed"][0]
|
||||||
|
return cli_args
|
||||||
|
|
||||||
def output_file_path(self):
|
def output_file_path(self):
|
||||||
if self.output_path is None:
|
if self.output_path is None:
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import dataclasses
|
|||||||
import multiprocessing as mp
|
import multiprocessing as mp
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
|
from contextlib import ExitStack
|
||||||
from typing import Any, List, Union
|
from typing import Any, List, Union
|
||||||
|
|
||||||
from sglang.multimodal_gen.configs.sample.sampling_params import (
|
from sglang.multimodal_gen.configs.sample.sampling_params import (
|
||||||
@@ -25,6 +26,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
|||||||
SetLoraReq,
|
SetLoraReq,
|
||||||
ShutdownReq,
|
ShutdownReq,
|
||||||
UnmergeLoraWeightsReq,
|
UnmergeLoraWeightsReq,
|
||||||
|
expand_request_outputs,
|
||||||
format_lora_message,
|
format_lora_message,
|
||||||
prepare_request,
|
prepare_request,
|
||||||
save_outputs,
|
save_outputs,
|
||||||
@@ -208,7 +210,7 @@ class DiffGenerator:
|
|||||||
**sampling_params_kwargs,
|
**sampling_params_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
requests: list[Req] = []
|
request_groups: list[list[Req]] = []
|
||||||
image_paths_per_prompt = self._resolve_image_paths_per_prompt(
|
image_paths_per_prompt = self._resolve_image_paths_per_prompt(
|
||||||
prompts, sampling_params_orig.image_path
|
prompts, sampling_params_orig.image_path
|
||||||
)
|
)
|
||||||
@@ -226,19 +228,31 @@ class DiffGenerator:
|
|||||||
sampling_params=sampling_params,
|
sampling_params=sampling_params,
|
||||||
external_trace_header=external_trace_header,
|
external_trace_header=external_trace_header,
|
||||||
)
|
)
|
||||||
requests.append(req)
|
request_groups.append(
|
||||||
|
expand_request_outputs(
|
||||||
|
req,
|
||||||
|
num_prompts=len(prompts),
|
||||||
|
prompt_index=i,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
results: list[GenerationResult] = []
|
results: list[GenerationResult] = []
|
||||||
total_start_time = time.perf_counter()
|
total_start_time = time.perf_counter()
|
||||||
|
global_output_index = 0
|
||||||
|
|
||||||
# 2. send requests to scheduler one at a time
|
for requests in request_groups:
|
||||||
# TODO: send batch when supported
|
|
||||||
for request_idx, req in enumerate(requests):
|
|
||||||
try:
|
try:
|
||||||
with trace_req(req.trace_ctx), log_generation_timer(
|
timer_prompt = [req.prompt for req in requests]
|
||||||
logger, req.prompt, request_idx + 1, len(requests)
|
logger.info("Processing %d grouped request(s)", len(requests))
|
||||||
) as timer:
|
with ExitStack() as stack:
|
||||||
output_batch = self._send_to_scheduler_and_wait_for_response([req])
|
for req in requests:
|
||||||
|
stack.enter_context(trace_req(req.trace_ctx))
|
||||||
|
timer = stack.enter_context(
|
||||||
|
log_generation_timer(logger, timer_prompt)
|
||||||
|
)
|
||||||
|
output_batch = self._send_to_scheduler_and_wait_for_response(
|
||||||
|
requests
|
||||||
|
)
|
||||||
if output_batch.error:
|
if output_batch.error:
|
||||||
raise Exception(f"{output_batch.error}")
|
raise Exception(f"{output_batch.error}")
|
||||||
|
|
||||||
@@ -246,97 +260,93 @@ class DiffGenerator:
|
|||||||
output_batch.output is None
|
output_batch.output is None
|
||||||
and output_batch.output_file_paths is None
|
and output_batch.output_file_paths is None
|
||||||
):
|
):
|
||||||
logger.error(
|
logger.error("Received empty output from scheduler")
|
||||||
"Received empty output from scheduler for prompt %d",
|
|
||||||
request_idx + 1,
|
|
||||||
)
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
common = dict(
|
if requests[0].save_output and requests[0].return_file_paths_only:
|
||||||
prompt=req.prompt,
|
output_file_paths = output_batch.output_file_paths or []
|
||||||
size=(req.height, req.width, req.num_frames),
|
self._validate_output_count(
|
||||||
generation_time=timer.duration,
|
len(output_file_paths), len(requests)
|
||||||
peak_memory_mb=output_batch.peak_memory_mb,
|
)
|
||||||
metrics=(
|
for idx, path in enumerate(output_file_paths):
|
||||||
output_batch.metrics.to_dict()
|
req = requests[idx]
|
||||||
if output_batch.metrics
|
|
||||||
else {}
|
|
||||||
),
|
|
||||||
trajectory_latents=output_batch.trajectory_latents,
|
|
||||||
trajectory_timesteps=output_batch.trajectory_timesteps,
|
|
||||||
rollout_trajectory_data=output_batch.rollout_trajectory_data,
|
|
||||||
trajectory_decoded=output_batch.trajectory_decoded,
|
|
||||||
)
|
|
||||||
|
|
||||||
if req.save_output and req.return_file_paths_only:
|
|
||||||
for idx, path in enumerate(output_batch.output_file_paths):
|
|
||||||
results.append(
|
results.append(
|
||||||
GenerationResult(
|
GenerationResult(
|
||||||
**common,
|
**self._result_common(
|
||||||
prompt_index=idx,
|
req, output_batch, timer.duration, idx
|
||||||
|
),
|
||||||
|
prompt_index=global_output_index + idx,
|
||||||
output_file_path=path,
|
output_file_path=path,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
continue
|
elif requests[0].data_type == DataType.MESH:
|
||||||
|
output_file_paths = output_batch.output_file_paths or []
|
||||||
if req.data_type == DataType.MESH:
|
self._validate_output_count(
|
||||||
for output_idx, sample in enumerate(
|
len(output_file_paths), len(requests)
|
||||||
output_batch.output_file_paths
|
)
|
||||||
):
|
for idx, sample in enumerate(output_file_paths):
|
||||||
|
req = requests[idx]
|
||||||
results.append(
|
results.append(
|
||||||
GenerationResult(
|
GenerationResult(
|
||||||
**common,
|
**self._result_common(
|
||||||
prompt_index=output_idx,
|
req, output_batch, timer.duration, idx
|
||||||
|
),
|
||||||
|
prompt_index=global_output_index + idx,
|
||||||
output_file_path=sample,
|
output_file_path=sample,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
continue
|
else:
|
||||||
|
self._validate_output_count(
|
||||||
samples_out: list[Any] = []
|
len(output_batch.output), len(requests)
|
||||||
audios_out: list[Any] = []
|
|
||||||
frames_out: list[Any] = []
|
|
||||||
num_outputs = len(output_batch.output)
|
|
||||||
save_outputs(
|
|
||||||
output_batch.output,
|
|
||||||
req.data_type,
|
|
||||||
req.fps,
|
|
||||||
req.save_output,
|
|
||||||
lambda idx: req.output_file_path(num_outputs, idx),
|
|
||||||
audio=output_batch.audio,
|
|
||||||
audio_sample_rate=output_batch.audio_sample_rate,
|
|
||||||
samples_out=samples_out,
|
|
||||||
audios_out=audios_out,
|
|
||||||
frames_out=frames_out,
|
|
||||||
output_compression=req.output_compression,
|
|
||||||
enable_frame_interpolation=req.enable_frame_interpolation,
|
|
||||||
frame_interpolation_exp=req.frame_interpolation_exp,
|
|
||||||
frame_interpolation_scale=req.frame_interpolation_scale,
|
|
||||||
frame_interpolation_model_path=req.frame_interpolation_model_path,
|
|
||||||
enable_upscaling=req.enable_upscaling,
|
|
||||||
upscaling_model_path=req.upscaling_model_path,
|
|
||||||
upscaling_scale=req.upscaling_scale,
|
|
||||||
)
|
|
||||||
|
|
||||||
for idx in range(len(samples_out)):
|
|
||||||
results.append(
|
|
||||||
GenerationResult(
|
|
||||||
**common,
|
|
||||||
samples=samples_out[idx],
|
|
||||||
frames=frames_out[idx],
|
|
||||||
audio=audios_out[idx],
|
|
||||||
prompt_index=idx,
|
|
||||||
output_file_path=req.output_file_path(num_outputs, idx),
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
samples_out: list[Any] = []
|
||||||
|
audios_out: list[Any] = []
|
||||||
|
frames_out: list[Any] = []
|
||||||
|
save_outputs(
|
||||||
|
output_batch.output,
|
||||||
|
requests[0].data_type,
|
||||||
|
requests[0].fps,
|
||||||
|
requests[0].save_output,
|
||||||
|
lambda idx: requests[idx].output_file_path(1, 0),
|
||||||
|
audio=output_batch.audio,
|
||||||
|
audio_sample_rate=output_batch.audio_sample_rate,
|
||||||
|
samples_out=samples_out,
|
||||||
|
audios_out=audios_out,
|
||||||
|
frames_out=frames_out,
|
||||||
|
output_compression=requests[0].output_compression,
|
||||||
|
enable_frame_interpolation=requests[
|
||||||
|
0
|
||||||
|
].enable_frame_interpolation,
|
||||||
|
frame_interpolation_exp=requests[0].frame_interpolation_exp,
|
||||||
|
frame_interpolation_scale=requests[
|
||||||
|
0
|
||||||
|
].frame_interpolation_scale,
|
||||||
|
frame_interpolation_model_path=requests[
|
||||||
|
0
|
||||||
|
].frame_interpolation_model_path,
|
||||||
|
enable_upscaling=requests[0].enable_upscaling,
|
||||||
|
upscaling_model_path=requests[0].upscaling_model_path,
|
||||||
|
upscaling_scale=requests[0].upscaling_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
for idx in range(len(samples_out)):
|
||||||
|
req = requests[idx]
|
||||||
|
results.append(
|
||||||
|
GenerationResult(
|
||||||
|
**self._result_common(
|
||||||
|
req, output_batch, timer.duration, idx
|
||||||
|
),
|
||||||
|
samples=samples_out[idx],
|
||||||
|
frames=frames_out[idx],
|
||||||
|
audio=audios_out[idx],
|
||||||
|
prompt_index=global_output_index + idx,
|
||||||
|
output_file_path=req.output_file_path(1, 0),
|
||||||
|
)
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(
|
logger.error("Generation failed: %s", e, exc_info=True)
|
||||||
"Generation failed for prompt %d/%d: %s",
|
finally:
|
||||||
request_idx + 1,
|
global_output_index += len(requests)
|
||||||
len(requests),
|
|
||||||
e,
|
|
||||||
exc_info=True,
|
|
||||||
)
|
|
||||||
continue
|
|
||||||
|
|
||||||
total_gen_time = time.perf_counter() - total_start_time
|
total_gen_time = time.perf_counter() - total_start_time
|
||||||
log_batch_completion(logger, len(results), total_gen_time)
|
log_batch_completion(logger, len(results), total_gen_time)
|
||||||
@@ -386,6 +396,39 @@ class DiffGenerator:
|
|||||||
f"Avg peak: {sum(peak_memories) / len(peak_memories):.2f} MB"
|
f"Avg peak: {sum(peak_memories) / len(peak_memories):.2f} MB"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _result_common(
|
||||||
|
req: Req,
|
||||||
|
output_batch: OutputBatch,
|
||||||
|
generation_time: float,
|
||||||
|
output_index: int | None = None,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
metrics = output_batch.metrics
|
||||||
|
if (
|
||||||
|
output_index is not None
|
||||||
|
and output_batch.metrics_list is not None
|
||||||
|
and output_index < len(output_batch.metrics_list)
|
||||||
|
):
|
||||||
|
metrics = output_batch.metrics_list[output_index]
|
||||||
|
return dict(
|
||||||
|
prompt=req.prompt,
|
||||||
|
size=(req.height, req.width, req.num_frames),
|
||||||
|
generation_time=generation_time,
|
||||||
|
peak_memory_mb=output_batch.peak_memory_mb,
|
||||||
|
metrics=metrics.to_dict() if metrics else {},
|
||||||
|
trajectory_latents=output_batch.trajectory_latents,
|
||||||
|
trajectory_timesteps=output_batch.trajectory_timesteps,
|
||||||
|
rollout_trajectory_data=output_batch.rollout_trajectory_data,
|
||||||
|
trajectory_decoded=output_batch.trajectory_decoded,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _validate_output_count(output_count: int, request_count: int) -> None:
|
||||||
|
if output_count != request_count:
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Expected {request_count} outputs, got {output_count} from scheduler"
|
||||||
|
)
|
||||||
|
|
||||||
def _send_to_scheduler_and_wait_for_response(self, batch: list[Req]) -> OutputBatch:
|
def _send_to_scheduler_and_wait_for_response(self, batch: list[Req]) -> OutputBatch:
|
||||||
"""
|
"""
|
||||||
Sends a request to the scheduler and waits for a response.
|
Sends a request to the scheduler and waits for a response.
|
||||||
|
|||||||
@@ -44,7 +44,7 @@ class ImageGenerationsRequest(BaseModel):
|
|||||||
true_cfg_scale: Optional[float] = (
|
true_cfg_scale: Optional[float] = (
|
||||||
None # for CFG vs guidance distillation (e.g., QwenImage)
|
None # for CFG vs guidance distillation (e.g., QwenImage)
|
||||||
)
|
)
|
||||||
seed: Optional[int] = 1024
|
seed: Optional[Union[int, List[int]]] = 1024
|
||||||
generator_device: Optional[str] = "cuda"
|
generator_device: Optional[str] = "cuda"
|
||||||
negative_prompt: Optional[str] = None
|
negative_prompt: Optional[str] = None
|
||||||
output_quality: Optional[str] = "default"
|
output_quality: Optional[str] = "default"
|
||||||
@@ -76,6 +76,8 @@ class VideoResponse(BaseModel):
|
|||||||
expires_at: Optional[int] = None
|
expires_at: Optional[int] = None
|
||||||
error: Optional[Dict[str, Any]] = None
|
error: Optional[Dict[str, Any]] = None
|
||||||
file_path: Optional[str] = None
|
file_path: Optional[str] = None
|
||||||
|
file_paths: Optional[List[str]] = None
|
||||||
|
num_outputs: Optional[int] = None
|
||||||
peak_memory_mb: Optional[float] = None
|
peak_memory_mb: Optional[float] = None
|
||||||
inference_time_s: Optional[float] = None
|
inference_time_s: Optional[float] = None
|
||||||
|
|
||||||
@@ -85,11 +87,13 @@ class VideoGenerationsRequest(BaseModel):
|
|||||||
input_reference: Optional[str] = None
|
input_reference: Optional[str] = None
|
||||||
reference_url: Optional[str] = None
|
reference_url: Optional[str] = None
|
||||||
model: Optional[str] = None
|
model: Optional[str] = None
|
||||||
|
n: Optional[int] = 1
|
||||||
|
num_outputs_per_prompt: Optional[int] = None
|
||||||
seconds: Optional[int] = 4
|
seconds: Optional[int] = 4
|
||||||
size: Optional[str] = ""
|
size: Optional[str] = ""
|
||||||
fps: Optional[int] = None
|
fps: Optional[int] = None
|
||||||
num_frames: Optional[int] = None
|
num_frames: Optional[int] = None
|
||||||
seed: Optional[int] = 1024
|
seed: Optional[Union[int, List[int]]] = 1024
|
||||||
generator_device: Optional[str] = "cuda"
|
generator_device: Optional[str] = "cuda"
|
||||||
# SGLang extensions
|
# SGLang extensions
|
||||||
width: Optional[int] = None
|
width: Optional[int] = None
|
||||||
@@ -151,7 +155,7 @@ class MeshGenerationsRequest(BaseModel):
|
|||||||
prompt: str = "generate 3d mesh"
|
prompt: str = "generate 3d mesh"
|
||||||
input_image: Optional[str] = None
|
input_image: Optional[str] = None
|
||||||
model: Optional[str] = None
|
model: Optional[str] = None
|
||||||
seed: Optional[int] = None
|
seed: Optional[Union[int, List[int]]] = None
|
||||||
generator_device: Optional[str] = "cuda"
|
generator_device: Optional[str] = "cuda"
|
||||||
num_inference_steps: Optional[int] = None
|
num_inference_steps: Optional[int] = None
|
||||||
guidance_scale: Optional[float] = None
|
guidance_scale: Optional[float] = None
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
|||||||
SetLoraReq,
|
SetLoraReq,
|
||||||
ShutdownReq,
|
ShutdownReq,
|
||||||
UnmergeLoraWeightsReq,
|
UnmergeLoraWeightsReq,
|
||||||
|
expand_request_outputs,
|
||||||
format_lora_message,
|
format_lora_message,
|
||||||
save_outputs,
|
save_outputs,
|
||||||
)
|
)
|
||||||
@@ -325,8 +326,9 @@ async def process_generation_batch(
|
|||||||
batch,
|
batch,
|
||||||
) -> tuple[list[str], OutputBatch]:
|
) -> tuple[list[str], OutputBatch]:
|
||||||
total_start_time = time.perf_counter()
|
total_start_time = time.perf_counter()
|
||||||
|
requests = expand_request_outputs(batch)
|
||||||
with trace_req(batch.trace_ctx), log_generation_timer(logger, batch.prompt):
|
with trace_req(batch.trace_ctx), log_generation_timer(logger, batch.prompt):
|
||||||
result = await scheduler_client.forward([batch])
|
result = await scheduler_client.forward(requests)
|
||||||
|
|
||||||
if result.output is None and result.output_file_paths 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"
|
||||||
@@ -336,14 +338,22 @@ async def process_generation_batch(
|
|||||||
|
|
||||||
if result.output_file_paths:
|
if result.output_file_paths:
|
||||||
save_file_path_list = result.output_file_paths
|
save_file_path_list = result.output_file_paths
|
||||||
|
if len(save_file_path_list) < len(requests):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Expected at least {len(requests)} output paths, "
|
||||||
|
f"got {len(save_file_path_list)}"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
num_outputs = len(result.output)
|
if len(result.output) != len(requests):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Expected {len(requests)} outputs, got {len(result.output)}"
|
||||||
|
)
|
||||||
save_file_path_list = save_outputs(
|
save_file_path_list = save_outputs(
|
||||||
result.output,
|
result.output,
|
||||||
batch.data_type,
|
batch.data_type,
|
||||||
batch.fps,
|
batch.fps,
|
||||||
batch.save_output,
|
batch.save_output,
|
||||||
lambda idx: str(batch.output_file_path(num_outputs, idx)),
|
lambda idx: str(requests[idx].output_file_path(1, 0)),
|
||||||
audio=result.audio,
|
audio=result.audio,
|
||||||
audio_sample_rate=result.audio_sample_rate,
|
audio_sample_rate=result.audio_sample_rate,
|
||||||
output_compression=batch.output_compression,
|
output_compression=batch.output_compression,
|
||||||
@@ -357,7 +367,7 @@ async def process_generation_batch(
|
|||||||
)
|
)
|
||||||
|
|
||||||
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, len(save_file_path_list), total_time)
|
||||||
|
|
||||||
if result.peak_memory_mb and result.peak_memory_mb > 0:
|
if result.peak_memory_mb and result.peak_memory_mb > 0:
|
||||||
logger.info(f"Peak memory usage: {result.peak_memory_mb:.2f} MB")
|
logger.info(f"Peak memory usage: {result.peak_memory_mb:.2f} MB")
|
||||||
|
|||||||
@@ -56,10 +56,14 @@ def _build_video_sampling_params(request_id: str, request: VideoGenerationsReque
|
|||||||
seconds = request.seconds if request.seconds is not None else DEFAULT_VIDEO_SECONDS
|
seconds = request.seconds if request.seconds is not None else DEFAULT_VIDEO_SECONDS
|
||||||
fps = request.fps if request.fps is not None else DEFAULT_FPS
|
fps = request.fps if request.fps is not None else DEFAULT_FPS
|
||||||
num_frames = request.num_frames if request.num_frames is not None else fps * seconds
|
num_frames = request.num_frames if request.num_frames is not None else fps * seconds
|
||||||
|
num_outputs = request.num_outputs_per_prompt
|
||||||
|
if num_outputs is None:
|
||||||
|
num_outputs = request.n or 1
|
||||||
|
|
||||||
return build_sampling_params(
|
return build_sampling_params(
|
||||||
request_id,
|
request_id,
|
||||||
prompt=request.prompt,
|
prompt=request.prompt,
|
||||||
|
num_outputs_per_prompt=max(1, min(int(num_outputs), 10)),
|
||||||
size=request.size,
|
size=request.size,
|
||||||
width=request.width,
|
width=request.width,
|
||||||
height=request.height,
|
height=request.height,
|
||||||
@@ -156,6 +160,12 @@ async def _dispatch_job_async(
|
|||||||
"completed_at": int(time.time()),
|
"completed_at": int(time.time()),
|
||||||
"url": cloud_url,
|
"url": cloud_url,
|
||||||
"file_path": persistent_path,
|
"file_path": persistent_path,
|
||||||
|
"file_paths": (
|
||||||
|
[os.path.abspath(path) for path in save_file_path_list]
|
||||||
|
if output_persistent
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
"num_outputs": len(save_file_path_list),
|
||||||
}
|
}
|
||||||
update_fields = add_common_data_to_response(
|
update_fields = add_common_data_to_response(
|
||||||
update_fields, request_id=job_id, result=result
|
update_fields, request_id=job_id, result=result
|
||||||
@@ -180,6 +190,8 @@ async def create_video(
|
|||||||
input_reference: Optional[UploadFile] = File(None),
|
input_reference: Optional[UploadFile] = File(None),
|
||||||
reference_url: Optional[str] = Form(None),
|
reference_url: Optional[str] = Form(None),
|
||||||
model: Optional[str] = Form(None),
|
model: Optional[str] = Form(None),
|
||||||
|
n: Optional[int] = Form(1),
|
||||||
|
num_outputs_per_prompt: Optional[int] = Form(None),
|
||||||
seconds: Optional[int] = Form(None),
|
seconds: Optional[int] = Form(None),
|
||||||
size: Optional[str] = Form(None),
|
size: Optional[str] = Form(None),
|
||||||
fps: Optional[int] = Form(None),
|
fps: Optional[int] = Form(None),
|
||||||
@@ -262,6 +274,8 @@ async def create_video(
|
|||||||
prompt=prompt,
|
prompt=prompt,
|
||||||
input_reference=input_path,
|
input_reference=input_path,
|
||||||
model=model,
|
model=model,
|
||||||
|
n=n,
|
||||||
|
num_outputs_per_prompt=num_outputs_per_prompt,
|
||||||
seconds=seconds if seconds is not None else 4,
|
seconds=seconds if seconds is not None else 4,
|
||||||
size=size,
|
size=size,
|
||||||
fps=fps_val,
|
fps=fps_val,
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ import os
|
|||||||
import shutil
|
import shutil
|
||||||
import subprocess
|
import subprocess
|
||||||
import tempfile
|
import tempfile
|
||||||
|
from copy import copy
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, Callable, List, Optional, Sequence, Union
|
from typing import Any, Callable, List, Optional, Sequence, Union
|
||||||
|
|
||||||
@@ -122,6 +123,124 @@ class GenerationResult:
|
|||||||
output_file_path: str | None = None
|
output_file_path: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def normalize_output_seeds(
|
||||||
|
seed: int | list[int],
|
||||||
|
*,
|
||||||
|
num_outputs_per_prompt: int,
|
||||||
|
num_prompts: int = 1,
|
||||||
|
prompt_index: int = 0,
|
||||||
|
) -> list[int]:
|
||||||
|
"""
|
||||||
|
return a list of seed with size equal to `num_outputs_per_prompt`
|
||||||
|
"""
|
||||||
|
if num_outputs_per_prompt <= 0:
|
||||||
|
raise ValueError(
|
||||||
|
f"num_outputs_per_prompt must be positive, got {num_outputs_per_prompt}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if isinstance(seed, list):
|
||||||
|
seeds = [int(item) for item in seed]
|
||||||
|
total_outputs = num_outputs_per_prompt * num_prompts
|
||||||
|
if len(seeds) == num_outputs_per_prompt:
|
||||||
|
return seeds
|
||||||
|
if len(seeds) == total_outputs:
|
||||||
|
start = prompt_index * num_outputs_per_prompt
|
||||||
|
return seeds[start : start + num_outputs_per_prompt]
|
||||||
|
raise ValueError(
|
||||||
|
"seed list length must match num_outputs_per_prompt "
|
||||||
|
f"({num_outputs_per_prompt}) or total outputs ({total_outputs}), "
|
||||||
|
f"got {len(seeds)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
base_seed = int(seed)
|
||||||
|
return [base_seed + i for i in range(num_outputs_per_prompt)]
|
||||||
|
|
||||||
|
|
||||||
|
def _with_output_index_suffix(output_file_name: str, output_index: int) -> str:
|
||||||
|
base, ext = os.path.splitext(output_file_name)
|
||||||
|
return f"{base}_{output_index}{ext}"
|
||||||
|
|
||||||
|
|
||||||
|
def _copy_trace_ctx_for_output(req: Req, request_id: str | None, output_index: int):
|
||||||
|
trace_ctx = req.trace_ctx
|
||||||
|
if output_index == 0 or not trace_ctx.tracing_enable:
|
||||||
|
return trace_ctx
|
||||||
|
|
||||||
|
output_trace_ctx = TraceReqContext(
|
||||||
|
rid=request_id,
|
||||||
|
module_name=trace_ctx.module_name,
|
||||||
|
external_trace_header=trace_ctx.external_trace_header,
|
||||||
|
)
|
||||||
|
output_trace_ctx.trace_req_start()
|
||||||
|
return output_trace_ctx
|
||||||
|
|
||||||
|
|
||||||
|
def _copy_req_for_output(
|
||||||
|
req: Req,
|
||||||
|
*,
|
||||||
|
request_id: str | None,
|
||||||
|
output_index: int,
|
||||||
|
) -> Req:
|
||||||
|
"""Create a lightweight per-output ``Req`` without deep-copying tensors."""
|
||||||
|
output_req = copy(req)
|
||||||
|
output_req.sampling_params = copy(req.sampling_params)
|
||||||
|
output_req.extra = dict(req.extra)
|
||||||
|
output_req.trace_ctx = _copy_trace_ctx_for_output(req, request_id, output_index)
|
||||||
|
return output_req
|
||||||
|
|
||||||
|
|
||||||
|
def expand_request_outputs(
|
||||||
|
req: Req,
|
||||||
|
*,
|
||||||
|
num_prompts: int = 1,
|
||||||
|
prompt_index: int = 0,
|
||||||
|
) -> list[Req]:
|
||||||
|
"""
|
||||||
|
Expand a req to a list with size equal to `num_prompts`
|
||||||
|
"""
|
||||||
|
num_outputs = int(req.num_outputs_per_prompt)
|
||||||
|
# each req must has different seed
|
||||||
|
seeds = normalize_output_seeds(
|
||||||
|
req.seed,
|
||||||
|
num_outputs_per_prompt=num_outputs,
|
||||||
|
num_prompts=num_prompts,
|
||||||
|
prompt_index=prompt_index,
|
||||||
|
)
|
||||||
|
|
||||||
|
if num_outputs == 1:
|
||||||
|
req.seed = seeds[0]
|
||||||
|
req.seeds = None
|
||||||
|
req.generator = None
|
||||||
|
return [req]
|
||||||
|
|
||||||
|
expanded: list[Req] = []
|
||||||
|
for output_index, seed in enumerate(seeds):
|
||||||
|
output_request_id = (
|
||||||
|
f"{req.request_id}:{output_index}" if req.request_id is not None else None
|
||||||
|
)
|
||||||
|
output_req = _copy_req_for_output(
|
||||||
|
req, request_id=output_request_id, output_index=output_index
|
||||||
|
)
|
||||||
|
output_req.seed = seed
|
||||||
|
output_req.num_outputs_per_prompt = 1
|
||||||
|
output_req.seeds = None
|
||||||
|
output_req.generator = None
|
||||||
|
output_req.extra["parent_request_id"] = req.request_id
|
||||||
|
output_req.extra["output_index"] = output_index
|
||||||
|
|
||||||
|
if output_request_id is not None:
|
||||||
|
output_req.request_id = output_request_id
|
||||||
|
|
||||||
|
if req.output_file_name:
|
||||||
|
output_req.output_file_name = _with_output_index_suffix(
|
||||||
|
req.output_file_name, output_index
|
||||||
|
)
|
||||||
|
output_req.validate()
|
||||||
|
expanded.append(output_req)
|
||||||
|
|
||||||
|
return expanded
|
||||||
|
|
||||||
|
|
||||||
def _normalize_audio_to_numpy(audio: Any) -> np.ndarray | None:
|
def _normalize_audio_to_numpy(audio: Any) -> np.ndarray | None:
|
||||||
"""Convert audio (torch / numpy) into a float32 numpy array in [-1, 1], best-effort."""
|
"""Convert audio (torch / numpy) into a float32 numpy array in [-1, 1], best-effort."""
|
||||||
if audio is None:
|
if audio is None:
|
||||||
|
|||||||
@@ -6,7 +6,9 @@ import logging
|
|||||||
import multiprocessing as mp
|
import multiprocessing as mp
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from typing import List, Union
|
from contextlib import ExitStack
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Callable, List, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from setproctitle import setproctitle
|
from setproctitle import setproctitle
|
||||||
@@ -65,6 +67,18 @@ from sglang.srt.utils.network import NetworkAddress
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class _ExpandedOutputParts:
|
||||||
|
tensor_outputs: list[torch.Tensor] = field(default_factory=list)
|
||||||
|
list_outputs: list[Any] = field(default_factory=list)
|
||||||
|
tensor_audio: list[torch.Tensor] = field(default_factory=list)
|
||||||
|
trajectory_latents: list[torch.Tensor] = field(default_factory=list)
|
||||||
|
noise_preds: list[torch.Tensor] = field(default_factory=list)
|
||||||
|
output_file_paths: list[str] = field(default_factory=list)
|
||||||
|
metrics_list: list[Any] = field(default_factory=list)
|
||||||
|
trajectory_decoded_parts: list[list[torch.Tensor]] | None = None
|
||||||
|
|
||||||
|
|
||||||
class GPUWorker:
|
class GPUWorker:
|
||||||
"""
|
"""
|
||||||
A worker that executes the model on a single GPU.
|
A worker that executes the model on a single GPU.
|
||||||
@@ -222,7 +236,57 @@ class GPUWorker:
|
|||||||
Used by disaggregated pipelines to access intermediate tensors.
|
Used by disaggregated pipelines to access intermediate tensors.
|
||||||
"""
|
"""
|
||||||
assert self.pipeline is not None
|
assert self.pipeline is not None
|
||||||
|
if len(batch) > 1:
|
||||||
|
if return_req:
|
||||||
|
raise ValueError(
|
||||||
|
"Grouped execute_forward does not support return_req=True"
|
||||||
|
)
|
||||||
|
# batched reqs is only possible with `num_outputs_per_prompt > 1` now
|
||||||
|
self._validate_group_forward_reqs(batch)
|
||||||
|
return self._execute_forward_batch(batch)
|
||||||
|
|
||||||
req = batch[0]
|
req = batch[0]
|
||||||
|
return self._execute_forward_common(
|
||||||
|
req,
|
||||||
|
forward_fn=lambda: self.pipeline.forward(req, self.server_args),
|
||||||
|
log_reqs=[req],
|
||||||
|
return_req=return_req,
|
||||||
|
save_output_paths=lambda output_batch: self._save_output_paths(
|
||||||
|
req, output_batch
|
||||||
|
),
|
||||||
|
error_context=f"request {req.request_id}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _execute_forward_batch(self, batch: list[Req]) -> OutputBatch:
|
||||||
|
"""Execute expanded multi-output requests as one grouped forward."""
|
||||||
|
# TODO: support early return or mix-stage execution for reqs in a group
|
||||||
|
assert self.pipeline is not None
|
||||||
|
req = batch[0]
|
||||||
|
return self._execute_forward_common(
|
||||||
|
req,
|
||||||
|
forward_fn=lambda: self._forward_group(batch),
|
||||||
|
log_reqs=batch,
|
||||||
|
return_req=False,
|
||||||
|
save_output_paths=lambda output_batch: self._save_group_output_paths(
|
||||||
|
batch, output_batch
|
||||||
|
),
|
||||||
|
error_context=f"grouped request {req.request_id}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _execute_forward_common(
|
||||||
|
self,
|
||||||
|
req: Req,
|
||||||
|
*,
|
||||||
|
forward_fn: Callable[[], Req | OutputBatch],
|
||||||
|
log_reqs: list[Req],
|
||||||
|
return_req: bool,
|
||||||
|
save_output_paths: Callable[[OutputBatch], None],
|
||||||
|
error_context: str,
|
||||||
|
) -> OutputBatch | Req:
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
forward_fn: the actual forward function for reqs
|
||||||
|
"""
|
||||||
output_batch = None
|
output_batch = None
|
||||||
try:
|
try:
|
||||||
if self.rank == 0 and not current_platform.is_cpu():
|
if self.rank == 0 and not current_platform.is_cpu():
|
||||||
@@ -230,47 +294,33 @@ class GPUWorker:
|
|||||||
|
|
||||||
start_time = time.monotonic()
|
start_time = time.monotonic()
|
||||||
|
|
||||||
# capture memory baseline before forward
|
request_metrics = [
|
||||||
if self.rank == 0 and req.metrics and not current_platform.is_cpu():
|
item.metrics for item in log_reqs if item.metrics is not None
|
||||||
|
]
|
||||||
|
if self.rank == 0 and request_metrics and not current_platform.is_cpu():
|
||||||
baseline_snapshot = capture_memory_snapshot()
|
baseline_snapshot = capture_memory_snapshot()
|
||||||
req.metrics.record_memory_snapshot("before_forward", baseline_snapshot)
|
for metrics in request_metrics:
|
||||||
|
metrics.record_memory_snapshot("before_forward", baseline_snapshot)
|
||||||
|
|
||||||
req.log(server_args=self.server_args)
|
for item in log_reqs:
|
||||||
with trace_slice(req.trace_ctx, DiffStage.GPU_FORWARD):
|
item.log(server_args=self.server_args)
|
||||||
result = self.pipeline.forward(req, self.server_args)
|
with ExitStack() as stack:
|
||||||
|
for item in log_reqs:
|
||||||
|
stack.enter_context(
|
||||||
|
trace_slice(item.trace_ctx, DiffStage.GPU_FORWARD)
|
||||||
|
)
|
||||||
|
result = forward_fn()
|
||||||
|
|
||||||
# For disagg roles, return raw Req to let the caller handle
|
|
||||||
# the role-to-role tensor transfer before OutputBatch conversion.
|
|
||||||
if return_req and isinstance(result, Req):
|
if return_req and isinstance(result, Req):
|
||||||
return result
|
return result
|
||||||
|
|
||||||
if isinstance(result, Req):
|
output_batch = self._to_output_batch(result)
|
||||||
output_batch = OutputBatch(
|
|
||||||
output=result.output,
|
|
||||||
audio=getattr(result, "audio", None),
|
|
||||||
audio_sample_rate=getattr(result, "audio_sample_rate", None),
|
|
||||||
metrics=result.metrics,
|
|
||||||
trajectory_timesteps=getattr(result, "trajectory_timesteps", None),
|
|
||||||
trajectory_latents=getattr(result, "trajectory_latents", None),
|
|
||||||
rollout_trajectory_data=getattr(
|
|
||||||
result, "rollout_trajectory_data", None
|
|
||||||
),
|
|
||||||
noise_pred=getattr(result, "noise_pred", None),
|
|
||||||
trajectory_decoded=getattr(result, "trajectory_decoded", None),
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
output_batch = result
|
|
||||||
|
|
||||||
# capture memory after forward (peak)
|
output_metrics = self._iter_output_metrics(output_batch)
|
||||||
if (
|
if self.rank == 0 and output_metrics and not current_platform.is_cpu():
|
||||||
self.rank == 0
|
|
||||||
and output_batch.metrics
|
|
||||||
and not current_platform.is_cpu()
|
|
||||||
):
|
|
||||||
peak_snapshot = capture_memory_snapshot()
|
peak_snapshot = capture_memory_snapshot()
|
||||||
output_batch.metrics.record_memory_snapshot(
|
for metrics in output_metrics:
|
||||||
"after_forward", peak_snapshot
|
metrics.record_memory_snapshot("after_forward", peak_snapshot)
|
||||||
)
|
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.rank == 0
|
self.rank == 0
|
||||||
@@ -281,34 +331,11 @@ class GPUWorker:
|
|||||||
self.do_mem_analysis(output_batch)
|
self.do_mem_analysis(output_batch)
|
||||||
|
|
||||||
duration_ms = (time.monotonic() - start_time) * 1000
|
duration_ms = (time.monotonic() - start_time) * 1000
|
||||||
if output_batch.metrics is not None:
|
for metrics in output_metrics:
|
||||||
output_batch.metrics.total_duration_ms = duration_ms
|
metrics.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:
|
if req.save_output and req.return_file_paths_only:
|
||||||
if self.rank == 0 and output_batch.output is not None:
|
save_output_paths(output_batch)
|
||||||
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_compression=req.output_compression,
|
|
||||||
enable_frame_interpolation=req.enable_frame_interpolation,
|
|
||||||
frame_interpolation_exp=req.frame_interpolation_exp,
|
|
||||||
frame_interpolation_scale=req.frame_interpolation_scale,
|
|
||||||
frame_interpolation_model_path=req.frame_interpolation_model_path,
|
|
||||||
enable_upscaling=req.enable_upscaling,
|
|
||||||
upscaling_model_path=req.upscaling_model_path,
|
|
||||||
upscaling_scale=req.upscaling_scale,
|
|
||||||
)
|
|
||||||
output_batch.output_file_paths = output_paths
|
|
||||||
|
|
||||||
# No rank needs to hold on to generated tensors once the file-path
|
|
||||||
# response has been materialized on rank 0
|
|
||||||
output_batch.output = None
|
output_batch.output = None
|
||||||
output_batch.audio = None
|
output_batch.audio = None
|
||||||
output_batch.audio_sample_rate = None
|
output_batch.audio_sample_rate = None
|
||||||
@@ -316,13 +343,13 @@ class GPUWorker:
|
|||||||
if torch.cuda.is_initialized():
|
if torch.cuda.is_initialized():
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
# TODO: extract to avoid duplication
|
if torch.cuda.is_initialized() and output_batch.output is None:
|
||||||
|
torch.cuda.empty_cache()
|
||||||
|
|
||||||
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.
|
|
||||||
if not req.is_warmup:
|
if not req.is_warmup:
|
||||||
PerformanceLogger.log_request_summary(metrics=output_batch.metrics)
|
PerformanceLogger.log_request_summary(metrics=output_batch.metrics)
|
||||||
|
|
||||||
# dump per-request perf report to specified file (server mode)
|
|
||||||
if (
|
if (
|
||||||
req.perf_dump_path is not None
|
req.perf_dump_path is not None
|
||||||
and not req.is_warmup
|
and not req.is_warmup
|
||||||
@@ -336,15 +363,240 @@ class GPUWorker:
|
|||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(
|
logger.error(
|
||||||
f"Error executing request {req.request_id}: {e}", exc_info=True
|
f"Error executing {error_context}: {e}",
|
||||||
|
exc_info=True,
|
||||||
)
|
)
|
||||||
if isinstance(e, _oom_exceptions()):
|
if isinstance(e, _oom_exceptions()):
|
||||||
logger.warning(OOM_MSG)
|
logger.warning(OOM_MSG)
|
||||||
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 {error_context}: {e}"
|
||||||
return output_batch
|
return output_batch
|
||||||
|
|
||||||
|
def _forward_group(self, batch: list[Req]) -> OutputBatch:
|
||||||
|
assert self.pipeline is not None
|
||||||
|
results = self.pipeline.forward_batch(batch, self.server_args)
|
||||||
|
output_batches = [self._to_output_batch(result) for result in results]
|
||||||
|
return self._merge_expanded_output_batches(output_batches)
|
||||||
|
|
||||||
|
def _save_output_paths(self, req: Req, output_batch: OutputBatch) -> None:
|
||||||
|
if self.rank != 0 or output_batch.output is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
output_batch.output_file_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_compression=req.output_compression,
|
||||||
|
enable_frame_interpolation=req.enable_frame_interpolation,
|
||||||
|
frame_interpolation_exp=req.frame_interpolation_exp,
|
||||||
|
frame_interpolation_scale=req.frame_interpolation_scale,
|
||||||
|
frame_interpolation_model_path=req.frame_interpolation_model_path,
|
||||||
|
enable_upscaling=req.enable_upscaling,
|
||||||
|
upscaling_model_path=req.upscaling_model_path,
|
||||||
|
upscaling_scale=req.upscaling_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _save_group_output_paths(
|
||||||
|
self,
|
||||||
|
reqs: list[Req],
|
||||||
|
output_batch: OutputBatch,
|
||||||
|
) -> None:
|
||||||
|
if self.rank != 0 or output_batch.output is None:
|
||||||
|
return
|
||||||
|
if len(output_batch.output) != len(reqs):
|
||||||
|
raise RuntimeError(
|
||||||
|
f"Expected {len(reqs)} grouped outputs, got {len(output_batch.output)}"
|
||||||
|
)
|
||||||
|
|
||||||
|
first_req = reqs[0]
|
||||||
|
output_batch.output_file_paths = save_outputs(
|
||||||
|
output_batch.output,
|
||||||
|
first_req.data_type,
|
||||||
|
first_req.fps,
|
||||||
|
True,
|
||||||
|
lambda idx: reqs[idx].output_file_path(1, 0),
|
||||||
|
audio=output_batch.audio,
|
||||||
|
audio_sample_rate=output_batch.audio_sample_rate,
|
||||||
|
output_compression=first_req.output_compression,
|
||||||
|
enable_frame_interpolation=first_req.enable_frame_interpolation,
|
||||||
|
frame_interpolation_exp=first_req.frame_interpolation_exp,
|
||||||
|
frame_interpolation_scale=first_req.frame_interpolation_scale,
|
||||||
|
frame_interpolation_model_path=first_req.frame_interpolation_model_path,
|
||||||
|
enable_upscaling=first_req.enable_upscaling,
|
||||||
|
upscaling_model_path=first_req.upscaling_model_path,
|
||||||
|
upscaling_scale=first_req.upscaling_scale,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _validate_group_forward_reqs(reqs: list[Req]) -> None:
|
||||||
|
"""Validate fields that the grouped output/save path treats as shared."""
|
||||||
|
first_req = reqs[0]
|
||||||
|
shared_output_fields = (
|
||||||
|
"save_output",
|
||||||
|
"return_file_paths_only",
|
||||||
|
"data_type",
|
||||||
|
"fps",
|
||||||
|
"output_compression",
|
||||||
|
"enable_frame_interpolation",
|
||||||
|
"frame_interpolation_exp",
|
||||||
|
"frame_interpolation_scale",
|
||||||
|
"frame_interpolation_model_path",
|
||||||
|
"enable_upscaling",
|
||||||
|
"upscaling_model_path",
|
||||||
|
"upscaling_scale",
|
||||||
|
)
|
||||||
|
for req in reqs[1:]:
|
||||||
|
mismatched = [
|
||||||
|
field
|
||||||
|
for field in shared_output_fields
|
||||||
|
if getattr(req, field) != getattr(first_req, field)
|
||||||
|
]
|
||||||
|
if mismatched:
|
||||||
|
raise ValueError(
|
||||||
|
"Grouped execute_forward requires matching output settings; "
|
||||||
|
f"mismatched fields: {mismatched}"
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _iter_output_metrics(output_batch: OutputBatch):
|
||||||
|
"""Return all metrics objects carried by an output batch."""
|
||||||
|
if output_batch.metrics_list is not None:
|
||||||
|
return [
|
||||||
|
metrics for metrics in output_batch.metrics_list if metrics is not None
|
||||||
|
]
|
||||||
|
if output_batch.metrics is not None:
|
||||||
|
return [output_batch.metrics]
|
||||||
|
return []
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _to_output_batch(result: Req | OutputBatch) -> OutputBatch:
|
||||||
|
if isinstance(result, Req):
|
||||||
|
return GPUWorker._req_to_output_batch(result)
|
||||||
|
return result
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _req_to_output_batch(result: Req) -> OutputBatch:
|
||||||
|
return OutputBatch(
|
||||||
|
output=result.output,
|
||||||
|
audio=getattr(result, "audio", None),
|
||||||
|
audio_sample_rate=getattr(result, "audio_sample_rate", None),
|
||||||
|
metrics=result.metrics,
|
||||||
|
trajectory_timesteps=getattr(result, "trajectory_timesteps", None),
|
||||||
|
trajectory_latents=getattr(result, "trajectory_latents", None),
|
||||||
|
rollout_trajectory_data=getattr(result, "rollout_trajectory_data", None),
|
||||||
|
noise_pred=getattr(result, "noise_pred", None),
|
||||||
|
trajectory_decoded=getattr(result, "trajectory_decoded", None),
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _merge_expanded_output_batches(
|
||||||
|
output_batches: list[OutputBatch],
|
||||||
|
) -> OutputBatch:
|
||||||
|
"""Merge per-output batches produced by grouped execution."""
|
||||||
|
merged = OutputBatch()
|
||||||
|
parts = _ExpandedOutputParts()
|
||||||
|
|
||||||
|
for output_batch in output_batches:
|
||||||
|
GPUWorker._merge_expanded_singletons(merged, output_batch)
|
||||||
|
GPUWorker._collect_expanded_parts(parts, output_batch)
|
||||||
|
|
||||||
|
GPUWorker._finalize_expanded_parts(
|
||||||
|
merged,
|
||||||
|
parts,
|
||||||
|
audio_sample_rate=output_batches[0].audio_sample_rate,
|
||||||
|
)
|
||||||
|
|
||||||
|
return merged
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _merge_expanded_singletons(
|
||||||
|
merged: OutputBatch, output_batch: OutputBatch
|
||||||
|
) -> None:
|
||||||
|
if output_batch.error is not None and merged.error is None:
|
||||||
|
merged.error = output_batch.error
|
||||||
|
merged.peak_memory_mb = max(merged.peak_memory_mb, output_batch.peak_memory_mb)
|
||||||
|
if (
|
||||||
|
merged.trajectory_timesteps is None
|
||||||
|
and output_batch.trajectory_timesteps is not None
|
||||||
|
):
|
||||||
|
merged.trajectory_timesteps = output_batch.trajectory_timesteps
|
||||||
|
if (
|
||||||
|
merged.rollout_trajectory_data is None
|
||||||
|
and output_batch.rollout_trajectory_data is not None
|
||||||
|
):
|
||||||
|
merged.rollout_trajectory_data = output_batch.rollout_trajectory_data
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _collect_expanded_parts(
|
||||||
|
parts: _ExpandedOutputParts, output_batch: OutputBatch
|
||||||
|
) -> None:
|
||||||
|
"""Collect expanded outputs"""
|
||||||
|
parts.metrics_list.append(output_batch.metrics)
|
||||||
|
if output_batch.output_file_paths:
|
||||||
|
parts.output_file_paths.extend(output_batch.output_file_paths)
|
||||||
|
if isinstance(output_batch.output, torch.Tensor):
|
||||||
|
parts.tensor_outputs.append(output_batch.output)
|
||||||
|
elif output_batch.output is not None:
|
||||||
|
parts.list_outputs.extend(output_batch.output)
|
||||||
|
if isinstance(output_batch.audio, torch.Tensor):
|
||||||
|
parts.tensor_audio.append(output_batch.audio)
|
||||||
|
if isinstance(output_batch.trajectory_latents, torch.Tensor):
|
||||||
|
parts.trajectory_latents.append(output_batch.trajectory_latents)
|
||||||
|
if isinstance(output_batch.noise_pred, torch.Tensor):
|
||||||
|
parts.noise_preds.append(output_batch.noise_pred)
|
||||||
|
if output_batch.trajectory_decoded:
|
||||||
|
GPUWorker._collect_trajectory_decoded(
|
||||||
|
parts, output_batch.trajectory_decoded
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _collect_trajectory_decoded(
|
||||||
|
parts: _ExpandedOutputParts, trajectory_decoded: list[torch.Tensor]
|
||||||
|
) -> None:
|
||||||
|
if parts.trajectory_decoded_parts is None:
|
||||||
|
parts.trajectory_decoded_parts = [[] for _ in trajectory_decoded]
|
||||||
|
for index, decoded in enumerate(trajectory_decoded):
|
||||||
|
parts.trajectory_decoded_parts[index].append(decoded)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _finalize_expanded_parts(
|
||||||
|
merged: OutputBatch,
|
||||||
|
parts: _ExpandedOutputParts,
|
||||||
|
*,
|
||||||
|
audio_sample_rate: int | None,
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
merge batched output
|
||||||
|
"""
|
||||||
|
if parts.output_file_paths:
|
||||||
|
merged.output_file_paths = parts.output_file_paths
|
||||||
|
if any(metrics is not None for metrics in parts.metrics_list):
|
||||||
|
merged.metrics_list = parts.metrics_list
|
||||||
|
merged.metrics = next(
|
||||||
|
metrics for metrics in parts.metrics_list if metrics is not None
|
||||||
|
)
|
||||||
|
if parts.tensor_outputs:
|
||||||
|
merged.output = torch.cat(parts.tensor_outputs, dim=0)
|
||||||
|
elif parts.list_outputs:
|
||||||
|
merged.output = parts.list_outputs
|
||||||
|
if parts.tensor_audio:
|
||||||
|
merged.audio = torch.cat(parts.tensor_audio, dim=0)
|
||||||
|
merged.audio_sample_rate = audio_sample_rate
|
||||||
|
if parts.trajectory_latents:
|
||||||
|
merged.trajectory_latents = torch.cat(parts.trajectory_latents, dim=0)
|
||||||
|
if parts.noise_preds:
|
||||||
|
merged.noise_pred = torch.cat(parts.noise_preds, dim=0)
|
||||||
|
if parts.trajectory_decoded_parts:
|
||||||
|
merged.trajectory_decoded = [
|
||||||
|
torch.cat(decoded_step, dim=0)
|
||||||
|
for decoded_step in parts.trajectory_decoded_parts
|
||||||
|
]
|
||||||
|
|
||||||
def get_can_stay_resident_components(
|
def get_can_stay_resident_components(
|
||||||
self, remaining_gpu_mem_gb: float
|
self, remaining_gpu_mem_gb: float
|
||||||
) -> List[str]:
|
) -> List[str]:
|
||||||
|
|||||||
@@ -115,7 +115,7 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
MergeLoraWeightsReq: self._handle_merge_lora,
|
MergeLoraWeightsReq: self._handle_merge_lora,
|
||||||
UnmergeLoraWeightsReq: self._handle_unmerge_lora,
|
UnmergeLoraWeightsReq: self._handle_unmerge_lora,
|
||||||
Req: self._handle_generation,
|
Req: self._handle_generation,
|
||||||
List[Req]: self._handle_generation,
|
list: self._handle_generation,
|
||||||
ListLorasReq: self._handle_list_loras,
|
ListLorasReq: self._handle_list_loras,
|
||||||
ShutdownReq: self._handle_shutdown,
|
ShutdownReq: self._handle_shutdown,
|
||||||
GetDisaggStatsReq: self._handle_get_disagg_stats,
|
GetDisaggStatsReq: self._handle_get_disagg_stats,
|
||||||
@@ -124,7 +124,7 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
}
|
}
|
||||||
|
|
||||||
# FIFO, new reqs are appended
|
# FIFO, new reqs are appended
|
||||||
self.waiting_queue: deque[tuple[bytes, Req]] = deque()
|
self.waiting_queue: deque[tuple[bytes, Any]] = deque()
|
||||||
|
|
||||||
# whether we've send the necessary warmup reqs
|
# whether we've send the necessary warmup reqs
|
||||||
self.warmed_up = False
|
self.warmed_up = False
|
||||||
@@ -195,7 +195,9 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
checksums = self.worker.get_weights_checksum(module_names=req.module_names)
|
checksums = self.worker.get_weights_checksum(module_names=req.module_names)
|
||||||
return OutputBatch(output=checksums)
|
return OutputBatch(output=checksums)
|
||||||
|
|
||||||
def _handle_generation(self, reqs: List[Req]):
|
def _handle_generation(self, reqs: List[Req] | list[list[Req]]):
|
||||||
|
if len(reqs) == 1 and isinstance(reqs[0], list):
|
||||||
|
reqs = reqs[0]
|
||||||
warmup_reqs = [req for req in reqs if req.is_warmup]
|
warmup_reqs = [req for req in reqs if req.is_warmup]
|
||||||
if warmup_reqs:
|
if warmup_reqs:
|
||||||
self._warmup_processed += len(warmup_reqs)
|
self._warmup_processed += len(warmup_reqs)
|
||||||
@@ -229,7 +231,7 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
if not is_warmup and self.receiver is not None and identity is not None:
|
if not is_warmup and self.receiver is not None and identity is not None:
|
||||||
self.receiver.send_multipart([identity, b"", pickle.dumps(output_batch)])
|
self.receiver.send_multipart([identity, b"", pickle.dumps(output_batch)])
|
||||||
|
|
||||||
def get_next_batch_to_run(self) -> list[tuple[bytes, Req]] | None:
|
def get_next_batch_to_run(self) -> list[tuple[bytes, Any]] | None:
|
||||||
"""pull a req from waiting_queue"""
|
"""pull a req from waiting_queue"""
|
||||||
if not self.waiting_queue:
|
if not self.waiting_queue:
|
||||||
return None
|
return None
|
||||||
@@ -342,7 +344,8 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
|
|
||||||
# handle server req-based warmup by inserting an identical req to the beginning of the waiting queue
|
# handle server req-based warmup by inserting an identical req to the beginning of the waiting queue
|
||||||
# only the very first req through server's lifetime will be warmed up
|
# only the very first req through server's lifetime will be warmed up
|
||||||
identity, req = recv_reqs[0]
|
identity, req_or_group = recv_reqs[0]
|
||||||
|
req = req_or_group[0] if isinstance(req_or_group, list) else req_or_group
|
||||||
if isinstance(req, Req):
|
if isinstance(req, Req):
|
||||||
warmup_req = req.copy_as_warmup(self.server_args.warmup_steps)
|
warmup_req = req.copy_as_warmup(self.server_args.warmup_steps)
|
||||||
recv_reqs.insert(0, (identity, warmup_req))
|
recv_reqs.insert(0, (identity, warmup_req))
|
||||||
@@ -371,12 +374,14 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
raise
|
raise
|
||||||
|
|
||||||
if recv_reqs:
|
if recv_reqs:
|
||||||
# Ensure recv_reqs is a list
|
if isinstance(recv_reqs, list) and all(
|
||||||
if not isinstance(recv_reqs, list):
|
isinstance(req, Req) for req in recv_reqs
|
||||||
recv_reqs = [recv_reqs]
|
):
|
||||||
|
recv_reqs = [(identity, recv_reqs)]
|
||||||
# Pack with identity for rank 0
|
else:
|
||||||
recv_reqs = [(identity, req) for req in recv_reqs]
|
if not isinstance(recv_reqs, list):
|
||||||
|
recv_reqs = [recv_reqs]
|
||||||
|
recv_reqs = [(identity, req) for req in recv_reqs]
|
||||||
else:
|
else:
|
||||||
recv_reqs = None
|
recv_reqs = None
|
||||||
|
|
||||||
@@ -463,9 +468,14 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
processed_req = reqs[0]
|
processed_req = reqs[0]
|
||||||
is_warmup = (
|
if isinstance(processed_req, list) and processed_req:
|
||||||
processed_req.is_warmup if isinstance(processed_req, Req) else False
|
is_warmup = processed_req[0].is_warmup
|
||||||
)
|
else:
|
||||||
|
is_warmup = (
|
||||||
|
processed_req.is_warmup
|
||||||
|
if isinstance(processed_req, Req)
|
||||||
|
else False
|
||||||
|
)
|
||||||
|
|
||||||
handler = self.request_handlers.get(type(processed_req))
|
handler = self.request_handlers.get(type(processed_req))
|
||||||
if handler:
|
if handler:
|
||||||
@@ -483,9 +493,14 @@ class Scheduler(SchedulerDisaggMixin):
|
|||||||
|
|
||||||
# 3. return results
|
# 3. return results
|
||||||
try:
|
try:
|
||||||
is_warmup = (
|
if isinstance(processed_req, list) and processed_req:
|
||||||
processed_req.is_warmup if isinstance(processed_req, Req) else False
|
is_warmup = processed_req[0].is_warmup
|
||||||
)
|
else:
|
||||||
|
is_warmup = (
|
||||||
|
processed_req.is_warmup
|
||||||
|
if isinstance(processed_req, Req)
|
||||||
|
else False
|
||||||
|
)
|
||||||
if is_warmup:
|
if is_warmup:
|
||||||
if output_batch.error is None:
|
if output_batch.error is None:
|
||||||
if self._warmup_total > 0:
|
if self._warmup_total > 0:
|
||||||
|
|||||||
@@ -739,3 +739,28 @@ class ComposedPipelineBase(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
return self.executor.execute_with_profiling(self.stages, batch, server_args)
|
return self.executor.execute_with_profiling(self.stages, batch, server_args)
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def forward_batch(
|
||||||
|
self,
|
||||||
|
batches: list[Req],
|
||||||
|
server_args: ServerArgs,
|
||||||
|
):
|
||||||
|
if len(batches) == 1:
|
||||||
|
return [self.forward(batches[0], server_args)]
|
||||||
|
|
||||||
|
if self.is_lora_set() and not self.is_lora_effective():
|
||||||
|
logger.warning(
|
||||||
|
"LoRA adapter is set, but not effective. Please make sure the LoRA weights are merged"
|
||||||
|
)
|
||||||
|
|
||||||
|
if not batches[0].is_warmup and not batches[0].suppress_logs:
|
||||||
|
logger.info(
|
||||||
|
"Running grouped pipeline stages: %s",
|
||||||
|
list(self._stage_name_mapping.keys()),
|
||||||
|
main_process_only=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
return self.executor.execute_group_with_profiling(
|
||||||
|
self.stages, batches, server_args
|
||||||
|
)
|
||||||
|
|||||||
+31
-15
@@ -1,6 +1,6 @@
|
|||||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||||
|
|
||||||
from typing import List
|
from typing import Any, Callable, List
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -52,15 +52,14 @@ class ParallelExecutor(PipelineExecutor):
|
|||||||
src=self.worker.cfg_group.ranks[0],
|
src=self.worker.cfg_group.ranks[0],
|
||||||
)
|
)
|
||||||
|
|
||||||
def _execute(
|
def _execute_stages(
|
||||||
self,
|
self,
|
||||||
stages: List[PipelineStage],
|
stages: List[PipelineStage],
|
||||||
batch: Req,
|
payload: Any,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
) -> OutputBatch:
|
run_stage: Callable[[PipelineStage, Any], Any],
|
||||||
"""
|
) -> Any:
|
||||||
Execute all pipeline stages respecting their declared parallelism type.
|
"""Execute stages while respecting their declared parallelism type."""
|
||||||
"""
|
|
||||||
if server_args.enable_cfg_parallel:
|
if server_args.enable_cfg_parallel:
|
||||||
rank = get_classifier_free_guidance_rank()
|
rank = get_classifier_free_guidance_rank()
|
||||||
else:
|
else:
|
||||||
@@ -74,23 +73,23 @@ class ParallelExecutor(PipelineExecutor):
|
|||||||
if paradigm == StageParallelismType.MAIN_RANK_ONLY:
|
if paradigm == StageParallelismType.MAIN_RANK_ONLY:
|
||||||
if rank == 0:
|
if rank == 0:
|
||||||
# Only main rank executes, others just wait
|
# Only main rank executes, others just wait
|
||||||
batch = stage(batch, server_args)
|
payload = run_stage(stage, payload)
|
||||||
torch.distributed.barrier()
|
torch.distributed.barrier()
|
||||||
|
|
||||||
elif paradigm == StageParallelismType.CFG_PARALLEL:
|
elif paradigm == StageParallelismType.CFG_PARALLEL:
|
||||||
obj_list = [batch] if rank == 0 else []
|
obj_list = [payload] if rank == 0 else []
|
||||||
broadcasted_list = broadcast_pyobj(
|
broadcasted_list = broadcast_pyobj(
|
||||||
obj_list, rank=rank, dist_group=cfg_group.cpu_group, src=0
|
obj_list, rank=rank, dist_group=cfg_group.cpu_group, src=0
|
||||||
)
|
)
|
||||||
if rank != 0:
|
if rank != 0:
|
||||||
batch = broadcasted_list[0]
|
payload = broadcasted_list[0]
|
||||||
batch = stage(batch, server_args)
|
payload = run_stage(stage, payload)
|
||||||
|
|
||||||
torch.distributed.barrier()
|
torch.distributed.barrier()
|
||||||
|
|
||||||
elif paradigm == StageParallelismType.REPLICATED:
|
elif paradigm == StageParallelismType.REPLICATED:
|
||||||
batch = stage(batch, server_args)
|
payload = run_stage(stage, payload)
|
||||||
return batch
|
return payload
|
||||||
|
|
||||||
def execute(
|
def execute(
|
||||||
self,
|
self,
|
||||||
@@ -98,5 +97,22 @@ class ParallelExecutor(PipelineExecutor):
|
|||||||
batch: Req,
|
batch: Req,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
) -> OutputBatch:
|
) -> OutputBatch:
|
||||||
batch = self._execute(stages, batch, server_args)
|
return self._execute_stages(
|
||||||
return batch
|
stages,
|
||||||
|
batch,
|
||||||
|
server_args,
|
||||||
|
lambda stage, current: stage(current, server_args),
|
||||||
|
)
|
||||||
|
|
||||||
|
def execute_group(
|
||||||
|
self,
|
||||||
|
stages: List[PipelineStage],
|
||||||
|
batches: list[Req],
|
||||||
|
server_args: ServerArgs,
|
||||||
|
):
|
||||||
|
return self._execute_stages(
|
||||||
|
stages,
|
||||||
|
batches,
|
||||||
|
server_args,
|
||||||
|
lambda stage, current: stage.run_grouped_requests(current, server_args),
|
||||||
|
)
|
||||||
|
|||||||
@@ -58,6 +58,17 @@ class PipelineExecutor(ABC):
|
|||||||
|
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def execute_group_with_profiling(
|
||||||
|
self,
|
||||||
|
stages: List["PipelineStage"],
|
||||||
|
batches: list[Req],
|
||||||
|
server_args: ServerArgs,
|
||||||
|
):
|
||||||
|
"""Execute a grouped request under the same profiler as a single request."""
|
||||||
|
with self.profile_execution(batches[0], dump_rank=0):
|
||||||
|
batches = self.execute_group(stages, batches, server_args)
|
||||||
|
return batches
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def execute(
|
def execute(
|
||||||
self,
|
self,
|
||||||
@@ -78,6 +89,22 @@ class PipelineExecutor(ABC):
|
|||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def execute_group(
|
||||||
|
self,
|
||||||
|
stages: List["PipelineStage"],
|
||||||
|
batches: list[Req],
|
||||||
|
server_args: ServerArgs,
|
||||||
|
):
|
||||||
|
"""Execute all pipeline stages over a group of independent requests.
|
||||||
|
|
||||||
|
Executors own cross-rank scheduling, while stages own whether duplicate
|
||||||
|
work can be removed. The base executor simply calls
|
||||||
|
``stage.run_grouped_requests`` for each stage in order.
|
||||||
|
"""
|
||||||
|
for stage in stages:
|
||||||
|
batches = stage.run_grouped_requests(batches, server_args)
|
||||||
|
return batches
|
||||||
|
|
||||||
@contextlib.contextmanager
|
@contextlib.contextmanager
|
||||||
def profile_execution(self, batch: Req, dump_rank: int = 0):
|
def profile_execution(self, batch: Req, dump_rank: int = 0):
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -5,7 +5,7 @@
|
|||||||
Synchronous pipeline executor implementation.
|
Synchronous pipeline executor implementation.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import List
|
from typing import Any, Callable, List
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
|
from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor import (
|
||||||
PipelineExecutor,
|
PipelineExecutor,
|
||||||
@@ -21,21 +21,33 @@ class SyncExecutor(PipelineExecutor):
|
|||||||
A simple synchronous executor that runs stages sequentially.
|
A simple synchronous executor that runs stages sequentially.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
def _run_profile_all_stages(
|
||||||
|
self,
|
||||||
|
stages: List[PipelineStage],
|
||||||
|
payload: Any,
|
||||||
|
server_args: ServerArgs,
|
||||||
|
run_stage: Callable[[PipelineStage, Any], Any],
|
||||||
|
) -> Any:
|
||||||
|
"""Execute all pipeline stages sequentially and step the profiler."""
|
||||||
|
for stage in stages:
|
||||||
|
payload = run_stage(stage, payload)
|
||||||
|
profiler = SGLDiffusionProfiler.get_instance()
|
||||||
|
if profiler:
|
||||||
|
profiler.step_stage()
|
||||||
|
return payload
|
||||||
|
|
||||||
def run_profile_all_stages(
|
def run_profile_all_stages(
|
||||||
self,
|
self,
|
||||||
stages: List[PipelineStage],
|
stages: List[PipelineStage],
|
||||||
batch: Req,
|
batch: Req,
|
||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
) -> OutputBatch:
|
) -> OutputBatch:
|
||||||
"""
|
return self._run_profile_all_stages(
|
||||||
Execute all pipeline stages sequentially.
|
stages,
|
||||||
"""
|
batch,
|
||||||
for stage in stages:
|
server_args,
|
||||||
batch = stage(batch, server_args)
|
lambda stage, current: stage(current, server_args),
|
||||||
profiler = SGLDiffusionProfiler.get_instance()
|
)
|
||||||
if profiler:
|
|
||||||
profiler.step_stage()
|
|
||||||
return batch
|
|
||||||
|
|
||||||
def execute(
|
def execute(
|
||||||
self,
|
self,
|
||||||
@@ -50,3 +62,16 @@ class SyncExecutor(PipelineExecutor):
|
|||||||
batch = self.run_profile_all_stages(stages, batch, server_args)
|
batch = self.run_profile_all_stages(stages, batch, server_args)
|
||||||
|
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def execute_group(
|
||||||
|
self,
|
||||||
|
stages: List[PipelineStage],
|
||||||
|
batches: list[Req],
|
||||||
|
server_args: ServerArgs,
|
||||||
|
):
|
||||||
|
return self._run_profile_all_stages(
|
||||||
|
stages,
|
||||||
|
batches,
|
||||||
|
server_args,
|
||||||
|
lambda stage, current: stage.run_grouped_requests(current, server_args),
|
||||||
|
)
|
||||||
|
|||||||
@@ -365,6 +365,7 @@ class OutputBatch:
|
|||||||
|
|
||||||
# logged metrics info, directly from Req.timings
|
# logged metrics info, directly from Req.timings
|
||||||
metrics: Optional["RequestMetrics"] = None
|
metrics: Optional["RequestMetrics"] = None
|
||||||
|
metrics_list: Optional[list[Optional["RequestMetrics"]]] = None
|
||||||
|
|
||||||
# For ComfyUI integration: noise prediction from denoising stage
|
# For ComfyUI integration: noise prediction from denoising stage
|
||||||
noise_pred: torch.Tensor | None = None
|
noise_pred: torch.Tensor | None = None
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ import torch
|
|||||||
|
|
||||||
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.dedup import StageDedupMixin
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||||
VerificationResult,
|
VerificationResult,
|
||||||
)
|
)
|
||||||
@@ -41,7 +42,7 @@ class StageVerificationError(Exception):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
class PipelineStage(ABC):
|
class PipelineStage(StageDedupMixin, ABC):
|
||||||
"""
|
"""
|
||||||
Abstract base class for all pipeline stages.
|
Abstract base class for all pipeline stages.
|
||||||
|
|
||||||
@@ -181,8 +182,6 @@ class PipelineStage(ABC):
|
|||||||
Execute the stage's processing on the batch with optional verification and logging.
|
Execute the stage's processing on the batch with optional verification and logging.
|
||||||
Should not be overridden by subclasses.
|
Should not be overridden by subclasses.
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
The updated batch information after this stage's processing.
|
The updated batch information after this stage's processing.
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -0,0 +1,186 @@
|
|||||||
|
"""Stage-local grouped-request dedup helpers."""
|
||||||
|
|
||||||
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from copy import deepcopy
|
||||||
|
from typing import TYPE_CHECKING, Any, ClassVar
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
|
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||||
|
|
||||||
|
|
||||||
|
class StageDedupMixin:
|
||||||
|
"""Mixin for stage-local grouped-request deduplication.
|
||||||
|
|
||||||
|
The mixin handles only stage-local reuse. It is not a global cache and does
|
||||||
|
not decide which requests are equivalent for a stage. A stage opts into the
|
||||||
|
common full-stage path by declaring the ``Req`` fields it writes through the
|
||||||
|
``deduplicated_*`` class attributes and by overriding
|
||||||
|
``build_dedup_fingerprint``.
|
||||||
|
|
||||||
|
Stages that can reuse only part of their work should override
|
||||||
|
``run_grouped_requests`` directly and may still use
|
||||||
|
``_group_requests_by_fingerprint`` for stable grouping.
|
||||||
|
"""
|
||||||
|
|
||||||
|
deduplicated_output_fields: ClassVar[tuple[str, ...]] = ()
|
||||||
|
deduplicated_tensor_tree_output_fields: ClassVar[tuple[str, ...]] = ()
|
||||||
|
deduplicated_deepcopy_output_fields: ClassVar[tuple[str, ...]] = ()
|
||||||
|
deduplicated_extra_tensor_tree_output_keys: ClassVar[tuple[str, ...]] = ()
|
||||||
|
|
||||||
|
def run_grouped_requests(
|
||||||
|
self,
|
||||||
|
batches: list["Req"],
|
||||||
|
server_args: "ServerArgs",
|
||||||
|
) -> list[Any]:
|
||||||
|
"""Run this stage for a group of independent requests.
|
||||||
|
|
||||||
|
A grouped request is still a list of normal ``Req`` objects. The group
|
||||||
|
boundary only gives a stage the opportunity to reduce duplicate work.
|
||||||
|
Stages that do not opt in keep the single-request behavior by running
|
||||||
|
``self(batch, server_args)`` for every request.
|
||||||
|
|
||||||
|
Full-stage dedup is declarative: declare the stage-owned output fields
|
||||||
|
and return a stage-local fingerprint. Partial reuse belongs in a custom
|
||||||
|
override, because the reusable unit is smaller than the whole stage.
|
||||||
|
"""
|
||||||
|
if self.has_deduplicated_output_fields():
|
||||||
|
return self.run_deduplicated_group(batches, server_args)
|
||||||
|
|
||||||
|
return [self(batch, server_args) for batch in batches]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def has_deduplicated_output_fields(cls) -> bool:
|
||||||
|
"""Return whether this stage opts into base full-stage dedup."""
|
||||||
|
return bool(
|
||||||
|
cls.deduplicated_output_fields
|
||||||
|
or cls.deduplicated_tensor_tree_output_fields
|
||||||
|
or cls.deduplicated_deepcopy_output_fields
|
||||||
|
or cls.deduplicated_extra_tensor_tree_output_keys
|
||||||
|
)
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(self, batch: "Req", server_args: "ServerArgs") -> Any:
|
||||||
|
"""Return this stage's semantic input fingerprint for grouped dedup.
|
||||||
|
|
||||||
|
A fingerprint is the stage-local set of input values that fully
|
||||||
|
determines the outputs this stage writes. The default is unique per
|
||||||
|
request, so dedup is explicit and safe by default.
|
||||||
|
|
||||||
|
Overrides should include every request/config field read by this stage
|
||||||
|
and exclude fields that only matter to later stages. If a field is a
|
||||||
|
tensor or nested container, use ``freeze_for_dedup`` so the fingerprint
|
||||||
|
remains hashable.
|
||||||
|
"""
|
||||||
|
return id(batch)
|
||||||
|
|
||||||
|
def run_deduplicated_group(
|
||||||
|
self,
|
||||||
|
batches: list["Req"],
|
||||||
|
server_args: "ServerArgs",
|
||||||
|
copy_outputs=None,
|
||||||
|
) -> list["Req"]:
|
||||||
|
"""Run full-stage-equivalent requests once and fan out stage outputs."""
|
||||||
|
if copy_outputs is None:
|
||||||
|
copy_outputs = self.copy_deduplicated_outputs
|
||||||
|
|
||||||
|
results: list[Req | None] = [None] * len(batches)
|
||||||
|
|
||||||
|
for _, group in self._group_requests_by_fingerprint(
|
||||||
|
batches, lambda batch: self.build_dedup_fingerprint(batch, server_args)
|
||||||
|
):
|
||||||
|
first_index, first_batch = group[0]
|
||||||
|
first_result = self(first_batch, server_args)
|
||||||
|
results[first_index] = first_result
|
||||||
|
|
||||||
|
for index, batch in group[1:]:
|
||||||
|
copy_outputs(first_result, batch)
|
||||||
|
results[index] = batch
|
||||||
|
|
||||||
|
return [result for result in results if result is not None]
|
||||||
|
|
||||||
|
def copy_deduplicated_outputs(self, src: "Req", dst: "Req") -> None:
|
||||||
|
"""Copy declared stage outputs from a computed request to a duplicate.
|
||||||
|
|
||||||
|
``deduplicated_output_fields`` uses shallow container copies and shares
|
||||||
|
tensor references, which is the low-overhead path for read-only outputs
|
||||||
|
such as embeddings. Tensor-tree fields recursively clone tensors.
|
||||||
|
Deepcopy fields are for mutable request-local runtime objects, such as
|
||||||
|
scheduler instances. Extra keys clone selected ``Req.extra`` entries
|
||||||
|
without replacing the destination extra dict.
|
||||||
|
"""
|
||||||
|
for field in self.deduplicated_output_fields:
|
||||||
|
setattr(dst, field, self.copy_stage_output(getattr(src, field)))
|
||||||
|
for field in self.deduplicated_tensor_tree_output_fields:
|
||||||
|
setattr(dst, field, self.clone_tensor_tree(getattr(src, field)))
|
||||||
|
for field in self.deduplicated_deepcopy_output_fields:
|
||||||
|
setattr(dst, field, deepcopy(getattr(src, field)))
|
||||||
|
for key in self.deduplicated_extra_tensor_tree_output_keys:
|
||||||
|
if key in src.extra:
|
||||||
|
dst.extra[key] = self.clone_tensor_tree(src.extra[key])
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def copy_stage_output(cls, value):
|
||||||
|
"""Shallow-copy reusable containers while preserving tensor ownership."""
|
||||||
|
if isinstance(value, list):
|
||||||
|
return list(value)
|
||||||
|
if isinstance(value, tuple):
|
||||||
|
return tuple(value)
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return dict(value)
|
||||||
|
return value
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def clone_tensor_tree(cls, value):
|
||||||
|
"""Recursively clone tensors in a small output tree."""
|
||||||
|
if isinstance(value, torch.Tensor):
|
||||||
|
return value.clone()
|
||||||
|
if isinstance(value, list):
|
||||||
|
return [cls.clone_tensor_tree(item) for item in value]
|
||||||
|
if isinstance(value, tuple):
|
||||||
|
return tuple(cls.clone_tensor_tree(item) for item in value)
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return {key: cls.clone_tensor_tree(item) for key, item in value.items()}
|
||||||
|
return value
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def freeze_for_dedup(value: Any) -> Any:
|
||||||
|
"""Convert common nested values into a hashable fingerprint fragment."""
|
||||||
|
if isinstance(value, torch.Tensor):
|
||||||
|
if value.numel() <= 256:
|
||||||
|
return (
|
||||||
|
"tensor",
|
||||||
|
tuple(value.shape),
|
||||||
|
str(value.dtype),
|
||||||
|
tuple(value.detach().cpu().reshape(-1).tolist()),
|
||||||
|
)
|
||||||
|
return ("tensor", tuple(value.shape), str(value.dtype), value.device.type)
|
||||||
|
if isinstance(value, dict):
|
||||||
|
return tuple(
|
||||||
|
sorted(
|
||||||
|
(key, StageDedupMixin.freeze_for_dedup(item))
|
||||||
|
for key, item in value.items()
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return tuple(StageDedupMixin.freeze_for_dedup(item) for item in value)
|
||||||
|
if isinstance(value, set):
|
||||||
|
return tuple(
|
||||||
|
sorted(StageDedupMixin.freeze_for_dedup(item) for item in value)
|
||||||
|
)
|
||||||
|
return value
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _group_requests_by_fingerprint(
|
||||||
|
batches: list["Req"],
|
||||||
|
fingerprint_fn,
|
||||||
|
) -> list[tuple[Any, list[tuple[int, "Req"]]]]:
|
||||||
|
"""Group requests by a stage-local fingerprint while preserving order."""
|
||||||
|
groups: dict[Any, list[tuple[int, "Req"]]] = {}
|
||||||
|
for index, batch in enumerate(batches):
|
||||||
|
fingerprint = fingerprint_fn(batch)
|
||||||
|
groups.setdefault(fingerprint, []).append((index, batch))
|
||||||
|
return list(groups.items())
|
||||||
@@ -8,6 +8,8 @@ This module contains implementations of image encoding stages for diffusion pipe
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import inspect
|
import inspect
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import PIL
|
import PIL
|
||||||
@@ -43,6 +45,67 @@ from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ImageEncodingFingerprint:
|
||||||
|
image_source: Any
|
||||||
|
prompt: Any
|
||||||
|
negative_prompt: Any
|
||||||
|
do_classifier_free_guidance: bool
|
||||||
|
height: int | None
|
||||||
|
width: int | None
|
||||||
|
num_frames: int | None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class LTX2ImageEncodingFingerprint:
|
||||||
|
image_source: Any
|
||||||
|
height: int | None
|
||||||
|
width: int | None
|
||||||
|
num_frames: int | None
|
||||||
|
latent_dtype: str
|
||||||
|
condition_encoder_subdir: str
|
||||||
|
encode_sample_mode: str
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ImageVAEEncodingFingerprint:
|
||||||
|
image_source: Any
|
||||||
|
height: int | None
|
||||||
|
width: int | None
|
||||||
|
num_frames: int | None
|
||||||
|
encode_sample_mode: str
|
||||||
|
vae_precision: Any
|
||||||
|
vae_tiling: bool
|
||||||
|
|
||||||
|
|
||||||
|
def _freeze_image_source_value(value):
|
||||||
|
"""Build a hashable identity fragment for image inputs.
|
||||||
|
|
||||||
|
Image inputs are often PIL/numpy/tensor objects. For file paths we can use
|
||||||
|
the path value; for in-memory objects we only dedup when the exact same
|
||||||
|
object instance is shared by multiple requests. This avoids expensive image
|
||||||
|
hashing and avoids treating two mutable image objects as equivalent just
|
||||||
|
because they currently have the same shape.
|
||||||
|
"""
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return tuple(_freeze_image_source_value(item) for item in value)
|
||||||
|
if isinstance(value, (str, int, float, bool, type(None))):
|
||||||
|
return value
|
||||||
|
return ("object", id(value))
|
||||||
|
|
||||||
|
|
||||||
|
def _build_image_source_fingerprint(batch: Req, *, prefer_vae_image: bool = False):
|
||||||
|
"""Return the image input fragment used by image encoding fingerprints."""
|
||||||
|
if batch.image_path is not None:
|
||||||
|
return ("path", PipelineStage.freeze_for_dedup(batch.image_path))
|
||||||
|
image = (
|
||||||
|
batch.vae_image if prefer_vae_image and batch.vae_image is not None else None
|
||||||
|
)
|
||||||
|
if image is None:
|
||||||
|
image = batch.condition_image
|
||||||
|
return ("image", _freeze_image_source_value(image))
|
||||||
|
|
||||||
|
|
||||||
class ImageEncodingStage(PipelineStage):
|
class ImageEncodingStage(PipelineStage):
|
||||||
"""
|
"""
|
||||||
Stage for encoding image prompts into embeddings for diffusion models.
|
Stage for encoding image prompts into embeddings for diffusion models.
|
||||||
@@ -51,6 +114,12 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
expected by the diffusion model.
|
expected by the diffusion model.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
deduplicated_output_fields = (
|
||||||
|
"image_embeds",
|
||||||
|
"prompt_embeds",
|
||||||
|
"negative_prompt_embeds",
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
image_processor,
|
image_processor,
|
||||||
@@ -208,6 +277,19 @@ class ImageEncodingStage(PipelineStage):
|
|||||||
|
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(
|
||||||
|
self, batch: Req, server_args: ServerArgs
|
||||||
|
) -> ImageEncodingFingerprint:
|
||||||
|
return ImageEncodingFingerprint(
|
||||||
|
image_source=_build_image_source_fingerprint(batch),
|
||||||
|
prompt=self.freeze_for_dedup(batch.prompt),
|
||||||
|
negative_prompt=self.freeze_for_dedup(batch.negative_prompt),
|
||||||
|
do_classifier_free_guidance=bool(batch.do_classifier_free_guidance),
|
||||||
|
height=batch.height,
|
||||||
|
width=batch.width,
|
||||||
|
num_frames=batch.num_frames,
|
||||||
|
)
|
||||||
|
|
||||||
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
||||||
"""Verify image encoding stage inputs."""
|
"""Verify image encoding stage inputs."""
|
||||||
result = VerificationResult()
|
result = VerificationResult()
|
||||||
@@ -234,6 +316,12 @@ class LTX2ImageEncodingStage(PipelineStage):
|
|||||||
- ``batch.ltx2_num_image_tokens``
|
- ``batch.ltx2_num_image_tokens``
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
deduplicated_output_fields = (
|
||||||
|
"condition_image",
|
||||||
|
"image_latent",
|
||||||
|
"ltx2_num_image_tokens",
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(self, vae=None, **kwargs) -> None:
|
def __init__(self, vae=None, **kwargs) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.vae = vae
|
self.vae = vae
|
||||||
@@ -558,6 +646,29 @@ class LTX2ImageEncodingStage(PipelineStage):
|
|||||||
self.offload_model()
|
self.offload_model()
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(
|
||||||
|
self, batch: Req, server_args: ServerArgs
|
||||||
|
) -> LTX2ImageEncodingFingerprint | int:
|
||||||
|
if batch.image_path is None or batch.image_latent is not None:
|
||||||
|
return id(batch)
|
||||||
|
|
||||||
|
sample_mode = server_args.pipeline_config.vae_config.encode_sample_mode()
|
||||||
|
arch_config = server_args.pipeline_config.vae_config.arch_config
|
||||||
|
encoder_subdir = str(getattr(arch_config, "condition_encoder_subdir", ""))
|
||||||
|
if not encoder_subdir and sample_mode == "sample":
|
||||||
|
return id(batch)
|
||||||
|
|
||||||
|
latent_dtype = batch.latents.dtype if batch.latents is not None else None
|
||||||
|
return LTX2ImageEncodingFingerprint(
|
||||||
|
image_source=_build_image_source_fingerprint(batch),
|
||||||
|
height=batch.height,
|
||||||
|
width=batch.width,
|
||||||
|
num_frames=batch.num_frames,
|
||||||
|
latent_dtype=str(latent_dtype),
|
||||||
|
condition_encoder_subdir=encoder_subdir,
|
||||||
|
encode_sample_mode=sample_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ImageVAEEncodingStage(PipelineStage):
|
class ImageVAEEncodingStage(PipelineStage):
|
||||||
"""
|
"""
|
||||||
@@ -567,6 +678,12 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
input format (e.g., image_latents).
|
input format (e.g., image_latents).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
deduplicated_output_fields = (
|
||||||
|
"image_latent",
|
||||||
|
"condition_image_latent_ids",
|
||||||
|
"vae_image_sizes",
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(self, vae: ParallelTiledVAE, **kwargs) -> None:
|
def __init__(self, vae: ParallelTiledVAE, **kwargs) -> None:
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.vae: ParallelTiledVAE = vae
|
self.vae: ParallelTiledVAE = vae
|
||||||
@@ -710,6 +827,26 @@ class ImageVAEEncodingStage(PipelineStage):
|
|||||||
self.offload_model()
|
self.offload_model()
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(
|
||||||
|
self, batch: Req, server_args: ServerArgs
|
||||||
|
) -> ImageVAEEncodingFingerprint | int:
|
||||||
|
if batch.condition_image is None:
|
||||||
|
return id(batch)
|
||||||
|
|
||||||
|
sample_mode = server_args.pipeline_config.vae_config.encode_sample_mode()
|
||||||
|
if sample_mode == "sample":
|
||||||
|
return id(batch)
|
||||||
|
|
||||||
|
return ImageVAEEncodingFingerprint(
|
||||||
|
image_source=_build_image_source_fingerprint(batch, prefer_vae_image=True),
|
||||||
|
height=batch.height,
|
||||||
|
width=batch.width,
|
||||||
|
num_frames=batch.num_frames,
|
||||||
|
encode_sample_mode=sample_mode,
|
||||||
|
vae_precision=server_args.pipeline_config.vae_precision,
|
||||||
|
vae_tiling=bool(server_args.pipeline_config.vae_tiling),
|
||||||
|
)
|
||||||
|
|
||||||
def retrieve_latents(
|
def retrieve_latents(
|
||||||
self,
|
self,
|
||||||
encoder_output: DiagonalGaussianDistribution,
|
encoder_output: DiagonalGaussianDistribution,
|
||||||
|
|||||||
@@ -73,7 +73,15 @@ class InputValidationStage(PipelineStage):
|
|||||||
num_videos_per_prompt = batch.num_outputs_per_prompt
|
num_videos_per_prompt = batch.num_outputs_per_prompt
|
||||||
|
|
||||||
assert seed is not None
|
assert seed is not None
|
||||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
if isinstance(seed, list):
|
||||||
|
if len(seed) != num_videos_per_prompt:
|
||||||
|
raise ValueError(
|
||||||
|
f"seed list length must match num_outputs_per_prompt "
|
||||||
|
f"({num_videos_per_prompt}), got {len(seed)}"
|
||||||
|
)
|
||||||
|
seeds = [int(item) for item in seed]
|
||||||
|
else:
|
||||||
|
seeds = [int(seed) + i for i in range(num_videos_per_prompt)]
|
||||||
batch.seeds = seeds
|
batch.seeds = seeds
|
||||||
|
|
||||||
# Create generators based on generator_device parameter
|
# Create generators based on generator_device parameter
|
||||||
@@ -393,7 +401,18 @@ class InputValidationStage(PipelineStage):
|
|||||||
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
||||||
"""Verify input validation stage inputs."""
|
"""Verify input validation stage inputs."""
|
||||||
result = VerificationResult()
|
result = VerificationResult()
|
||||||
result.add_check("seed", batch.seed, [V.not_none, V.non_negative_int])
|
result.add_check(
|
||||||
|
"seed",
|
||||||
|
batch.seed,
|
||||||
|
[
|
||||||
|
V.not_none,
|
||||||
|
lambda x: (
|
||||||
|
V.non_negative_int(x)
|
||||||
|
if not isinstance(x, list)
|
||||||
|
else bool(x) and all(V.non_negative_int(item) for item in x)
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
result.add_check(
|
result.add_check(
|
||||||
"num_videos_per_prompt", batch.num_outputs_per_prompt, V.positive_int
|
"num_videos_per_prompt", batch.num_outputs_per_prompt, V.positive_int
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -5,6 +5,10 @@
|
|||||||
Latent preparation stage for diffusion pipelines.
|
Latent preparation stage for diffusion pipelines.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import torch
|
||||||
from diffusers.utils.torch_utils import randn_tensor
|
from diffusers.utils.torch_utils import randn_tensor
|
||||||
|
|
||||||
from sglang.multimodal_gen.runtime.distributed import (
|
from sglang.multimodal_gen.runtime.distributed import (
|
||||||
@@ -24,6 +28,16 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class LatentPreparationFingerprint:
|
||||||
|
height: int | None
|
||||||
|
width: int | None
|
||||||
|
num_frames: int | None
|
||||||
|
latent_num_frames: int | None
|
||||||
|
prompt_dtype: Any
|
||||||
|
generator_device: str | None
|
||||||
|
|
||||||
|
|
||||||
class LatentPreparationStage(PipelineStage):
|
class LatentPreparationStage(PipelineStage):
|
||||||
"""
|
"""
|
||||||
Stage for preparing initial latent variables for the diffusion process.
|
Stage for preparing initial latent variables for the diffusion process.
|
||||||
@@ -115,6 +129,151 @@ class LatentPreparationStage(PipelineStage):
|
|||||||
batch.raw_latent_shape = latents.shape
|
batch.raw_latent_shape = latents.shape
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def run_grouped_requests(
|
||||||
|
self,
|
||||||
|
batches: list[Req],
|
||||||
|
server_args: ServerArgs,
|
||||||
|
) -> list[Req]:
|
||||||
|
"""Group only the deterministic latent-preparation subprocess.
|
||||||
|
|
||||||
|
Latent preparation is not a pure full-stage copy: each request still
|
||||||
|
owns its RNG stream, so raw noise must be drawn once per request with
|
||||||
|
that request's generator. The reusable part is the deterministic work
|
||||||
|
after raw noise generation, such as packing latent tokens and applying
|
||||||
|
scheduler scaling. For that reason this stage uses the common
|
||||||
|
fingerprint grouping helper but implements its own grouped execution
|
||||||
|
instead of ``run_deduplicated_group``.
|
||||||
|
"""
|
||||||
|
results: list[Req | None] = [None] * len(batches)
|
||||||
|
|
||||||
|
for _, group in self._group_requests_by_fingerprint(
|
||||||
|
batches, lambda batch: self.build_dedup_fingerprint(batch, server_args)
|
||||||
|
):
|
||||||
|
indexed_batches = group
|
||||||
|
group_batches = [batch for _, batch in indexed_batches]
|
||||||
|
if len(group_batches) == 1 or any(
|
||||||
|
batch.latents is not None for batch in group_batches
|
||||||
|
):
|
||||||
|
for index, batch in indexed_batches:
|
||||||
|
results[index] = self(batch, server_args)
|
||||||
|
continue
|
||||||
|
|
||||||
|
first_index, first_batch = indexed_batches[0]
|
||||||
|
first_result = self._prepare_grouped_latents(group_batches, server_args)
|
||||||
|
self._split_batched_latents(first_result, group_batches)
|
||||||
|
results[first_index] = first_batch
|
||||||
|
for index, batch in indexed_batches[1:]:
|
||||||
|
results[index] = batch
|
||||||
|
|
||||||
|
return [result for result in results if result is not None]
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(
|
||||||
|
self, batch: Req, server_args: ServerArgs
|
||||||
|
) -> LatentPreparationFingerprint:
|
||||||
|
prompt_dtype = (
|
||||||
|
batch.prompt_embeds[0].dtype
|
||||||
|
if isinstance(batch.prompt_embeds, list) and batch.prompt_embeds
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
latent_num_frames = self.adjust_video_length(batch, server_args)
|
||||||
|
return LatentPreparationFingerprint(
|
||||||
|
height=batch.height,
|
||||||
|
width=batch.width,
|
||||||
|
num_frames=batch.num_frames,
|
||||||
|
latent_num_frames=latent_num_frames,
|
||||||
|
prompt_dtype=prompt_dtype,
|
||||||
|
generator_device=batch.generator_device,
|
||||||
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _single_generator(batch: Req):
|
||||||
|
if isinstance(batch.generator, list):
|
||||||
|
assert len(batch.generator) == 1
|
||||||
|
return batch.generator[0]
|
||||||
|
return batch.generator
|
||||||
|
|
||||||
|
def _prepare_grouped_latents(
|
||||||
|
self,
|
||||||
|
batches: list[Req],
|
||||||
|
server_args: ServerArgs,
|
||||||
|
) -> Req:
|
||||||
|
"""Prepare grouped random latents without changing per-request RNG streams.
|
||||||
|
|
||||||
|
``randn_tensor`` accepts a list of generators, but its batched draw is not
|
||||||
|
guaranteed to match drawing each request independently. For multi-output
|
||||||
|
requests we need exact equivalence to the sequential seed path, so this
|
||||||
|
helper draws one raw latent tensor per request and only batches the
|
||||||
|
deterministic packing/scaling work.
|
||||||
|
"""
|
||||||
|
first_batch = batches[0]
|
||||||
|
latent_num_frames = self.adjust_video_length(first_batch, server_args)
|
||||||
|
batch_size = len(batches)
|
||||||
|
|
||||||
|
dtype = self._get_latent_dtype(first_batch, server_args)
|
||||||
|
device = get_local_torch_device()
|
||||||
|
num_frames = (
|
||||||
|
latent_num_frames
|
||||||
|
if latent_num_frames is not None
|
||||||
|
else first_batch.num_frames
|
||||||
|
)
|
||||||
|
height = first_batch.height
|
||||||
|
width = first_batch.width
|
||||||
|
|
||||||
|
if height is None or width is None:
|
||||||
|
raise ValueError("Height and width must be provided")
|
||||||
|
|
||||||
|
raw_latents = []
|
||||||
|
for batch in batches:
|
||||||
|
shape = server_args.pipeline_config.prepare_latent_shape(
|
||||||
|
batch, 1, num_frames
|
||||||
|
)
|
||||||
|
raw_latents.append(
|
||||||
|
randn_tensor(
|
||||||
|
shape,
|
||||||
|
generator=self._single_generator(batch),
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
latents = torch.cat(raw_latents, dim=0)
|
||||||
|
latent_ids = server_args.pipeline_config.maybe_prepare_latent_ids(latents)
|
||||||
|
if latent_ids is not None:
|
||||||
|
first_batch.latent_ids = latent_ids.to(device=device)
|
||||||
|
|
||||||
|
original_num_outputs = first_batch.num_outputs_per_prompt
|
||||||
|
try:
|
||||||
|
first_batch.num_outputs_per_prompt = batch_size
|
||||||
|
latents = server_args.pipeline_config.maybe_pack_latents(
|
||||||
|
latents, batch_size, first_batch
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
first_batch.num_outputs_per_prompt = original_num_outputs
|
||||||
|
|
||||||
|
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||||
|
latents = latents * self.scheduler.init_noise_sigma
|
||||||
|
|
||||||
|
first_batch.latents = latents
|
||||||
|
first_batch.raw_latent_shape = latents.shape
|
||||||
|
return first_batch
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _slice_batch_tensor(tensor: torch.Tensor, index: int, total: int):
|
||||||
|
if tensor.shape[0] == total:
|
||||||
|
return tensor[index : index + 1].contiguous()
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
def _split_batched_latents(self, src: Req, batches: list[Req]) -> None:
|
||||||
|
total = len(batches)
|
||||||
|
assert src.latents is not None
|
||||||
|
latents = src.latents
|
||||||
|
latent_ids = src.latent_ids
|
||||||
|
for index, batch in enumerate(batches):
|
||||||
|
batch.latents = self._slice_batch_tensor(latents, index, total)
|
||||||
|
batch.raw_latent_shape = batch.latents.shape
|
||||||
|
if latent_ids is not None:
|
||||||
|
batch.latent_ids = self._slice_batch_tensor(latent_ids, index, total)
|
||||||
|
|
||||||
def adjust_video_length(self, batch: Req, server_args: ServerArgs) -> int:
|
def adjust_video_length(self, batch: Req, server_args: ServerArgs) -> int:
|
||||||
"""
|
"""
|
||||||
Adjust video length based on VAE version.
|
Adjust video length based on VAE version.
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import inspect
|
import inspect
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -28,6 +30,15 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class TextEncodingFingerprint:
|
||||||
|
prompt: Any
|
||||||
|
negative_prompt: Any
|
||||||
|
do_classifier_free_guidance: bool
|
||||||
|
prompt_template: Any
|
||||||
|
max_sequence_length: int | None
|
||||||
|
|
||||||
|
|
||||||
class TextEncodingStage(PipelineStage):
|
class TextEncodingStage(PipelineStage):
|
||||||
"""
|
"""
|
||||||
Stage for encoding text prompts into embeddings for diffusion models.
|
Stage for encoding text prompts into embeddings for diffusion models.
|
||||||
@@ -36,6 +47,18 @@ class TextEncodingStage(PipelineStage):
|
|||||||
expected by the diffusion model.
|
expected by the diffusion model.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
deduplicated_output_fields = (
|
||||||
|
"prompt_embeds",
|
||||||
|
"negative_prompt_embeds",
|
||||||
|
"prompt_attention_mask",
|
||||||
|
"negative_attention_mask",
|
||||||
|
"pooled_embeds",
|
||||||
|
"neg_pooled_embeds",
|
||||||
|
"clip_embedding_pos",
|
||||||
|
"clip_embedding_neg",
|
||||||
|
"is_prompt_processed",
|
||||||
|
)
|
||||||
|
|
||||||
def __init__(self, text_encoders, tokenizers) -> None:
|
def __init__(self, text_encoders, tokenizers) -> None:
|
||||||
"""
|
"""
|
||||||
Initialize the prompt encoding stage.
|
Initialize the prompt encoding stage.
|
||||||
@@ -107,6 +130,17 @@ class TextEncodingStage(PipelineStage):
|
|||||||
|
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(
|
||||||
|
self, batch: Req, server_args: ServerArgs
|
||||||
|
) -> TextEncodingFingerprint:
|
||||||
|
return TextEncodingFingerprint(
|
||||||
|
prompt=self.freeze_for_dedup(batch.prompt),
|
||||||
|
negative_prompt=self.freeze_for_dedup(batch.negative_prompt),
|
||||||
|
do_classifier_free_guidance=bool(batch.do_classifier_free_guidance),
|
||||||
|
prompt_template=self.freeze_for_dedup(batch.prompt_template),
|
||||||
|
max_sequence_length=batch.max_sequence_length,
|
||||||
|
)
|
||||||
|
|
||||||
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
||||||
"""Verify text encoding stage inputs."""
|
"""Verify text encoding stage inputs."""
|
||||||
result = VerificationResult()
|
result = VerificationResult()
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ This module contains implementations of timestep preparation stages for diffusio
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import inspect
|
import inspect
|
||||||
|
from dataclasses import dataclass
|
||||||
from typing import Any, Callable, Tuple
|
from typing import Any, Callable, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
@@ -33,6 +34,17 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
|||||||
logger = init_logger(__name__)
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class TimestepPreparationFingerprint:
|
||||||
|
num_inference_steps: int
|
||||||
|
timesteps: Any
|
||||||
|
sigmas: Any
|
||||||
|
n_tokens: int | None
|
||||||
|
height: int | None
|
||||||
|
width: int | None
|
||||||
|
num_frames: int | None
|
||||||
|
|
||||||
|
|
||||||
class TimestepPreparationStage(PipelineStage):
|
class TimestepPreparationStage(PipelineStage):
|
||||||
"""
|
"""
|
||||||
Stage for preparing timesteps for the diffusion process.
|
Stage for preparing timesteps for the diffusion process.
|
||||||
@@ -41,6 +53,10 @@ class TimestepPreparationStage(PipelineStage):
|
|||||||
during the diffusion process.
|
during the diffusion process.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
deduplicated_tensor_tree_output_fields = ("timesteps", "sigmas")
|
||||||
|
deduplicated_deepcopy_output_fields = ("scheduler",)
|
||||||
|
deduplicated_extra_tensor_tree_output_keys = ("mu",)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
scheduler,
|
scheduler,
|
||||||
@@ -144,6 +160,19 @@ class TimestepPreparationStage(PipelineStage):
|
|||||||
self.log_debug("timesteps: %s", timesteps)
|
self.log_debug("timesteps: %s", timesteps)
|
||||||
return batch
|
return batch
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(
|
||||||
|
self, batch: Req, server_args: ServerArgs
|
||||||
|
) -> TimestepPreparationFingerprint:
|
||||||
|
return TimestepPreparationFingerprint(
|
||||||
|
num_inference_steps=batch.num_inference_steps,
|
||||||
|
timesteps=self.freeze_for_dedup(batch.timesteps),
|
||||||
|
sigmas=self.freeze_for_dedup(batch.sigmas),
|
||||||
|
n_tokens=batch.n_tokens,
|
||||||
|
height=batch.height,
|
||||||
|
width=batch.width,
|
||||||
|
num_frames=batch.num_frames,
|
||||||
|
)
|
||||||
|
|
||||||
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
def verify_input(self, batch: Req, server_args: ServerArgs) -> VerificationResult:
|
||||||
"""Verify timestep preparation stage inputs."""
|
"""Verify timestep preparation stage inputs."""
|
||||||
result = VerificationResult()
|
result = VerificationResult()
|
||||||
|
|||||||
@@ -0,0 +1,198 @@
|
|||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||||
|
expand_request_outputs,
|
||||||
|
normalize_output_seeds,
|
||||||
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage
|
||||||
|
from sglang.multimodal_gen.runtime.pipelines_core.stages.latent_preparation import (
|
||||||
|
LatentPreparationStage,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class CountingDedupStage(PipelineStage):
|
||||||
|
deduplicated_output_fields = ("prompt_embeds",)
|
||||||
|
deduplicated_tensor_tree_output_fields = ("timesteps",)
|
||||||
|
deduplicated_deepcopy_output_fields = ("scheduler",)
|
||||||
|
deduplicated_extra_tensor_tree_output_keys = ("mu",)
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.server_args = SimpleNamespace(comfyui_mode=True)
|
||||||
|
self.forward_calls = 0
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(self, batch: Req, server_args):
|
||||||
|
return batch.prompt
|
||||||
|
|
||||||
|
def forward(self, batch: Req, server_args) -> Req:
|
||||||
|
self.forward_calls += 1
|
||||||
|
value = float(self.forward_calls)
|
||||||
|
batch.prompt_embeds = [torch.tensor([value])]
|
||||||
|
batch.timesteps = torch.tensor([value])
|
||||||
|
batch.scheduler = {"state": [value]}
|
||||||
|
batch.extra["mu"] = torch.tensor([value])
|
||||||
|
return batch
|
||||||
|
|
||||||
|
|
||||||
|
class CountingLatentStage(LatentPreparationStage):
|
||||||
|
def __init__(self):
|
||||||
|
self.server_args = SimpleNamespace(comfyui_mode=True)
|
||||||
|
self.prepare_group_calls = 0
|
||||||
|
self.forward_calls = 0
|
||||||
|
|
||||||
|
def build_dedup_fingerprint(self, batch: Req, server_args):
|
||||||
|
return batch.prompt
|
||||||
|
|
||||||
|
def _prepare_grouped_latents(
|
||||||
|
self,
|
||||||
|
batches: list[Req],
|
||||||
|
server_args,
|
||||||
|
) -> Req:
|
||||||
|
self.prepare_group_calls += 1
|
||||||
|
first_batch = batches[0]
|
||||||
|
first_batch.latents = torch.arange(len(batches), dtype=torch.float32).reshape(
|
||||||
|
len(batches), 1, 1
|
||||||
|
)
|
||||||
|
first_batch.latent_ids = first_batch.latents + 10
|
||||||
|
first_batch.raw_latent_shape = first_batch.latents.shape
|
||||||
|
return first_batch
|
||||||
|
|
||||||
|
def forward(self, batch: Req, server_args) -> Req:
|
||||||
|
self.forward_calls += 1
|
||||||
|
batch.latents = torch.tensor([[[100.0 + self.forward_calls]]])
|
||||||
|
return batch
|
||||||
|
|
||||||
|
|
||||||
|
class TestMultiOutputGrouping(unittest.TestCase):
|
||||||
|
def test_normalize_output_seeds_from_int(self):
|
||||||
|
self.assertEqual(
|
||||||
|
normalize_output_seeds(10, num_outputs_per_prompt=3),
|
||||||
|
[10, 11, 12],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_normalize_output_seeds_from_per_prompt_list(self):
|
||||||
|
self.assertEqual(
|
||||||
|
normalize_output_seeds([3, 5], num_outputs_per_prompt=2),
|
||||||
|
[3, 5],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_normalize_output_seeds_from_total_list(self):
|
||||||
|
self.assertEqual(
|
||||||
|
normalize_output_seeds(
|
||||||
|
[1, 2, 3, 4],
|
||||||
|
num_outputs_per_prompt=2,
|
||||||
|
num_prompts=2,
|
||||||
|
prompt_index=1,
|
||||||
|
),
|
||||||
|
[3, 4],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_normalize_output_seeds_rejects_mismatched_list(self):
|
||||||
|
with self.assertRaisesRegex(ValueError, r"seed list length"):
|
||||||
|
normalize_output_seeds(
|
||||||
|
[1, 2, 3],
|
||||||
|
num_outputs_per_prompt=2,
|
||||||
|
num_prompts=2,
|
||||||
|
prompt_index=0,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_expand_request_outputs_splits_seed_and_output_name(self):
|
||||||
|
req = Req(
|
||||||
|
sampling_params=SamplingParams(
|
||||||
|
request_id="rid",
|
||||||
|
prompt="p",
|
||||||
|
output_path="/tmp",
|
||||||
|
output_file_name="image.png",
|
||||||
|
num_outputs_per_prompt=2,
|
||||||
|
seed=[100, 101],
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
outputs = expand_request_outputs(req)
|
||||||
|
|
||||||
|
self.assertEqual([item.seed for item in outputs], [100, 101])
|
||||||
|
self.assertEqual([item.num_outputs_per_prompt for item in outputs], [1, 1])
|
||||||
|
self.assertEqual(
|
||||||
|
[item.output_file_name for item in outputs],
|
||||||
|
["image_0.png", "image_1.png"],
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
[item.request_id for item in outputs],
|
||||||
|
["rid:0", "rid:1"],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_split_batched_latents_uses_original_batched_tensor(self):
|
||||||
|
stage = LatentPreparationStage.__new__(LatentPreparationStage)
|
||||||
|
src = Req(sampling_params=SamplingParams(prompt="p"))
|
||||||
|
dst = Req(sampling_params=SamplingParams(prompt="p"))
|
||||||
|
src.latents = torch.tensor([[[1.0]], [[2.0]]])
|
||||||
|
src.latent_ids = torch.tensor([[[10.0]], [[20.0]]])
|
||||||
|
|
||||||
|
stage._split_batched_latents(src, [src, dst])
|
||||||
|
|
||||||
|
self.assertTrue(torch.equal(src.latents, torch.tensor([[[1.0]]])))
|
||||||
|
self.assertTrue(torch.equal(dst.latents, torch.tensor([[[2.0]]])))
|
||||||
|
self.assertTrue(torch.equal(src.latent_ids, torch.tensor([[[10.0]]])))
|
||||||
|
self.assertTrue(torch.equal(dst.latent_ids, torch.tensor([[[20.0]]])))
|
||||||
|
|
||||||
|
def test_declarative_stage_dedup_runs_equivalent_request_once(self):
|
||||||
|
stage = CountingDedupStage()
|
||||||
|
reqs = [
|
||||||
|
Req(sampling_params=SamplingParams(prompt="same")),
|
||||||
|
Req(sampling_params=SamplingParams(prompt="same")),
|
||||||
|
Req(sampling_params=SamplingParams(prompt="same")),
|
||||||
|
]
|
||||||
|
|
||||||
|
results = stage.run_grouped_requests(reqs, SimpleNamespace())
|
||||||
|
|
||||||
|
self.assertEqual(stage.forward_calls, 1)
|
||||||
|
self.assertEqual(results, reqs)
|
||||||
|
for req in reqs:
|
||||||
|
self.assertTrue(torch.equal(req.prompt_embeds[0], torch.tensor([1.0])))
|
||||||
|
self.assertTrue(torch.equal(req.timesteps, torch.tensor([1.0])))
|
||||||
|
self.assertEqual(req.scheduler, {"state": [1.0]})
|
||||||
|
self.assertTrue(torch.equal(req.extra["mu"], torch.tensor([1.0])))
|
||||||
|
|
||||||
|
self.assertIsNot(reqs[0].prompt_embeds, reqs[1].prompt_embeds)
|
||||||
|
self.assertIs(reqs[0].prompt_embeds[0], reqs[1].prompt_embeds[0])
|
||||||
|
self.assertIsNot(reqs[0].timesteps, reqs[1].timesteps)
|
||||||
|
self.assertIsNot(reqs[0].scheduler, reqs[1].scheduler)
|
||||||
|
self.assertIsNot(reqs[0].extra["mu"], reqs[1].extra["mu"])
|
||||||
|
|
||||||
|
def test_declarative_stage_dedup_runs_distinct_fingerprints_separately(self):
|
||||||
|
stage = CountingDedupStage()
|
||||||
|
reqs = [
|
||||||
|
Req(sampling_params=SamplingParams(prompt="a")),
|
||||||
|
Req(sampling_params=SamplingParams(prompt="b")),
|
||||||
|
]
|
||||||
|
|
||||||
|
stage.run_grouped_requests(reqs, SimpleNamespace())
|
||||||
|
|
||||||
|
self.assertEqual(stage.forward_calls, 2)
|
||||||
|
self.assertTrue(torch.equal(reqs[0].prompt_embeds[0], torch.tensor([1.0])))
|
||||||
|
self.assertTrue(torch.equal(reqs[1].prompt_embeds[0], torch.tensor([2.0])))
|
||||||
|
|
||||||
|
def test_latent_grouped_path_batches_equivalent_requests_once(self):
|
||||||
|
stage = CountingLatentStage()
|
||||||
|
reqs = [
|
||||||
|
Req(sampling_params=SamplingParams(prompt="same")),
|
||||||
|
Req(sampling_params=SamplingParams(prompt="same")),
|
||||||
|
]
|
||||||
|
|
||||||
|
results = stage.run_grouped_requests(reqs, SimpleNamespace())
|
||||||
|
|
||||||
|
self.assertEqual(results, reqs)
|
||||||
|
self.assertEqual(stage.prepare_group_calls, 1)
|
||||||
|
self.assertEqual(stage.forward_calls, 0)
|
||||||
|
self.assertTrue(torch.equal(reqs[0].latents, torch.tensor([[[0.0]]])))
|
||||||
|
self.assertTrue(torch.equal(reqs[1].latents, torch.tensor([[[1.0]]])))
|
||||||
|
self.assertTrue(torch.equal(reqs[0].latent_ids, torch.tensor([[[10.0]]])))
|
||||||
|
self.assertTrue(torch.equal(reqs[1].latent_ids, torch.tensor([[[11.0]]])))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -40,6 +40,14 @@ class TestSamplingParamsValidate(unittest.TestCase):
|
|||||||
with self.assertRaisesRegex(ValueError, r"num_outputs_per_prompt"):
|
with self.assertRaisesRegex(ValueError, r"num_outputs_per_prompt"):
|
||||||
SamplingParams(num_outputs_per_prompt=0)
|
SamplingParams(num_outputs_per_prompt=0)
|
||||||
|
|
||||||
|
def test_seed_accepts_int_or_non_empty_int_list(self):
|
||||||
|
self.assertEqual(SamplingParams(seed=7).seed, 7)
|
||||||
|
self.assertEqual(SamplingParams(seed=[7, 8]).seed, [7, 8])
|
||||||
|
with self.assertRaisesRegex(ValueError, r"seed list"):
|
||||||
|
SamplingParams(seed=[])
|
||||||
|
with self.assertRaisesRegex(ValueError, r"seed"):
|
||||||
|
SamplingParams(seed=[1, -1])
|
||||||
|
|
||||||
def test_fps_must_be_positive_int(self):
|
def test_fps_must_be_positive_int(self):
|
||||||
with self.assertRaisesRegex(ValueError, r"\bfps\b"):
|
with self.assertRaisesRegex(ValueError, r"\bfps\b"):
|
||||||
SamplingParams(fps=0)
|
SamplingParams(fps=0)
|
||||||
@@ -203,6 +211,13 @@ class TestSamplingParamsCliArgs(unittest.TestCase):
|
|||||||
self.assertEqual(kwargs["negative_prompt"], SamplingParams.negative_prompt)
|
self.assertEqual(kwargs["negative_prompt"], SamplingParams.negative_prompt)
|
||||||
self.assertTrue(kwargs["save_output"])
|
self.assertTrue(kwargs["save_output"])
|
||||||
|
|
||||||
|
def test_get_cli_args_accepts_seed_list(self):
|
||||||
|
self.assertEqual(self._parse_cli_kwargs(["--seed", "7"])["seed"], 7)
|
||||||
|
self.assertEqual(
|
||||||
|
self._parse_cli_kwargs(["--seed", "7", "8"])["seed"],
|
||||||
|
[7, 8],
|
||||||
|
)
|
||||||
|
|
||||||
def test_qwen_image_cli_path_preserves_model_defaults(self):
|
def test_qwen_image_cli_path_preserves_model_defaults(self):
|
||||||
params = self._make_qwen_image_params([])
|
params = self._make_qwen_image_params([])
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user