Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
151 changes: 151 additions & 0 deletions benchmarks/tts/benchmark_cosyvoice3_trt_streams.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,151 @@
"""Benchmark CosyVoice3-style CUDA stream handoff strategies."""

import argparse
import json
import statistics
import time

import torch

from vllm_omni.platforms import current_omni_platform


def run_case(
*,
device: torch.device,
steps: int,
warmup_steps: int,
producer_cycles: int,
estimator_cycles: int,
use_stream_dependencies: bool,
) -> dict[str, float | int | str]:
caller_stream = torch.cuda.Stream(device=device)
estimator_stream = torch.cuda.Stream(device=device)
input_tensor = torch.zeros(1024, device=device)
output_tensor = torch.empty_like(input_tensor)
submission_ms = []

def step(measure: bool) -> None:
with torch.cuda.stream(caller_stream):
torch.cuda._sleep(producer_cycles)
input_tensor.add_(1)

start = time.perf_counter_ns()
if use_stream_dependencies:
estimator_stream.wait_stream(caller_stream)
else:
caller_stream.synchronize()

with torch.cuda.stream(estimator_stream):
torch.cuda._sleep(estimator_cycles)
output_tensor.copy_(input_tensor)

if use_stream_dependencies:
caller_stream.wait_stream(estimator_stream)
else:
estimator_stream.synchronize()

if measure:
submission_ms.append((time.perf_counter_ns() - start) / 1e6)

for _ in range(warmup_steps):
step(measure=False)
caller_stream.synchronize()
estimator_stream.synchronize()
input_tensor.zero_()
torch.accelerator.reset_peak_memory_stats(device)
initial_allocated = torch.accelerator.memory_allocated(device)

start_event = torch.cuda.Event(enable_timing=True)
end_event = torch.cuda.Event(enable_timing=True)
wall_start = time.perf_counter()
with torch.cuda.stream(caller_stream):
start_event.record()
for _ in range(steps):
step(measure=True)
with torch.cuda.stream(caller_stream):
end_event.record()
end_event.synchronize()
wall_ms = (time.perf_counter() - wall_start) * 1e3
peak_delta_bytes = torch.accelerator.max_memory_allocated(device) - initial_allocated

expected = torch.full_like(output_tensor, steps)
torch.testing.assert_close(output_tensor, expected, rtol=0, atol=0)
return {
"mode": "wait_stream" if use_stream_dependencies else "synchronize",
"steps": steps,
"submission_median_ms": statistics.median(submission_ms),
"submission_p95_ms": statistics.quantiles(submission_ms, n=20)[18],
"submission_mean_ms": statistics.mean(submission_ms),
"gpu_elapsed_ms": start_event.elapsed_time(end_event),
"wall_elapsed_ms": wall_ms,
"peak_allocation_delta_bytes": peak_delta_bytes,
"final_value": int(output_tensor[0].item()),
}


def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--device", default="cuda:0")
parser.add_argument("--steps", type=int, default=100)
parser.add_argument("--warmup-steps", type=int, default=10)
parser.add_argument("--repeats", type=int, default=7)
parser.add_argument("--producer-cycles", type=int, default=1_000_000)
parser.add_argument("--estimator-cycles", type=int, default=2_000_000)
return parser.parse_args()


def main() -> None:
args = parse_args()
device = torch.device(args.device)
current_omni_platform.set_device(device)
torch.manual_seed(0)

old_results = []
new_results = []
for repeat in range(args.repeats):
common_args = {
"device": device,
"steps": args.steps,
"warmup_steps": args.warmup_steps,
"producer_cycles": args.producer_cycles,
"estimator_cycles": args.estimator_cycles,
}
if repeat % 2 == 0:
old_results.append(run_case(**common_args, use_stream_dependencies=False))
new_results.append(run_case(**common_args, use_stream_dependencies=True))
else:
new_results.append(run_case(**common_args, use_stream_dependencies=True))
old_results.append(run_case(**common_args, use_stream_dependencies=False))

