[smg][ci] preserve model launch order with test collected (#16618)
This commit is contained in:
@@ -175,15 +175,25 @@ from infra import (
|
||||
PARAM_MODEL,
|
||||
PARAM_SETUP_BACKEND,
|
||||
ConnectionMode,
|
||||
WorkerIdentity,
|
||||
WorkerType,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Test collection: scan for required backends
|
||||
# Test collection: scan for required workers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Global storage for scanned requirements
|
||||
_scanned_backends: set[str] = set() # {"grpc", "http", "openai", ...}
|
||||
_scanned_models: set[str] = set() # Models needed by tests
|
||||
# Track max worker counts: (model_id, mode, worker_type) -> max_count
|
||||
# This unified approach handles regular, prefill, and decode workers the same way
|
||||
_worker_counts: dict[tuple[str, ConnectionMode, WorkerType], int] = {}
|
||||
|
||||
# Track first-seen order to preserve test collection order
|
||||
_first_seen_order: list[tuple[str, ConnectionMode, WorkerType]] = []
|
||||
|
||||
# Track max GPU requirement for any single test (for validation)
|
||||
_max_test_gpu_requirement: int = 0
|
||||
_max_test_name: str = ""
|
||||
|
||||
_needs_default_model: bool = False # True if any e2e test lacks explicit model marker
|
||||
|
||||
|
||||
@@ -192,93 +202,259 @@ def pytest_collection_modifyitems(
|
||||
config: pytest.Config,
|
||||
items: list[pytest.Item],
|
||||
) -> None:
|
||||
"""Scan collected tests to determine required backends and models.
|
||||
"""Scan collected tests to determine required workers.
|
||||
|
||||
This runs after test collection but before tests execute.
|
||||
It extracts backend requirements from @pytest.mark.parametrize markers.
|
||||
It extracts worker requirements from markers in test collection order,
|
||||
tracking the max count needed for each (model, mode, worker_type) combination.
|
||||
|
||||
Also tracks the max GPU requirement for any single test for validation.
|
||||
"""
|
||||
global _scanned_backends, _scanned_models, _needs_default_model
|
||||
global _worker_counts, _first_seen_order, _needs_default_model
|
||||
global _max_test_gpu_requirement, _max_test_name
|
||||
|
||||
from infra import MODEL_SPECS
|
||||
|
||||
def track_worker(
|
||||
model_id: str, mode: ConnectionMode, worker_type: WorkerType, count: int
|
||||
) -> None:
|
||||
"""Track a worker requirement, updating max count if needed."""
|
||||
key = (model_id, mode, worker_type)
|
||||
if key not in _worker_counts:
|
||||
_first_seen_order.append(key)
|
||||
_worker_counts[key] = count
|
||||
else:
|
||||
_worker_counts[key] = max(_worker_counts[key], count)
|
||||
|
||||
def calculate_test_gpus(
|
||||
model_id: str, prefill: int, decode: int, regular: int
|
||||
) -> int:
|
||||
"""Calculate GPU requirement for a single test."""
|
||||
if model_id not in MODEL_SPECS:
|
||||
return 0
|
||||
tp = MODEL_SPECS[model_id].get("tp", 1)
|
||||
return tp * (prefill + decode + regular)
|
||||
|
||||
for item in items:
|
||||
# Track if this test has an explicit model marker
|
||||
has_model_marker = False
|
||||
# Extract model from marker or use default
|
||||
model_marker = item.get_closest_marker(PARAM_MODEL)
|
||||
model_id = model_marker.args[0] if model_marker and model_marker.args else None
|
||||
|
||||
# Scan parametrize markers for setup_backend
|
||||
# Check parametrize for model
|
||||
if model_id is None:
|
||||
for marker in item.iter_markers("parametrize"):
|
||||
if marker.args and len(marker.args) >= 2:
|
||||
param_name = marker.args[0]
|
||||
if param_name == PARAM_MODEL or PARAM_MODEL in param_name:
|
||||
param_values = marker.args[1]
|
||||
if isinstance(param_values, (list, tuple)) and param_values:
|
||||
model_id = param_values[0] # First model in parametrize
|
||||
break
|
||||
|
||||
# Extract backends from parametrize
|
||||
backends: list[str] = []
|
||||
for marker in item.iter_markers("parametrize"):
|
||||
if marker.args and len(marker.args) >= 2:
|
||||
param_name = marker.args[0]
|
||||
param_values = marker.args[1]
|
||||
|
||||
if param_name == PARAM_SETUP_BACKEND:
|
||||
# Extract backend names from parametrize values
|
||||
if isinstance(param_values, (list, tuple)):
|
||||
_scanned_backends.update(param_values)
|
||||
backends.extend(param_values)
|
||||
|
||||
elif param_name == PARAM_MODEL or PARAM_MODEL in param_name:
|
||||
# Extract model names from parametrize
|
||||
if isinstance(param_values, (list, tuple)):
|
||||
_scanned_models.update(param_values)
|
||||
has_model_marker = True
|
||||
# Check for workers marker (@pytest.mark.workers(...))
|
||||
workers_marker = item.get_closest_marker("workers")
|
||||
prefill_count = 0
|
||||
decode_count = 0
|
||||
regular_count = 1 # Default to 1 regular worker
|
||||
if workers_marker:
|
||||
prefill_count = workers_marker.kwargs.get("prefill") or 0
|
||||
decode_count = workers_marker.kwargs.get("decode") or 0
|
||||
regular_count = workers_marker.kwargs.get("count") or 1
|
||||
|
||||
# Check for @pytest.mark.model("name") markers
|
||||
model_marker = item.get_closest_marker(PARAM_MODEL)
|
||||
if model_marker and model_marker.args:
|
||||
model_name = model_marker.args[0]
|
||||
_scanned_models.add(model_name)
|
||||
has_model_marker = True
|
||||
|
||||
# Check if this is an e2e test without an explicit model marker
|
||||
# Such tests need the DEFAULT_MODEL
|
||||
if not has_model_marker and item.get_closest_marker("e2e"):
|
||||
# Track if this test needs default model
|
||||
is_e2e = item.get_closest_marker("e2e") is not None
|
||||
if model_id is None and is_e2e:
|
||||
_needs_default_model = True
|
||||
model_id = DEFAULT_MODEL
|
||||
|
||||
logger.info(
|
||||
"Scanned test requirements - backends: %s, models: %s, needs default: %s",
|
||||
_scanned_backends or {"(none)"},
|
||||
_scanned_models or {"(none)"},
|
||||
_needs_default_model,
|
||||
)
|
||||
# Track worker requirements and calculate this test's GPU requirement
|
||||
test_gpus = 0
|
||||
if model_id and backends:
|
||||
for backend in backends:
|
||||
# "pd" backend means PD workers
|
||||
if backend == "pd":
|
||||
mode = ConnectionMode.HTTP # PD uses HTTP mode
|
||||
# Default to 1 prefill + 1 decode if not specified
|
||||
p_count = prefill_count if prefill_count > 0 else 1
|
||||
d_count = decode_count if decode_count > 0 else 1
|
||||
track_worker(model_id, mode, WorkerType.PREFILL, p_count)
|
||||
track_worker(model_id, mode, WorkerType.DECODE, d_count)
|
||||
test_gpus = max(
|
||||
test_gpus, calculate_test_gpus(model_id, p_count, d_count, 0)
|
||||
)
|
||||
else:
|
||||
try:
|
||||
mode = ConnectionMode(backend)
|
||||
except ValueError:
|
||||
# Cloud backend (openai, xai, etc.) - skip
|
||||
continue
|
||||
|
||||
# Check if this backend also has PD workers
|
||||
if prefill_count > 0 or decode_count > 0:
|
||||
track_worker(model_id, mode, WorkerType.PREFILL, prefill_count)
|
||||
track_worker(model_id, mode, WorkerType.DECODE, decode_count)
|
||||
test_gpus = max(
|
||||
test_gpus,
|
||||
calculate_test_gpus(
|
||||
model_id, prefill_count, decode_count, 0
|
||||
),
|
||||
)
|
||||
else:
|
||||
# Regular worker
|
||||
track_worker(model_id, mode, WorkerType.REGULAR, regular_count)
|
||||
test_gpus = max(
|
||||
test_gpus,
|
||||
calculate_test_gpus(model_id, 0, 0, regular_count),
|
||||
)
|
||||
|
||||
elif model_id and is_e2e:
|
||||
# E2E test without explicit backend - will use HTTP by default
|
||||
track_worker(model_id, ConnectionMode.HTTP, WorkerType.REGULAR, 1)
|
||||
test_gpus = calculate_test_gpus(model_id, 0, 0, 1)
|
||||
|
||||
# Track max GPU requirement across all tests
|
||||
if test_gpus > _max_test_gpu_requirement:
|
||||
_max_test_gpu_requirement = test_gpus
|
||||
_max_test_name = item.nodeid
|
||||
|
||||
# Log results
|
||||
if _worker_counts:
|
||||
summary = []
|
||||
for key in _first_seen_order:
|
||||
model_id, mode, worker_type = key
|
||||
count = _worker_counts[key]
|
||||
if worker_type == WorkerType.REGULAR:
|
||||
summary.append(f"{model_id}:{mode.value}x{count}")
|
||||
else:
|
||||
summary.append(f"{model_id}:{mode.value}:{worker_type.value}x{count}")
|
||||
logger.info("Scanned worker requirements (in test order): %s", summary)
|
||||
logger.info(
|
||||
"Max GPU requirement for single test: %d (%s)",
|
||||
_max_test_gpu_requirement,
|
||||
_max_test_name,
|
||||
)
|
||||
else:
|
||||
logger.info("Scanned worker requirements: (none)")
|
||||
|
||||
|
||||
def get_pool_requirements() -> list[tuple[str, ConnectionMode]]:
|
||||
def get_pool_requirements() -> list[WorkerIdentity]:
|
||||
"""Build pool requirements from scanned test markers.
|
||||
|
||||
Returns:
|
||||
List of (model_id, ConnectionMode) tuples to try to pre-launch.
|
||||
Models that don't fit will be launched on-demand.
|
||||
List of WorkerIdentity objects to pre-launch.
|
||||
Each WorkerIdentity has (model_id, mode, worker_type, index).
|
||||
Requirements are ordered by first appearance in test collection order,
|
||||
so workers needed by earlier tests are launched first.
|
||||
|
||||
Note:
|
||||
If a model's first test needs PD workers (prefill/decode), we skip
|
||||
pre-launching regular workers for that model (they'd be evicted
|
||||
immediately when PD workers are launched).
|
||||
"""
|
||||
models = set(_scanned_models)
|
||||
# Track which models have PD workers as their first requirement
|
||||
# These models shouldn't have regular workers pre-launched
|
||||
models_with_pd_first: set[str] = set()
|
||||
first_worker_type_per_model: dict[str, WorkerType] = {}
|
||||
|
||||
# Add DEFAULT_MODEL if any e2e test lacks an explicit model marker,
|
||||
# or if no models were specified at all
|
||||
if _needs_default_model or not models:
|
||||
models.add(DEFAULT_MODEL)
|
||||
for model_id, mode, worker_type in _first_seen_order:
|
||||
if model_id not in first_worker_type_per_model:
|
||||
first_worker_type_per_model[model_id] = worker_type
|
||||
if worker_type in (WorkerType.PREFILL, WorkerType.DECODE):
|
||||
models_with_pd_first.add(model_id)
|
||||
logger.info(
|
||||
"Model %s has PD test first - skipping regular worker pre-launch",
|
||||
model_id,
|
||||
)
|
||||
|
||||
# Convert scanned string backends to ConnectionMode enums
|
||||
# Filter to local backends only (grpc, http) - cloud backends don't need workers
|
||||
local_modes: set[ConnectionMode] = set()
|
||||
for backend in _scanned_backends:
|
||||
try:
|
||||
mode = ConnectionMode(backend)
|
||||
if mode in LOCAL_MODES:
|
||||
local_modes.add(mode)
|
||||
except ValueError:
|
||||
# Not a ConnectionMode (e.g., "openai", "xai", "pd") - skip
|
||||
pass
|
||||
# Generate individual WorkerIdentity objects in first-seen order
|
||||
requirements: list[WorkerIdentity] = []
|
||||
for model_id, mode, worker_type in _first_seen_order:
|
||||
# Skip regular workers for models that have PD first
|
||||
if model_id in models_with_pd_first and worker_type == WorkerType.REGULAR:
|
||||
continue
|
||||
|
||||
# Default to HTTP if no local backends specified
|
||||
if not local_modes:
|
||||
local_modes = {ConnectionMode.HTTP}
|
||||
count = _worker_counts.get((model_id, mode, worker_type), 1)
|
||||
for i in range(count):
|
||||
requirements.append(WorkerIdentity(model_id, mode, worker_type, i))
|
||||
|
||||
# Build requirements: each model needs each mode
|
||||
requirements: list[tuple[str, ConnectionMode]] = []
|
||||
for model in models:
|
||||
for mode in local_modes:
|
||||
requirements.append((model, mode))
|
||||
# Add default if no requirements
|
||||
if not requirements:
|
||||
requirements.append(WorkerIdentity(DEFAULT_MODEL, ConnectionMode.HTTP))
|
||||
|
||||
return requirements
|
||||
|
||||
|
||||
def validate_gpu_requirements() -> tuple[int, int]:
|
||||
"""Check if there are enough GPUs for any single test.
|
||||
|
||||
Returns:
|
||||
Tuple of (max_required_gpus, available_gpus).
|
||||
|
||||
Note:
|
||||
We check the max requirement for any single test, not the sum.
|
||||
Workers can be evicted between tests, so we only need enough GPUs
|
||||
for the most demanding test.
|
||||
"""
|
||||
# Count available GPUs
|
||||
available_gpus = 0
|
||||
try:
|
||||
import torch
|
||||
|
||||
if torch.cuda.is_available():
|
||||
available_gpus = torch.cuda.device_count()
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
return _max_test_gpu_requirement, available_gpus
|
||||
|
||||
|
||||
def pytest_collection_finish(session: pytest.Session) -> None:
|
||||
"""Validate GPU requirements after test collection.
|
||||
|
||||
This runs after all tests are collected but before any tests execute.
|
||||
Fails fast if any single test requires more GPUs than available.
|
||||
"""
|
||||
if not _worker_counts:
|
||||
return
|
||||
|
||||
# Skip validation if model pool is disabled
|
||||
if os.environ.get(ENV_SKIP_MODEL_POOL, "").lower() in ("1", "true", "yes"):
|
||||
return
|
||||
|
||||
max_required, available_gpus = validate_gpu_requirements()
|
||||
|
||||
if max_required > available_gpus:
|
||||
raise pytest.UsageError(
|
||||
f"\n{'='*60}\n"
|
||||
f"GPU REQUIREMENTS EXCEEDED\n"
|
||||
f"{'='*60}\n"
|
||||
f"Test '{_max_test_name}' requires {max_required} GPUs\n"
|
||||
f"Available: {available_gpus} GPUs\n"
|
||||
f"\nOptions:\n"
|
||||
f" 1. Run tests that fit: pytest -k 'not {_max_test_name.split('::')[0]}'\n"
|
||||
f" 2. Reduce workers: @pytest.mark.workers(prefill=1, decode=1)\n"
|
||||
f" 3. Skip GPU tests: SKIP_MODEL_POOL=1 pytest\n"
|
||||
f"{'='*60}"
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"GPU validation passed: max %d required (by %s), %d available",
|
||||
max_required,
|
||||
_max_test_name,
|
||||
available_gpus,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Custom pytest markers
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -333,13 +509,15 @@ def model_pool(request: pytest.FixtureRequest) -> "ModelPool":
|
||||
routers (~1-2s) pointing to these workers.
|
||||
|
||||
Startup behavior:
|
||||
- Scans test markers to determine required (model, mode) combinations
|
||||
- Launches workers sequentially, but they boot up concurrently
|
||||
- Scans test markers to determine required workers (model, mode, type, count)
|
||||
- Launches workers in test collection order
|
||||
- Waits for all workers to become healthy before returning
|
||||
|
||||
Test requirements are auto-detected from:
|
||||
- @pytest.mark.parametrize("setup_backend", ["grpc", "http"])
|
||||
- @pytest.mark.parametrize("setup_backend", ["grpc", "http", "pd"])
|
||||
- @pytest.mark.model("model-name")
|
||||
- @pytest.mark.workers(count=N) for regular workers
|
||||
- @pytest.mark.workers(prefill=N, decode=N) for PD workers
|
||||
|
||||
Environment variable overrides:
|
||||
- E2E_MODELS: Comma-separated model IDs (e.g., "llama-8b,qwen-7b")
|
||||
@@ -388,15 +566,20 @@ def model_pool(request: pytest.FixtureRequest) -> "ModelPool":
|
||||
if not backend_modes:
|
||||
backend_modes = {ConnectionMode.HTTP}
|
||||
|
||||
requirements = [(m, b) for m in models for b in backend_modes]
|
||||
logger.info("Using env var requirements: %s", requirements)
|
||||
# Create WorkerIdentity objects (regular workers only from env vars)
|
||||
requirements = [
|
||||
WorkerIdentity(m, b, WorkerType.REGULAR, 0)
|
||||
for m in models
|
||||
for b in backend_modes
|
||||
]
|
||||
logger.info("Using env var requirements: %s", [str(r) for r in requirements])
|
||||
else:
|
||||
# Use scanned requirements from test markers
|
||||
requirements = get_pool_requirements()
|
||||
logger.info("Using scanned requirements: %s", requirements)
|
||||
logger.info("Using scanned requirements: %s", [str(r) for r in requirements])
|
||||
|
||||
# Filter to valid models
|
||||
requirements = [(m, b) for m, b in requirements if m in MODEL_SPECS]
|
||||
requirements = [r for r in requirements if r.model_id in MODEL_SPECS]
|
||||
|
||||
if not requirements:
|
||||
logger.warning("No valid requirements, model pool will be empty")
|
||||
@@ -408,7 +591,10 @@ def model_pool(request: pytest.FixtureRequest) -> "ModelPool":
|
||||
_model_pool = ModelPool(allocator)
|
||||
|
||||
startup_timeout = int(os.environ.get(ENV_STARTUP_TIMEOUT, "300"))
|
||||
_model_pool.startup(requirements=requirements, startup_timeout=startup_timeout)
|
||||
_model_pool.startup(
|
||||
requirements=requirements,
|
||||
startup_timeout=startup_timeout,
|
||||
)
|
||||
|
||||
# Log final GPU allocation summary
|
||||
logger.info(_model_pool.allocator.summary())
|
||||
@@ -624,15 +810,16 @@ def setup_backend(request: pytest.FixtureRequest, model_pool: "ModelPool"):
|
||||
f"({num_prefill} prefill + {num_decode} decode), found {gpu_count}"
|
||||
)
|
||||
|
||||
# Try to use pre-launched PD workers, or launch new ones if needed
|
||||
# Try to use pre-launched PD workers, or launch additional ones if needed
|
||||
existing_prefills = model_pool.get_workers_by_type(model_id, WorkerType.PREFILL)
|
||||
existing_decodes = model_pool.get_workers_by_type(model_id, WorkerType.DECODE)
|
||||
|
||||
if (
|
||||
len(existing_prefills) >= num_prefill
|
||||
and len(existing_decodes) >= num_decode
|
||||
):
|
||||
# Use pre-launched workers
|
||||
# Calculate how many more we need (if any)
|
||||
missing_prefill = max(0, num_prefill - len(existing_prefills))
|
||||
missing_decode = max(0, num_decode - len(existing_decodes))
|
||||
|
||||
if missing_prefill == 0 and missing_decode == 0:
|
||||
# Use pre-launched workers (we have enough)
|
||||
prefills = existing_prefills[:num_prefill]
|
||||
decodes = existing_decodes[:num_decode]
|
||||
logger.info(
|
||||
@@ -641,13 +828,48 @@ def setup_backend(request: pytest.FixtureRequest, model_pool: "ModelPool"):
|
||||
len(decodes),
|
||||
)
|
||||
else:
|
||||
# Launch new PD workers (custom config or not pre-launched)
|
||||
prefills, decodes = model_pool.launch_pd_workers(
|
||||
model_id=model_id,
|
||||
num_prefill=num_prefill,
|
||||
num_decode=num_decode,
|
||||
startup_timeout=300,
|
||||
# Build WorkerIdentity list for missing workers
|
||||
workers_to_launch: list[WorkerIdentity] = []
|
||||
for i in range(missing_prefill):
|
||||
workers_to_launch.append(
|
||||
WorkerIdentity(
|
||||
model_id,
|
||||
ConnectionMode.HTTP,
|
||||
WorkerType.PREFILL,
|
||||
len(existing_prefills) + i,
|
||||
)
|
||||
)
|
||||
for i in range(missing_decode):
|
||||
workers_to_launch.append(
|
||||
WorkerIdentity(
|
||||
model_id,
|
||||
ConnectionMode.HTTP,
|
||||
WorkerType.DECODE,
|
||||
len(existing_decodes) + i,
|
||||
)
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"Have %d/%d prefill, %d/%d decode. Launching %d more workers",
|
||||
len(existing_prefills),
|
||||
num_prefill,
|
||||
len(existing_decodes),
|
||||
num_decode,
|
||||
len(workers_to_launch),
|
||||
)
|
||||
new_instances = model_pool.launch_workers(
|
||||
workers_to_launch, startup_timeout=300
|
||||
)
|
||||
|
||||
# Combine existing + newly launched
|
||||
new_prefills = [
|
||||
w for w in new_instances if w.worker_type == WorkerType.PREFILL
|
||||
]
|
||||
new_decodes = [
|
||||
w for w in new_instances if w.worker_type == WorkerType.DECODE
|
||||
]
|
||||
prefills = existing_prefills + new_prefills
|
||||
decodes = existing_decodes + new_decodes
|
||||
|
||||
model_path = prefills[0].model_path if prefills else None
|
||||
|
||||
@@ -698,17 +920,31 @@ def setup_backend(request: pytest.FixtureRequest, model_pool: "ModelPool"):
|
||||
|
||||
try:
|
||||
if num_workers > 1:
|
||||
# Launch multiple workers on separate GPUs
|
||||
instances = model_pool.launch_regular_workers(
|
||||
model_id=model_id,
|
||||
num_workers=num_workers,
|
||||
mode=connection_mode,
|
||||
startup_timeout=300,
|
||||
)
|
||||
if not instances:
|
||||
pytest.fail(
|
||||
f"Failed to launch {num_workers} workers for {model_id}"
|
||||
# Check existing workers
|
||||
existing = model_pool.get_workers_by_type(model_id, WorkerType.REGULAR)
|
||||
existing_for_mode = [w for w in existing if w.mode == connection_mode]
|
||||
|
||||
if len(existing_for_mode) >= num_workers:
|
||||
instances = existing_for_mode[:num_workers]
|
||||
else:
|
||||
# Launch missing workers
|
||||
missing = num_workers - len(existing_for_mode)
|
||||
workers_to_launch = [
|
||||
WorkerIdentity(
|
||||
model_id,
|
||||
connection_mode,
|
||||
WorkerType.REGULAR,
|
||||
len(existing_for_mode) + i,
|
||||
)
|
||||
for i in range(missing)
|
||||
]
|
||||
new_instances = model_pool.launch_workers(
|
||||
workers_to_launch, startup_timeout=300
|
||||
)
|
||||
instances = existing_for_mode + new_instances
|
||||
|
||||
if not instances:
|
||||
pytest.fail(f"Failed to get {num_workers} workers for {model_id}")
|
||||
worker_urls = [inst.worker_url for inst in instances]
|
||||
model_path = instances[0].model_path
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user