[smg][ci] preserve model launch order with test collected (#16618)
This commit is contained in:
@@ -37,7 +37,7 @@ from .gpu_allocator import (
|
||||
)
|
||||
from .gpu_monitor import GPUMonitor
|
||||
from .gpu_monitor import should_monitor as should_monitor_gpu
|
||||
from .model_pool import ModelInstance, ModelPool
|
||||
from .model_pool import ModelInstance, ModelPool, WorkerIdentity
|
||||
from .model_specs import ( # Default model paths; Model groups
|
||||
CHAT_MODELS,
|
||||
DEFAULT_EMBEDDING_MODEL_PATH,
|
||||
@@ -63,10 +63,11 @@ from .process_utils import (
|
||||
from .run_eval import run_eval
|
||||
|
||||
__all__ = [
|
||||
# Enums
|
||||
# Enums and Identity
|
||||
"ConnectionMode",
|
||||
"WorkerType",
|
||||
"Runtime",
|
||||
"WorkerIdentity",
|
||||
# Convenience sets
|
||||
"LOCAL_MODES",
|
||||
"LOCAL_RUNTIMES",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""Constants and enums for E2E test infrastructure."""
|
||||
|
||||
from enum import Enum, auto
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class ConnectionMode(str, Enum):
|
||||
|
||||
@@ -244,19 +244,27 @@ class GPUAllocator:
|
||||
logger.warning("Failed to detect GPUs: %s", e)
|
||||
return []
|
||||
|
||||
def allocate_slots(self, model_specs: dict[str, dict]) -> list[GPUSlot]:
|
||||
def allocate_slots(
|
||||
self, model_specs: dict[str, dict], preserve_order: bool = False
|
||||
) -> list[GPUSlot]:
|
||||
"""Allocate GPU slots based on model memory requirements.
|
||||
|
||||
Uses a first-fit decreasing bin-packing algorithm:
|
||||
Uses a first-fit decreasing bin-packing algorithm by default:
|
||||
1. Sort models by memory requirement (largest first)
|
||||
2. For each model, find the first GPU(s) that can fit it
|
||||
3. For multi-GPU models, find consecutive GPUs
|
||||
|
||||
When preserve_order=True, processes models in dict insertion order
|
||||
(test collection order) instead of sorting by memory. This ensures
|
||||
models needed by earlier tests are allocated first.
|
||||
|
||||
Note: This method tracks used GPUs across multiple calls, so subsequent
|
||||
allocations will use different GPUs than previous ones.
|
||||
|
||||
Args:
|
||||
model_specs: Dict of model_id -> spec dict with 'memory_gb' and 'tp' keys
|
||||
preserve_order: If True, allocate in dict order (test order) instead
|
||||
of sorting by memory size. Default False.
|
||||
|
||||
Returns:
|
||||
List of GPUSlots with assigned models (only the newly allocated slots)
|
||||
@@ -265,17 +273,21 @@ class GPUAllocator:
|
||||
logger.warning("No GPUs available for allocation")
|
||||
return []
|
||||
|
||||
# Sort models by memory requirement (largest first for better packing)
|
||||
sorted_models = sorted(
|
||||
model_specs.items(),
|
||||
key=lambda x: x[1].get("memory_gb", 0),
|
||||
reverse=True,
|
||||
)
|
||||
if preserve_order:
|
||||
# Process in dict insertion order (test collection order)
|
||||
ordered_models = list(model_specs.items())
|
||||
else:
|
||||
# Sort models by memory requirement (largest first for better packing)
|
||||
ordered_models = sorted(
|
||||
model_specs.items(),
|
||||
key=lambda x: x[1].get("memory_gb", 0),
|
||||
reverse=True,
|
||||
)
|
||||
|
||||
# Track new slots allocated in this call
|
||||
new_slots: list[GPUSlot] = []
|
||||
|
||||
for model_id, spec in sorted_models:
|
||||
for model_id, spec in ordered_models:
|
||||
memory_gb = spec.get("memory_gb", 16)
|
||||
tp_size = spec.get("tp", 1)
|
||||
|
||||
|
||||
@@ -31,9 +31,63 @@ from .process_utils import detect_ib_device
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class WorkerIdentity:
|
||||
"""Unique identity for a single worker instance.
|
||||
|
||||
Each worker is uniquely identified by (model_id, mode, worker_type, index).
|
||||
For example:
|
||||
- llama-8b:http (regular worker, index 0)
|
||||
- llama-8b:http:prefill_0 (first prefill worker)
|
||||
- llama-8b:http:prefill_1 (second prefill worker)
|
||||
- llama-8b:http:decode_0 (first decode worker)
|
||||
|
||||
Frozen/hashable so it can be used in sets and as dict keys for deduplication.
|
||||
"""
|
||||
|
||||
model_id: str
|
||||
mode: ConnectionMode = ConnectionMode.HTTP
|
||||
worker_type: WorkerType = WorkerType.REGULAR
|
||||
index: int = 0
|
||||
|
||||
@property
|
||||
def is_prefill(self) -> bool:
|
||||
"""Check if this is a prefill worker."""
|
||||
return self.worker_type == WorkerType.PREFILL
|
||||
|
||||
@property
|
||||
def is_decode(self) -> bool:
|
||||
"""Check if this is a decode worker."""
|
||||
return self.worker_type == WorkerType.DECODE
|
||||
|
||||
@property
|
||||
def is_regular(self) -> bool:
|
||||
"""Check if this is a regular worker."""
|
||||
return self.worker_type == WorkerType.REGULAR
|
||||
|
||||
@property
|
||||
def key(self) -> str:
|
||||
"""Unique key for this worker instance."""
|
||||
if self.worker_type == WorkerType.REGULAR:
|
||||
if self.index == 0:
|
||||
return f"{self.model_id}:{self.mode.value}"
|
||||
return f"{self.model_id}:{self.mode.value}:{self.index}"
|
||||
return (
|
||||
f"{self.model_id}:{self.mode.value}:{self.worker_type.value}_{self.index}"
|
||||
)
|
||||
|
||||
def __str__(self) -> str:
|
||||
"""String representation for logging."""
|
||||
return self.key
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelInstance:
|
||||
"""A running model instance."""
|
||||
"""A running model instance.
|
||||
|
||||
Contains both identity (model_id, mode, worker_type) and runtime state
|
||||
(process, port, gpu_slot, etc.).
|
||||
"""
|
||||
|
||||
model_id: str
|
||||
mode: ConnectionMode
|
||||
@@ -42,21 +96,20 @@ class ModelInstance:
|
||||
port: int
|
||||
process: subprocess.Popen
|
||||
gpu_slot: GPUSlot | None
|
||||
key: str # Unique instance key (e.g., "llama-8b:http:prefill_0")
|
||||
worker_type: WorkerType = WorkerType.REGULAR
|
||||
bootstrap_port: int | None = None # For prefill workers in PD mode
|
||||
last_used: float = 0.0 # Timestamp for MRU eviction
|
||||
_healthy: bool = False # Track if initial health check passed
|
||||
|
||||
@property
|
||||
def key(self) -> str:
|
||||
"""Unique key for this instance.
|
||||
|
||||
Regular: 'model_id:mode' (e.g., 'llama-8b:http')
|
||||
PD workers: 'model_id:mode:worker_type' (e.g., 'llama-8b:http:prefill')
|
||||
"""
|
||||
if self.worker_type == WorkerType.REGULAR:
|
||||
return f"{self.model_id}:{self.mode.value}"
|
||||
return f"{self.model_id}:{self.mode.value}:{self.worker_type.value}"
|
||||
def identity(self) -> WorkerIdentity:
|
||||
"""Get the identity (model_id, mode, worker_type) of this instance."""
|
||||
return WorkerIdentity(
|
||||
model_id=self.model_id,
|
||||
mode=self.mode,
|
||||
worker_type=self.worker_type,
|
||||
)
|
||||
|
||||
@property
|
||||
def worker_url(self) -> str:
|
||||
@@ -200,86 +253,120 @@ class ModelPool:
|
||||
|
||||
def startup(
|
||||
self,
|
||||
requirements: list[tuple[str, ConnectionMode]] | None = None,
|
||||
requirements: list[WorkerIdentity] | None = None,
|
||||
startup_timeout: int = DEFAULT_STARTUP_TIMEOUT,
|
||||
) -> None:
|
||||
"""Start worker processes for the required models.
|
||||
"""Start worker processes for the required workers in order.
|
||||
|
||||
Workers are launched sequentially (one Popen at a time) but boot up
|
||||
concurrently since model loading happens in parallel across processes.
|
||||
This method blocks until all workers pass health checks.
|
||||
|
||||
All worker types (regular, prefill, decode) are handled uniformly.
|
||||
Each WorkerIdentity uniquely identifies a worker by (model_id, mode,
|
||||
worker_type, index).
|
||||
|
||||
Args:
|
||||
requirements: List of (model_id, mode) tuples specifying what to start.
|
||||
mode is ConnectionMode.HTTP or ConnectionMode.GRPC.
|
||||
requirements: List of WorkerIdentity specifying what to start.
|
||||
If None, starts default model in HTTP mode.
|
||||
startup_timeout: Timeout in seconds for all models to become healthy.
|
||||
"""
|
||||
self._startup_timeout = startup_timeout
|
||||
|
||||
if requirements is None:
|
||||
requirements = [(DEFAULT_MODEL, ConnectionMode.HTTP)]
|
||||
requirements = [WorkerIdentity(DEFAULT_MODEL, ConnectionMode.HTTP)]
|
||||
|
||||
# Deduplicate and validate
|
||||
requirements = list(set(requirements))
|
||||
valid_requirements = []
|
||||
for model_id, mode in requirements:
|
||||
if model_id not in MODEL_SPECS:
|
||||
logger.warning("Unknown model %s, skipping", model_id)
|
||||
# Validate requirements
|
||||
valid_requirements: list[WorkerIdentity] = []
|
||||
for identity in requirements:
|
||||
if identity.model_id not in MODEL_SPECS:
|
||||
logger.warning("Unknown model %s, skipping", identity.model_id)
|
||||
continue
|
||||
if mode not in LOCAL_MODES:
|
||||
logger.warning("Invalid mode %s for %s, skipping", mode, model_id)
|
||||
if identity.mode not in LOCAL_MODES:
|
||||
logger.warning(
|
||||
"Invalid mode %s for %s, skipping", identity.mode, identity.model_id
|
||||
)
|
||||
continue
|
||||
valid_requirements.append((model_id, mode))
|
||||
valid_requirements.append(identity)
|
||||
|
||||
if not valid_requirements:
|
||||
logger.warning("No valid requirements to start")
|
||||
return
|
||||
|
||||
logger.info("Starting model pool with: %s", valid_requirements)
|
||||
logger.info(
|
||||
"Starting model pool with %d workers: %s",
|
||||
len(valid_requirements),
|
||||
[str(r) for r in valid_requirements],
|
||||
)
|
||||
|
||||
# Build allocation specs - each (model, mode) combo needs its own slot
|
||||
# Use "model_id:mode" as the allocation key
|
||||
allocation_specs = {}
|
||||
for model_id, mode in valid_requirements:
|
||||
spec = MODEL_SPECS[model_id]
|
||||
key = f"{model_id}:{mode.value}"
|
||||
allocation_specs[key] = {
|
||||
"model": spec["model"],
|
||||
"memory_gb": spec.get("memory_gb", 16),
|
||||
"tp": spec.get("tp", 1),
|
||||
# Detect IB device once for PD workers
|
||||
has_pd = any(r.is_prefill or r.is_decode for r in valid_requirements)
|
||||
ib_device = detect_ib_device() if has_pd else None
|
||||
if ib_device:
|
||||
logger.info("Detected InfiniBand device: %s", ib_device)
|
||||
|
||||
# Track bootstrap ports for PD groups (all PD workers of same model/mode share one)
|
||||
pd_bootstrap_ports: dict[tuple[str, ConnectionMode], int] = {}
|
||||
|
||||
deferred: list[str] = []
|
||||
|
||||
# Process requirements in order - all workers treated uniformly
|
||||
for identity in valid_requirements:
|
||||
spec = get_model_spec(identity.model_id)
|
||||
tp = spec.get("tp", 1)
|
||||
|
||||
# Check if we have enough GPUs
|
||||
available_gpus = self.allocator.available_gpus()
|
||||
if len(available_gpus) < tp:
|
||||
logger.info(
|
||||
"Not enough GPUs for %s (need %d, have %d), deferring",
|
||||
identity,
|
||||
tp,
|
||||
len(available_gpus),
|
||||
)
|
||||
deferred.append(str(identity))
|
||||
continue
|
||||
|
||||
# Allocate GPU slot
|
||||
allocation_specs = {
|
||||
identity.key: {
|
||||
"model": spec["model"],
|
||||
"memory_gb": spec.get("memory_gb", 16),
|
||||
"tp": tp,
|
||||
}
|
||||
}
|
||||
slots = self.allocator.allocate_slots(allocation_specs, preserve_order=True)
|
||||
if not slots:
|
||||
deferred.append(str(identity))
|
||||
continue
|
||||
|
||||
# Allocate GPU slots
|
||||
slots = self.allocator.allocate_slots(allocation_specs)
|
||||
# Get bootstrap port for PD workers (shared within model/mode group)
|
||||
bootstrap_port = None
|
||||
if identity.is_prefill or identity.is_decode:
|
||||
pd_key = (identity.model_id, identity.mode)
|
||||
if pd_key not in pd_bootstrap_ports:
|
||||
pd_bootstrap_ports[pd_key] = get_open_port()
|
||||
bootstrap_port = pd_bootstrap_ports[pd_key]
|
||||
|
||||
# Track which models got slots
|
||||
launched_keys = set()
|
||||
# Launch the worker
|
||||
self._launch_model(
|
||||
model_id=identity.model_id,
|
||||
mode=identity.mode,
|
||||
gpu_slot=slots[0],
|
||||
worker_type=identity.worker_type,
|
||||
bootstrap_port=bootstrap_port if identity.is_prefill else None,
|
||||
ib_device=(
|
||||
ib_device if (identity.is_prefill or identity.is_decode) else None
|
||||
),
|
||||
instance_key=identity.key,
|
||||
)
|
||||
|
||||
if not slots:
|
||||
logger.warning("No GPU slots allocated, launching without GPU assignment")
|
||||
# Fallback: launch without specific GPU assignment
|
||||
for model_id, mode in valid_requirements:
|
||||
self._launch_model(model_id, mode, gpu_slot=None)
|
||||
launched_keys.add(f"{model_id}:{mode.value}")
|
||||
else:
|
||||
# Launch on allocated slots
|
||||
for slot in slots:
|
||||
if slot.assigned_model:
|
||||
# Parse "model_id:mode" back
|
||||
model_id, mode_str = slot.assigned_model.rsplit(":", 1)
|
||||
mode = ConnectionMode(mode_str)
|
||||
self._launch_model(model_id, mode, gpu_slot=slot)
|
||||
launched_keys.add(slot.assigned_model)
|
||||
|
||||
# Log models that will be launched on-demand (not enough GPUs to pre-launch)
|
||||
all_keys = set(allocation_specs.keys())
|
||||
deferred_keys = all_keys - launched_keys
|
||||
if deferred_keys:
|
||||
# Log deferred workers
|
||||
if deferred:
|
||||
logger.info(
|
||||
"%d models deferred for on-demand launch: %s",
|
||||
len(deferred_keys),
|
||||
deferred_keys,
|
||||
"%d workers deferred for on-demand launch: %s",
|
||||
len(deferred),
|
||||
deferred,
|
||||
)
|
||||
|
||||
# Wait for all launched models to be healthy
|
||||
@@ -391,6 +478,7 @@ class ModelPool:
|
||||
port=port,
|
||||
process=proc,
|
||||
gpu_slot=gpu_slot,
|
||||
key=key,
|
||||
worker_type=worker_type,
|
||||
bootstrap_port=bootstrap_port,
|
||||
last_used=time.time(),
|
||||
@@ -714,230 +802,119 @@ class ModelPool:
|
||||
if inst.model_id == model_id and inst.worker_type == worker_type
|
||||
]
|
||||
|
||||
def launch_regular_workers(
|
||||
def launch_workers(
|
||||
self,
|
||||
model_id: str,
|
||||
num_workers: int,
|
||||
mode: ConnectionMode = ConnectionMode.HTTP,
|
||||
workers: list[WorkerIdentity],
|
||||
startup_timeout: int = DEFAULT_STARTUP_TIMEOUT,
|
||||
allow_eviction: bool = True,
|
||||
) -> list[ModelInstance]:
|
||||
"""Launch multiple regular workers for load balancing.
|
||||
"""Launch workers of any type.
|
||||
|
||||
This is the unified method for launching workers. It handles all worker
|
||||
types (regular, prefill, decode) uniformly.
|
||||
|
||||
Args:
|
||||
model_id: Model identifier from MODEL_SPECS.
|
||||
num_workers: Number of workers to launch.
|
||||
mode: Connection mode (HTTP or GRPC).
|
||||
workers: List of WorkerIdentity objects specifying workers to launch.
|
||||
startup_timeout: Timeout for workers to become healthy.
|
||||
allow_eviction: If True, evict MRU models to free GPUs.
|
||||
|
||||
Returns:
|
||||
List of ModelInstance objects.
|
||||
List of launched ModelInstance objects.
|
||||
"""
|
||||
if not workers:
|
||||
return []
|
||||
|
||||
self._startup_timeout = startup_timeout
|
||||
|
||||
if model_id not in MODEL_SPECS:
|
||||
raise ValueError(f"Unknown model: {model_id}")
|
||||
# Validate all workers
|
||||
valid_workers: list[WorkerIdentity] = []
|
||||
for w in workers:
|
||||
if w.model_id not in MODEL_SPECS:
|
||||
logger.warning("Unknown model %s, skipping", w.model_id)
|
||||
continue
|
||||
if w.mode not in LOCAL_MODES:
|
||||
logger.warning("Invalid mode %s, skipping", w.mode)
|
||||
continue
|
||||
valid_workers.append(w)
|
||||
|
||||
spec = get_model_spec(model_id)
|
||||
tp = spec.get("tp", 1)
|
||||
required_gpus = num_workers * tp
|
||||
if not valid_workers:
|
||||
return []
|
||||
|
||||
# Calculate total GPUs needed
|
||||
total_gpus = 0
|
||||
for w in valid_workers:
|
||||
spec = get_model_spec(w.model_id)
|
||||
total_gpus += spec.get("tp", 1)
|
||||
|
||||
# Check if we have enough GPUs
|
||||
available = self.allocator.available_gpus()
|
||||
if len(available) < required_gpus:
|
||||
if len(available) < total_gpus:
|
||||
if allow_eviction:
|
||||
logger.info(
|
||||
"Need %d GPUs for %d workers, only %d available. Evicting MRU models...",
|
||||
required_gpus,
|
||||
num_workers,
|
||||
"Need %d GPUs for %d workers, only %d available. Evicting...",
|
||||
total_gpus,
|
||||
len(valid_workers),
|
||||
len(available),
|
||||
)
|
||||
# Exclude REGULAR workers of same model/mode from eviction
|
||||
self._evict_for_gpus(
|
||||
required_gpus,
|
||||
exclude_model_id=model_id,
|
||||
exclude_mode=mode,
|
||||
exclude_worker_types={WorkerType.REGULAR},
|
||||
)
|
||||
self._evict_for_gpus(total_gpus)
|
||||
else:
|
||||
logger.info(
|
||||
"Need %d GPUs for %d workers, only %d available. "
|
||||
"Skipping (eviction not allowed).",
|
||||
required_gpus,
|
||||
num_workers,
|
||||
logger.warning(
|
||||
"Need %d GPUs, only %d available. Skipping launch.",
|
||||
total_gpus,
|
||||
len(available),
|
||||
)
|
||||
return []
|
||||
|
||||
# Build allocation specs for all workers
|
||||
# Build allocation specs
|
||||
allocation_specs = {}
|
||||
for i in range(num_workers):
|
||||
key = f"{model_id}:{mode.value}:{i}"
|
||||
allocation_specs[key] = {
|
||||
for w in valid_workers:
|
||||
spec = get_model_spec(w.model_id)
|
||||
allocation_specs[w.key] = {
|
||||
"model": spec["model"],
|
||||
"memory_gb": spec.get("memory_gb", 16),
|
||||
"tp": tp,
|
||||
"tp": spec.get("tp", 1),
|
||||
}
|
||||
|
||||
# Allocate GPU slots
|
||||
slots = self.allocator.allocate_slots(allocation_specs)
|
||||
slot_map = {slot.assigned_model: slot for slot in slots}
|
||||
slots = self.allocator.allocate_slots(allocation_specs, preserve_order=True)
|
||||
slot_map = {s.assigned_model: s for s in slots}
|
||||
|
||||
if not slots:
|
||||
raise RuntimeError(
|
||||
f"Failed to allocate GPU slots for {num_workers} workers after eviction. "
|
||||
f"Need {required_gpus} GPUs."
|
||||
f"Failed to allocate GPU slots for {len(valid_workers)} workers"
|
||||
)
|
||||
|
||||
# Detect IB device for PD workers
|
||||
has_pd = any(w.is_prefill or w.is_decode for w in valid_workers)
|
||||
ib_device = detect_ib_device() if has_pd else None
|
||||
|
||||
# Track bootstrap ports for PD groups (shared within model/mode)
|
||||
pd_bootstrap_ports: dict[tuple[str, ConnectionMode], int] = {}
|
||||
|
||||
instances: list[ModelInstance] = []
|
||||
for w in valid_workers:
|
||||
# Get bootstrap port for PD workers
|
||||
bootstrap_port = None
|
||||
if w.is_prefill or w.is_decode:
|
||||
pd_key = (w.model_id, w.mode)
|
||||
if pd_key not in pd_bootstrap_ports:
|
||||
pd_bootstrap_ports[pd_key] = get_open_port()
|
||||
bootstrap_port = pd_bootstrap_ports[pd_key]
|
||||
|
||||
# Launch workers
|
||||
for i in range(num_workers):
|
||||
key = f"{model_id}:{mode.value}:{i}"
|
||||
gpu_slot = slot_map.get(key)
|
||||
instance = self._launch_model(
|
||||
model_id=model_id,
|
||||
mode=mode,
|
||||
gpu_slot=gpu_slot,
|
||||
worker_type=WorkerType.REGULAR,
|
||||
instance_key=key,
|
||||
model_id=w.model_id,
|
||||
mode=w.mode,
|
||||
gpu_slot=slot_map.get(w.key),
|
||||
worker_type=w.worker_type,
|
||||
bootstrap_port=bootstrap_port if w.is_prefill else None,
|
||||
ib_device=ib_device if (w.is_prefill or w.is_decode) else None,
|
||||
instance_key=w.key,
|
||||
)
|
||||
instances.append(instance)
|
||||
|
||||
# Wait for all to be healthy
|
||||
self._wait_all_healthy()
|
||||
|
||||
return instances
|
||||
|
||||
def launch_pd_workers(
|
||||
self,
|
||||
model_id: str,
|
||||
num_prefill: int = 1,
|
||||
num_decode: int = 1,
|
||||
mode: ConnectionMode = ConnectionMode.HTTP,
|
||||
startup_timeout: int = DEFAULT_STARTUP_TIMEOUT,
|
||||
allow_eviction: bool = True,
|
||||
) -> tuple[list[ModelInstance], list[ModelInstance]]:
|
||||
"""Launch prefill and decode workers for PD disaggregation.
|
||||
|
||||
Args:
|
||||
model_id: Model identifier from MODEL_SPECS.
|
||||
num_prefill: Number of prefill workers to launch. Defaults to 1.
|
||||
num_decode: Number of decode workers to launch. Defaults to 1.
|
||||
mode: Connection mode (HTTP or GRPC).
|
||||
startup_timeout: Timeout for workers to become healthy.
|
||||
allow_eviction: If True, evict MRU models to free GPUs. If False,
|
||||
return empty lists when not enough GPUs available.
|
||||
|
||||
Returns:
|
||||
Tuple of (prefill_instances, decode_instances).
|
||||
"""
|
||||
self._startup_timeout = startup_timeout
|
||||
|
||||
if model_id not in MODEL_SPECS:
|
||||
raise ValueError(f"Unknown model: {model_id}")
|
||||
|
||||
spec = get_model_spec(model_id)
|
||||
ib_device = detect_ib_device()
|
||||
if ib_device:
|
||||
logger.info("Detected InfiniBand device: %s", ib_device)
|
||||
|
||||
# Calculate total GPUs needed for PD workers
|
||||
tp = spec.get("tp", 1)
|
||||
required_gpus = (num_prefill + num_decode) * tp
|
||||
|
||||
# Check if we have enough GPUs
|
||||
available = self.allocator.available_gpus()
|
||||
if len(available) < required_gpus:
|
||||
if allow_eviction:
|
||||
logger.info(
|
||||
"Need %d GPUs for PD workers, only %d available. Evicting MRU models...",
|
||||
required_gpus,
|
||||
len(available),
|
||||
)
|
||||
# Exclude PD workers of same model/mode, but evict REGULAR workers
|
||||
self._evict_for_gpus(
|
||||
required_gpus,
|
||||
exclude_model_id=model_id,
|
||||
exclude_mode=mode,
|
||||
exclude_worker_types={WorkerType.PREFILL, WorkerType.DECODE},
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"Need %d GPUs for PD workers, only %d available. "
|
||||
"Skipping pre-launch (eviction not allowed).",
|
||||
required_gpus,
|
||||
len(available),
|
||||
)
|
||||
return [], []
|
||||
|
||||
# Build allocation specs for all PD workers
|
||||
# Each worker needs its own GPU slot
|
||||
allocation_specs = {}
|
||||
for i in range(num_prefill):
|
||||
key = f"{model_id}:{mode.value}:prefill_{i}"
|
||||
allocation_specs[key] = {
|
||||
"model": spec["model"],
|
||||
"memory_gb": spec.get("memory_gb", 16),
|
||||
"tp": tp,
|
||||
}
|
||||
for i in range(num_decode):
|
||||
key = f"{model_id}:{mode.value}:decode_{i}"
|
||||
allocation_specs[key] = {
|
||||
"model": spec["model"],
|
||||
"memory_gb": spec.get("memory_gb", 16),
|
||||
"tp": tp,
|
||||
}
|
||||
|
||||
# Allocate GPU slots
|
||||
slots = self.allocator.allocate_slots(allocation_specs)
|
||||
slot_map = {slot.assigned_model: slot for slot in slots}
|
||||
|
||||
if not slots:
|
||||
raise RuntimeError(
|
||||
f"Failed to allocate GPU slots for PD workers after eviction. "
|
||||
f"Need {required_gpus} GPUs."
|
||||
)
|
||||
|
||||
prefill_instances: list[ModelInstance] = []
|
||||
decode_instances: list[ModelInstance] = []
|
||||
|
||||
# Launch prefill workers
|
||||
for i in range(num_prefill):
|
||||
key = f"{model_id}:{mode.value}:prefill_{i}"
|
||||
gpu_slot = slot_map.get(key)
|
||||
bootstrap_port = get_open_port()
|
||||
instance = self._launch_model(
|
||||
model_id=model_id,
|
||||
mode=mode,
|
||||
gpu_slot=gpu_slot,
|
||||
worker_type=WorkerType.PREFILL,
|
||||
bootstrap_port=bootstrap_port,
|
||||
ib_device=ib_device,
|
||||
instance_key=key,
|
||||
)
|
||||
prefill_instances.append(instance)
|
||||
|
||||
# Launch decode workers
|
||||
for i in range(num_decode):
|
||||
key = f"{model_id}:{mode.value}:decode_{i}"
|
||||
gpu_slot = slot_map.get(key)
|
||||
instance = self._launch_model(
|
||||
model_id=model_id,
|
||||
mode=mode,
|
||||
gpu_slot=gpu_slot,
|
||||
worker_type=WorkerType.DECODE,
|
||||
ib_device=ib_device,
|
||||
instance_key=key,
|
||||
)
|
||||
decode_instances.append(instance)
|
||||
|
||||
# Wait for all to be healthy
|
||||
self._wait_all_healthy()
|
||||
|
||||
return prefill_instances, decode_instances
|
||||
|
||||
def get_client(
|
||||
self, model_id: str, mode: ConnectionMode | str = ConnectionMode.HTTP
|
||||
) -> "openai.OpenAI":
|
||||
|
||||
Reference in New Issue
Block a user