old_submission_ms = statistics.median(result["submission_median_ms"] for result in old_results)
new_submission_ms = statistics.median(result["submission_median_ms"] for result in new_results)
old_gpu_ms = statistics.median(result["gpu_elapsed_ms"] for result in old_results)
new_gpu_ms = statistics.median(result["gpu_elapsed_ms"] for result in new_results)
print(
json.dumps(
{
"device": torch.cuda.get_device_name(device),
"torch": torch.__version__,
"cuda": torch.version.cuda,
"repeats": args.repeats,
"warmup_steps": args.warmup_steps,
"producer_cycles": args.producer_cycles,
"estimator_cycles": args.estimator_cycles,
"old_submission_median_ms": old_submission_ms,
"new_submission_median_ms": new_submission_ms,
"submission_median_speedup": old_submission_ms / new_submission_ms,
"old_gpu_elapsed_median_ms": old_gpu_ms,
"new_gpu_elapsed_median_ms": new_gpu_ms,
"gpu_elapsed_change_percent": (new_gpu_ms - old_gpu_ms) / old_gpu_ms * 100,
"parity": all(result["final_value"] == args.steps for result in old_results + new_results),
"old_runs": old_results,
"new_runs": new_results,
},
indent=2,
)
)


if __name__ == "__main__":
main()
147 changes: 147 additions & 0 deletions tests/model_executor/models/cosyvoice3/test_cosyvoice3_components.py
Original file line number Diff line number Diff line change
Expand Up @@ -281,6 +281,153 @@ def test_causal_conditional_cfm_forward(self, dummy_estimator):

assert out.shape == mu.shape

@pytest.mark.core_model
@pytest.mark.cpu
def test_trt_estimator_uses_stream_dependencies(self, monkeypatch):
from omegaconf import DictConfig

from vllm_omni.model_executor.models.cosyvoice3.code2wav_core.cfm import ConditionalCFM

class FakeStream:
def __init__(self, name, state):
self.name = name
self.state = state
self.cuda_stream = hash(name)
self.synchronize_calls = 0
self.waited_on = []

def synchronize(self):
self.synchronize_calls += 1

def wait_stream(self, stream):
self.waited_on.append(stream)

class FakeStreamContext:
def __init__(self, stream):
self.stream = stream

def __enter__(self):
self.previous = self.stream.state.current
self.stream.state.current = self.stream
return self.stream

def __exit__(self, exc_type, exc, traceback):
self.stream.state.current = self.previous

state = SimpleNamespace(current=None)
caller_stream = FakeStream("caller", state)
estimator_stream = FakeStream("estimator", state)
state.current = caller_stream
monkeypatch.setattr(torch.cuda, "current_stream", lambda *args, **kwargs: state.current)
stream_contexts = []

def stream_context(stream):
stream_contexts.append(stream)
return FakeStreamContext(stream)

monkeypatch.setattr(torch.cuda, "stream", stream_context)

class FakeContext:
def __init__(self):
self.execute_stream = None

def set_input_shape(self, name, shape):
pass

def set_tensor_address(self, name, address):
pass

def execute_async_v3(self, stream):
self.execute_stream = stream
return True

class FakeEngine:
@staticmethod
def get_tensor_name(index):
return f"tensor_{index}"

context = FakeContext()

class FakeEstimatorPool:
io_dtype = torch.float32

def __init__(self):
self.released = []

def acquire_estimator(self):
return [context, estimator_stream], FakeEngine()

def release_estimator(self, released_context, released_stream):
self.released.append((released_context, released_stream))

estimator_pool = FakeEstimatorPool()
cfm = ConditionalCFM(
in_channels=80,
cfm_params=DictConfig(
{
"sigma_min": 1e-6,
"solver": "euler",
"t_scheduler": "cosine",
"training_cfg_rate": 0.2,
"inference_cfg_rate": 0.7,
}
),
n_spks=1,
spk_emb_dim=80,
estimator=estimator_pool,
)

