项目文件夹

文件
wehub-resource-sync eec33d25b2
Build Wheel / build (3.11) (push) Failing after 1s
Build Wheel / build (3.12) (push) Failing after 0s
pre-commit / pre-commit (push) Failing after 1s
chore: import upstream snapshot with attribution
2026-07-13 12:29:08 +08:00

236 行
9.1 KiB
Python

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Perf smoke tests for Ulysses advanced_uaa communication overhead.
This test is intended for CI monitoring only:
- Print per-iteration timings and ratio vs strict Ulysses all-to-all.
- Use a loose sanity bound to catch gross regressions without flakiness.
"""
from __future__ import annotations
import os
import socket
from dataclasses import dataclass
import pytest
import torch
import torch.distributed as dist
from tests.helpers.mark import hardware_test
from vllm_omni.diffusion.attention.parallel.ulysses import (
_all_gather_int,
_ulysses_all_to_all_any_o,
_ulysses_all_to_all_any_qkv,
)
from vllm_omni.diffusion.distributed.comm import SeqAllToAll4D
from vllm_omni.diffusion.distributed.parallel_state import (
destroy_distributed_env,
get_sp_group,
init_distributed_environment,
initialize_model_parallel,
)
from vllm_omni.platforms import current_omni_platform
def _find_free_port() -> int:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("127.0.0.1", 0))
return int(s.getsockname()[1])
def _set_dist_env(*, rank: int, world_size: int, master_port: int) -> None:
os.environ["RANK"] = str(rank)
os.environ["LOCAL_RANK"] = str(rank)
os.environ["WORLD_SIZE"] = str(world_size)
os.environ["MASTER_ADDR"] = "127.0.0.1"
os.environ["MASTER_PORT"] = str(master_port)
def _max_all_reduce(pg: dist.ProcessGroup, value: float, *, device: torch.device) -> float:
t = torch.tensor([value], device=device, dtype=torch.float32)
dist.all_reduce(t, op=dist.ReduceOp.MAX, group=pg)
return float(t.item())
@dataclass(frozen=True, slots=True)
class _PerfCase:
ulysses_degree: int
ring_degree: int
@property
def world_size(self) -> int:
return int(self.ulysses_degree * self.ring_degree)
PERF_CASES: list[_PerfCase] = [
_PerfCase(ulysses_degree=4, ring_degree=1),
_PerfCase(ulysses_degree=2, ring_degree=2),
]
@pytest.mark.parametrize("case", PERF_CASES)
@pytest.mark.core_model
@hardware_test(res={"cuda": "L4"}, num_cards=4)
def test_ulysses_advanced_uaa_comm_overhead(case: _PerfCase) -> None:
available_gpus = current_omni_platform.get_device_count()
if available_gpus < case.world_size:
pytest.skip(f"Requires {case.world_size} GPUs, got {available_gpus}")
master_port = _find_free_port()
torch.multiprocessing.spawn(
_perf_worker,
args=(case.world_size, master_port, case.ulysses_degree, case.ring_degree),
nprocs=case.world_size,
)
def _perf_worker(local_rank: int, world_size: int, master_port: int, ulysses_degree: int, ring_degree: int) -> None:
device = torch.device(f"{current_omni_platform.device_type}:{local_rank}")
current_omni_platform.set_device(device)
_set_dist_env(rank=local_rank, world_size=world_size, master_port=master_port)
try:
init_distributed_environment(world_size=world_size, rank=local_rank)
initialize_model_parallel(
data_parallel_size=1,
cfg_parallel_size=1,
sequence_parallel_size=world_size,
ulysses_degree=ulysses_degree,
ring_degree=ring_degree,
tensor_parallel_size=1,
pipeline_parallel_size=1,
)
sp_group = get_sp_group()
ulysses_pg = sp_group.ulysses_group
ring_pg = sp_group.ring_group
ulysses_world_size = dist.get_world_size(ulysses_pg)
ring_world_size = dist.get_world_size(ring_pg)
# A moderate tensor size to reduce timing noise while staying fast in CI.
bsz = 1
s_local = 256
head_cnt = 32 # divisible by ulysses_degree in both cases above
head_dim = 128
dtype = torch.float16
use_sync = False
torch.manual_seed(1234 + local_rank)
q = torch.randn(bsz, s_local, head_cnt, head_dim, device=device, dtype=dtype)
k = torch.randn_like(q)
v = torch.randn_like(q)
warmup_iters = 10
iters = 200
def strict_comm_step(q_in: torch.Tensor, k_in: torch.Tensor, v_in: torch.Tensor) -> torch.Tensor:
q_out = SeqAllToAll4D.apply(ulysses_pg, q_in, 2, 1, use_sync)
SeqAllToAll4D.apply(ulysses_pg, k_in, 2, 1, use_sync)
SeqAllToAll4D.apply(ulysses_pg, v_in, 2, 1, use_sync)
o_in = q_out # dummy
return SeqAllToAll4D.apply(ulysses_pg, o_in, 1, 2, use_sync)
def uaa_comm_step(q_in: torch.Tensor, k_in: torch.Tensor, v_in: torch.Tensor) -> torch.Tensor:
seq_lens = _all_gather_int(ulysses_pg, int(q_in.shape[1]), device=q_in.device)
s_global = int(sum(seq_lens))
if ring_world_size > 1:
ring_s_globals = _all_gather_int(ring_pg, s_global, device=q_in.device)
if len(set(ring_s_globals)) != 1:
raise RuntimeError(f"Unexpected hybrid ring post-Ulysses seq_len mismatch: {ring_s_globals}.")
q_out, orig_head_cnt = _ulysses_all_to_all_any_qkv(ulysses_pg, q_in, seq_lens=seq_lens, use_sync=use_sync)
_ulysses_all_to_all_any_qkv(ulysses_pg, k_in, seq_lens=seq_lens, use_sync=use_sync)
_ulysses_all_to_all_any_qkv(ulysses_pg, v_in, seq_lens=seq_lens, use_sync=use_sync)
o_in = q_out # dummy
return _ulysses_all_to_all_any_o(
ulysses_pg,
o_in,
seq_lens=seq_lens,
local_seq_len=int(s_local),
orig_head_cnt=int(orig_head_cnt),
use_sync=use_sync,
)
with torch.no_grad():
# Warmup (strict)
for _ in range(warmup_iters):
_ = strict_comm_step(q, k, v)
current_omni_platform.synchronize()
# Timed (strict)
t0 = torch.cuda.Event(enable_timing=True)
t1 = torch.cuda.Event(enable_timing=True)
t0.record()
for _ in range(iters):
_ = strict_comm_step(q, k, v)
t1.record()
current_omni_platform.synchronize()
strict_ms = float(t0.elapsed_time(t1)) / float(iters)
# Warmup (UAA)
for _ in range(warmup_iters):
_ = uaa_comm_step(q, k, v)
current_omni_platform.synchronize()
# Timed (UAA)
u0 = torch.cuda.Event(enable_timing=True)
u1 = torch.cuda.Event(enable_timing=True)
u0.record()
for _ in range(iters):
_ = uaa_comm_step(q, k, v)
u1.record()
current_omni_platform.synchronize()
uaa_ms = float(u0.elapsed_time(u1)) / float(iters)
# Reduce across ranks (use worst-rank to be conservative).
strict_ms_max = _max_all_reduce(dist.group.WORLD, strict_ms, device=device)
uaa_ms_max = _max_all_reduce(dist.group.WORLD, uaa_ms, device=device)
ratio = (uaa_ms_max / strict_ms_max) if strict_ms_max > 0 else float("inf")
# Approx bytes moved per iteration per-rank (send+recv) for 4x all-to-all.
elem_size = torch.tensor([], dtype=dtype).element_size()
per_a2a_bytes = int(q.numel()) * elem_size * 2
comm_bytes = int(4 * per_a2a_bytes)
strict_gbps = (comm_bytes / (strict_ms_max / 1000.0)) / 1e9 if strict_ms_max > 0 else 0.0
uaa_gbps = (comm_bytes / (uaa_ms_max / 1000.0)) / 1e9 if uaa_ms_max > 0 else 0.0
if dist.get_rank() == 0:
payload = {
"name": "ulysses_uaa_comm_perf",
"world_size": int(world_size),
"ulysses_degree": int(ulysses_degree),
"ring_degree": int(ring_degree),
"ulysses_world_size": int(ulysses_world_size),
"ring_world_size": int(ring_world_size),
"shape": {
"bsz": int(bsz),
"s_local": int(s_local),
"head_cnt": int(head_cnt),
"head_dim": int(head_dim),
},
"dtype": str(dtype),
"iters": int(iters),
"strict_ms_per_iter_max": float(strict_ms_max),
"uaa_ms_per_iter_max": float(uaa_ms_max),
"uaa_over_strict_ratio": float(ratio),
"comm_bytes_per_iter_per_rank": int(comm_bytes),
"strict_effective_gbps_per_rank": float(strict_gbps),
"uaa_effective_gbps_per_rank": float(uaa_gbps),
}
print(f"UAA_COMM_PERF_JSON={payload}")
# Loose bound: we only want to catch severe regressions. The printed JSON
# payload is used for monitoring smaller changes, while this threshold is
# intentionally generous to avoid flakiness across GPU types/drivers.
max_uaa_ms_per_iter = 10.0
assert uaa_ms_max < max_uaa_ms_per_iter, (
f"UAA comm too slow: uaa={uaa_ms_max:.3f}ms/iter, strict={strict_ms_max:.3f}ms/iter "
f"(cap={max_uaa_ms_per_iter:.1f}ms/iter)."
)
assert ratio < 3.0, f"UAA comm overhead too high: ratio={ratio:.3f}x (strict={strict_ms_max:.3f}ms)."
finally:
destroy_distributed_env()