[smg][ci] preserve model launch order with test collected (#16618)

This commit is contained in:
Simo Lin
2026-01-07 06:16:59 -08:00
committed by GitHub
parent 4d902c8211
commit e432057381
6 changed files with 570 additions and 344 deletions
+217 -240
View File
@@ -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":