x = torch.randn(2, 80, 4)
mask = torch.ones(2, 1, 4)
mu = torch.randn(2, 80, 4)
timestep = torch.randn(2)
speakers = torch.randn(2, 80)
condition = torch.randn(2, 80, 4)

output = cfm.forward_estimator(x, mask, mu, timestep, speakers, condition)

assert output.shape == x.shape
assert caller_stream.synchronize_calls == 0
assert estimator_stream.synchronize_calls == 0
assert estimator_stream.waited_on == [caller_stream]
assert caller_stream.waited_on == [estimator_stream]
assert stream_contexts == [estimator_stream]
assert context.execute_stream == estimator_stream.cuda_stream
assert estimator_pool.released == [(context, estimator_stream)]

@pytest.mark.core_model
@pytest.mark.cpu
def test_trt_context_pool_stores_cuda_stream(self, monkeypatch):
from vllm_omni.model_executor.models.cosyvoice3.flow_estimator_trt import TrtContextWrapper

execution_context = object()

class FakeEngine:
@staticmethod
def create_execution_context():
return execution_context

cuda_stream = object()
monkeypatch.setattr(torch.cuda, "Stream", lambda device: cuda_stream)

def reject_stream_context(stream):
pytest.fail("TrtContextWrapper must store the CUDA stream, not a StreamContext")

monkeypatch.setattr(torch.cuda, "stream", reject_stream_context)

engine = FakeEngine()
wrapper = TrtContextWrapper(engine, device="cuda:0")
[context, stream], acquired_engine = wrapper.acquire_estimator()

assert context is execution_context
assert stream is cuda_stream
assert acquired_engine is engine

wrapper.release_estimator(context, stream)
[reused_context, reused_stream], _ = wrapper.acquire_estimator()
assert reused_context is context
assert reused_stream is stream


class TestSDPAFallback:
"""Test SDPA fallback for float32 inputs."""
Expand Down
15 changes: 10 additions & 5 deletions vllm_omni/model_executor/models/cosyvoice3/code2wav_core/cfm.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,9 +148,9 @@ def forward_estimator(self, x, mask, mu, t, spks, cond):
# ``.contiguous().data_ptr()`` could free the temp -> dangling ptr).
io_dtype = getattr(self.estimator, "io_dtype", x.dtype)
[estimator, stream], trt_engine = self.estimator.acquire_estimator()
# NOTE need to synchronize when switching stream
torch.cuda.current_stream().synchronize()
with stream:
caller_stream = torch.cuda.current_stream(x.device)
stream.wait_stream(caller_stream)
with torch.cuda.stream(stream):
x_e = x.to(io_dtype).contiguous()
mask_e = mask.to(io_dtype).contiguous()
mu_e = mu.to(io_dtype).contiguous()
Expand All @@ -176,8 +176,13 @@ def forward_estimator(self, x, mask, mu, t, spks, cond):
for i, j in enumerate(data_ptrs):
estimator.set_tensor_address(trt_engine.get_tensor_name(i), j)
# run trt engine
assert estimator.execute_async_v3(torch.cuda.current_stream().cuda_stream) is True
torch.cuda.current_stream().synchronize()
assert estimator.execute_async_v3(stream.cuda_stream) is True
for tensor in (x_e, mask_e, mu_e, t_e, spks_e, cond_e, out_e):
if tensor.is_cuda:
tensor.record_stream(stream)
caller_stream.wait_stream(stream)
if out_e.is_cuda:
out_e.record_stream(caller_stream)
self.estimator.release_estimator(estimator, stream)
return out_e.to(x.dtype)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -121,7 +121,7 @@ def __init__(
for _ in range(trt_concurrent):
ctx = engine.create_execution_context()
assert ctx is not None, "failed to create TRT execution context (out of memory?)"
stream = torch.cuda.stream(torch.cuda.Stream(torch.device(device)))
stream = torch.cuda.Stream(torch.device(device))
self._pool.put([ctx, stream])

def acquire_estimator(self):
Expand Down
Loading