diff --git a/.gitignore b/.gitignore index 003992b..25c35b8 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,6 @@ IMPLEMENTATION_SUMMARY.md +docs/project_structure.md +docs/rollout_scheduling_improvements.md # Python and test caches __pycache__/ diff --git a/config.py b/config.py index e7e6297..40f1472 100644 --- a/config.py +++ b/config.py @@ -301,6 +301,56 @@ class SchedulingConfig: max_queue_length: int = 100 + # Load metric for least-connections style selection: "requests" (legacy + # request-count based), "tokens" (estimated in-flight token load, + # prompt + EMA-expected output), or "kv_tokens" (the same token load + # normalized by each TP bucket's estimated KV-cache capacity in tokens). + # Applies to load_balance / length_aware / la_mlfq schedulers; cmlfq + # keeps its own workload-aware routing. + load_metric: str = "requests" + + # Explicit KV capacity override (tokens) per TP degree, e.g. {1: 120000, + # 2: 230000, 4: 450000}. When empty and load_metric is "kv_tokens", + # capacities are auto-estimated at config-load time from the hardware / + # model-arch configs and the rollout gpu_memory_utilization. + kv_capacity_tokens_by_tp: dict = field(default_factory=dict) + + # Per-GPU bytes reserved for activations / workspace / fragmentation when + # auto-estimating KV capacity. + kv_activation_reserve_gib: float = 4.0 + + # Closed-loop /metrics feedback (engine-layer poller feeding the + # schedulers). Disabled by default; when enabled, observed vLLM gauges + # (running/waiting, gpu cache occupancy, preemptions) drive an additive + # bias correction of the local load estimate plus occupancy admission + # with hysteresis. Preemption counters are cumulative: events penalize + # an instance for the TTL window instead of blacklisting forever. + enable_metrics_feedback: bool = False + metrics_poll_interval_s: float = 3.0 + metrics_request_timeout_s: float = 1.0 + metrics_staleness_ttl_s: float = 10.0 + metrics_admission_enter: float = 0.90 + metrics_admission_exit: float = 0.75 + metrics_preemption_penalty_ttl_s: float = 60.0 + metrics_bias_alpha: float = 0.3 + + # Cross-rank load aggregation: every training rank publishes its + # per-instance in-flight accounting to this shared directory (atomic + # renames + heartbeat TTL); schedulers rank instances by the summed + # cluster load instead of their own partial view. Empty disables. + shared_load_dir: str = "" + shared_load_ttl_s: float = 30.0 + shared_load_heartbeat_s: float = 10.0 + shared_load_cache_ttl_s: float = 1.0 + + # Prefix affinity: bind prompt_id -> instance (bounded LRU) so repeated + # prompts (n_samples, multi-turn replays, shared system prompts) land + # on the instance whose prefix cache already holds them. Affinity only + # reorders preference among candidates that survived readiness / queue / + # capacity / admission filters; it never overrides them. + prefix_affinity: bool = False + prefix_affinity_max_entries: int = 8192 + enable_fallback: bool = True adaptive_routing: bool = True @@ -950,6 +1000,57 @@ def __post_init__(self): else: self.master_addr = "localhost" + + self._maybe_estimate_kv_capacities() + + + def _maybe_estimate_kv_capacities(self) -> None: + """Auto-fill scheduling.kv_capacity_tokens_by_tp when needed. + + Only runs for heterogeneous rollouts using the capacity-normalized + ``kv_tokens`` load metric without an explicit capacity table. The + estimate follows vLLM's budget rule: per GPU, the utilization + fraction of HBM minus sharded weights and an activation reserve, + divided by the per-token KV bytes of the bucket's TP layout. + """ + sched = self.heterogeneous_rollout.scheduling + if not self.heterogeneous_rollout.enabled: + return + if getattr(sched, "load_metric", "") != "kv_tokens": + return + if sched.kv_capacity_tokens_by_tp: + return + + from RL_Framework.infra.scheduling.base import estimate_kv_capacity_tokens + + arch = self.model_arch + weights_bytes = float(arch.num_params) * max(1, int(arch.dtype_bytes)) + reserve = max(0.0, float(sched.kv_activation_reserve_gib)) * 1024**3 + tp_degrees = sorted({int(i.tp) for i in self.heterogeneous_rollout.instances if int(i.tp) > 0}) + if not tp_degrees: + return + + capacities: dict[int, int] = {} + for tp in tp_degrees: + capacity = estimate_kv_capacity_tokens( + tp_degree=tp, + mem_capacity_bytes=self.hardware.mem_capacity, + gpu_memory_utilization=self.heterogeneous_rollout.gpu_memory_utilization, + weights_bytes=weights_bytes, + n_layers=arch.n_layers, + n_kv_heads=arch.n_kv_heads, + n_heads=arch.n_heads, + d_model=arch.d_model, + head_dim=arch.head_dim, + dtype_bytes=arch.dtype_bytes, + activation_reserve_bytes=reserve, + ) + if capacity > 0: + capacities[tp] = capacity + + if capacities: + sched.kv_capacity_tokens_by_tp = capacities + # ---------------------------------------------------------------- # ---------------------------------------------------------------- diff --git a/configs/scheduling_adaptive.yaml b/configs/scheduling_adaptive.yaml new file mode 100644 index 0000000..3dcd648 --- /dev/null +++ b/configs/scheduling_adaptive.yaml @@ -0,0 +1,20 @@ +# Portable rollout-scheduling preset. Model, endpoints and GPU placement +# belong to the deployment config or the validation command, not this file. +# Example (existing vLLM services): +# python examples/validate_rollout_scheduling.py --config configs/scheduling_adaptive.yaml \ +# --model --endpoint ,, \ +# --endpoint ,, --output +# To use for training: inherit this file via base_config from a full experiment +# config, and supply model_path and heterogeneous_rollout.instances there. +heterogeneous_rollout: + enabled: true + scheduling: + scheduler_type: load_balance + load_balance_strategy: least_connections + load_metric: kv_tokens + enable_metrics_feedback: true + prefix_affinity: true + # For multi-rank runs only, use a unique run directory visible to ALL ranks: + # shared_load_dir: //rollout-load + # Measured per-instance capacities from /metrics take precedence over + # analytic estimates; the latter require correct hardware/model_arch inputs. diff --git a/docs/configuration_reference.md b/docs/configuration_reference.md index e5e2e8d..b0f7d0b 100644 --- a/docs/configuration_reference.md +++ b/docs/configuration_reference.md @@ -78,6 +78,66 @@ phase names allow the same converter to compare baseline runs. See the [cluster manual](manual.md) for recommended combinations and launch examples. +## Rollout Scheduling Options + +These options live under `heterogeneous_rollout.scheduling` and control how +requests are routed across the heterogeneous TP buckets. + +### Load metric + +| Option | Meaning | +| --- | --- | +| `load_metric` | Selection signal for least-connections routing: `requests` (legacy request-count, default), `tokens` (estimated in-flight token load = prompt + EMA-expected output), or `kv_tokens` (that load normalized by each TP bucket's KV-cache capacity, i.e. occupancy ratio) | +| `kv_capacity_tokens_by_tp` | Explicit KV capacity per TP degree, e.g. `{1: 120000, 4: 450000}`. When empty and `load_metric: kv_tokens`, capacities are auto-estimated from the hardware / model-arch configs and the rollout `gpu_memory_utilization`; at engine startup they are additionally calibrated against vLLM's own profiled `kv_cache_size_tokens` from `/metrics` | +| `kv_activation_reserve_gib` | Per-GPU reserve subtracted when auto-estimating KV capacity (default 4) | +| `load_balance_strategy` | `least_connections` (default), `round_robin`, or `weighted` | + +Under `kv_tokens`, instances whose TP degree has no configured capacity are +excluded from selection (with a warning) rather than comparing raw token +counts against ratios. All strategies rotate among equal-load endpoints to +avoid starving later-registered instances. + +The expected-output component of the estimate is a per-category EMA learned +from observed completion lengths; it survives weight-sync rebinding, and is +clamped by each request's `max_new_tokens` (callers that pass exact +`input_tokens` — all bundled workflows do — skip the chars-based fallback). + +### Closed-loop /metrics feedback + +| Option | Meaning | +| --- | --- | +| `enable_metrics_feedback` | Master switch (default false). An engine-layer poller scrapes each instance's `/metrics`; observed gauges drive an additive bias correction of the local load estimate plus occupancy admission | +| `metrics_poll_interval_s` | Scrape interval per instance (default 3; measured interference is <1% at 5Hz) | +| `metrics_request_timeout_s` | Per-scrape HTTP timeout (default 1) | +| `metrics_staleness_ttl_s` | Snapshots older than this are ignored (default 10) | +| `metrics_admission_enter` / `metrics_admission_exit` | Occupancy hysteresis band: block new routes above `enter` (default 0.90), unblock below `exit` (default 0.75) | +| `metrics_preemption_penalty_ttl_s` | Preemption counters are cumulative; a detected increase penalizes the instance for this window (default 60) instead of blacklisting forever | +| `metrics_bias_alpha` | EMA weight for the additive bias `bias = EMA(observed - local)` (default 0.3) | + +The poller runs at the engine layer (it survives scheduler rebinding during +weight sync), staggers endpoints with random phases (multi-rank safety), and +any feed outage degrades transparently back to the open-loop estimate. + +### Cross-rank load aggregation + +| Option | Meaning | +| --- | --- | +| `shared_load_dir` | Shared directory where every training rank publishes its per-instance in-flight accounting (atomic renames + heartbeat TTL). Schedulers then rank instances by the summed cluster load instead of their own partial view, eliminating multi-rank herd. Empty disables | +| `shared_load_ttl_s` / `shared_load_heartbeat_s` | Liveness window / refresh rate for rank files (defaults 30 / 10) | +| `shared_load_cache_ttl_s` | Aggregate-read cache window (default 1); a rank's own writes invalidate immediately | + +### Prefix affinity + +| Option | Meaning | +| --- | --- | +| `prefix_affinity` | Bind `prompt_id` -> instance (bounded LRU) so repeated prompts (GRPO `n_samples`, multi-turn replays, shared system prompts) land on the instance whose prefix cache already holds them (default false) | +| `prefix_affinity_max_entries` | LRU bound (default 8192) | + +Affinity only reorders preference among candidates that survived the +readiness / queue / capacity / admission filters — a sticky instance at its +queue cap or blocked by occupancy admission is bypassed and the mapping +remaps. The affinity table survives weight-sync rebinding. + ## Rollout Weight Reloads When `rollout_weight_sync_mode` is `restart`, the trainer publishes one reload diff --git a/engine/heterogeneous_engine.py b/engine/heterogeneous_engine.py index ff89a1c..2677b51 100644 --- a/engine/heterogeneous_engine.py +++ b/engine/heterogeneous_engine.py @@ -47,10 +47,140 @@ def __init__( self._pending_futures: dict[str, list[asyncio.Future]] = {} self._cmlfq_backend = CMLFQGenerationBackend() + # Engine-owned telemetry. Whole-engine replacement stops these + # resources before creating new ones; they are not global singletons. + self._metrics_poller: Any = None + # Engine-layer cross-rank load publisher/aggregator (same + # lifetime rules as the poller). + self._shared_state: Any = None + # ---------------------------------------------------------------- # ---------------------------------------------------------------- + def _setup_metrics_polling(self, hetero_cfg: Any) -> None: + """Create/replace/stop the /metrics poller per scheduling config.""" + from RL_Framework.infra.scheduling.base import MetricsFeedbackConfig + from RL_Framework.infra.scheduling.metrics_feed import VLLMMetricsPoller + + feedback = MetricsFeedbackConfig.from_scheduling( + getattr(hetero_cfg, "scheduling", None) + ) + if not feedback.enabled: + self._stop_metrics_poller() + return + self._stop_metrics_poller() + poller = VLLMMetricsPoller( + interval_s=feedback.poll_interval_s, + timeout_s=feedback.request_timeout_s, + ttl_s=feedback.staleness_ttl_s, + ) + self._metrics_poller = poller + self._refresh_metrics_endpoints() + poller.start() + if self.scheduler is not None: + self.scheduler.attach_metrics_feed(poller) + + def _refresh_metrics_endpoints(self) -> None: + if self._metrics_poller is None: + return + urls = { + cfg["instance_id"]: f"http://{cfg['host']}:{cfg['port']}/metrics" + for cfg in self.instance_configs + } + self._metrics_poller.set_endpoints(urls) + + def _calibrate_capacities_from_metrics(self) -> None: + """Retain fresh profiled capacities per instance, never max by TP. + + Same-TP endpoints can have different budgets. The scheduler uses + fresh snapshots first, then this last profiled value, then the + explicitly configured/analytic per-TP estimate. + """ + if self._metrics_poller is None or self.scheduler is None: + return + calibrated: dict[str, int] = {} + for cfg, snap in ( + (cfg, snap) + for cfg in self.instance_configs + for snap in [self._metrics_poller.get(cfg["instance_id"])] + if snap is not None + ): + capacity = int(getattr(snap, "kv_capacity_tokens", -1)) + if capacity > 0: + calibrated[cfg["instance_id"]] = capacity + with self.scheduler._lock: + for handle in self.scheduler._instances: + if handle.instance_id in calibrated: + handle.kv_capacity_tokens = calibrated[handle.instance_id] + if calibrated: + logger.info("[MetricsFeedback] profiled capacity by instance: %s", calibrated) + + def _stop_metrics_poller(self) -> None: + if self._metrics_poller is not None: + try: + self._metrics_poller.stop() + except Exception as exc: + logger.warning("Failed stopping metrics poller: %s", exc) + self._metrics_poller = None + self.scheduler.attach_metrics_feed(None) + + def _setup_shared_state(self, hetero_cfg: Any) -> None: + """Create the cross-rank load state once per engine lifetime.""" + sched = getattr(hetero_cfg, "scheduling", None) + directory = str(getattr(sched, "shared_load_dir", "") or "") + if not directory: + return + if hasattr(self.scheduler, "get_request_route"): + logger.warning("C-MLFQ uses cmlfq_shared_load_dir, ignoring shared_load_dir") + return + from RL_Framework.infra.scheduling.shared_token_state import ( + SharedTokenLoadState, + ) + + self._shared_state = SharedTokenLoadState( + directory=directory, + ttl_s=float(getattr(sched, "shared_load_ttl_s", 30.0)), + heartbeat_interval_s=float( + getattr(sched, "shared_load_heartbeat_s", 10.0) + ), + cache_ttl_s=float(getattr(sched, "shared_load_cache_ttl_s", 1.0)), + ) + if self.scheduler is not None: + self.scheduler.attach_shared_state(self._shared_state) + logger.info( + "Attached cross-rank shared load state: dir=%s writer=%s", + directory, + self._shared_state.writer_id, + ) + + def _reattach_shared_state(self) -> None: + """Wire the existing shared state into a freshly built scheduler. + + The new scheduler's local counters start at zero, so this rank's + published totals must be zeroed to match (other ranks untouched). + """ + if self._shared_state is None or self.scheduler is None: + return + try: + self._shared_state.reset() + except Exception as exc: + logger.warning("Shared-state reset on reattach failed: %s", exc) + self.scheduler.attach_shared_state(self._shared_state) + + def _close_shared_state(self) -> None: + if self._shared_state is not None: + try: + self._shared_state.close() + except Exception as exc: + logger.warning("Failed closing shared load state: %s", exc) + self._shared_state = None + + def metrics_snapshots(self) -> dict[str, Any]: + if self._metrics_poller is None: + return {} + return self._metrics_poller.snapshots() + def add_instance( self, instance_id: str, @@ -74,6 +204,7 @@ def add_instance( "tp_degree": tp_degree, "gpu_ids": gpu_ids or [], }) + self._refresh_metrics_endpoints() self.scheduler.register_instance( @@ -109,6 +240,10 @@ def reconfigure_from_plan(self, plan: Any, config: Any): start or stop workers before this method is called. """ with self._lock: + if any(h.active_requests for h in self.scheduler._instances) or any( + not f.done() for futures in self._pending_futures.values() for f in futures + ): + raise RuntimeError("Drain rollout requests before reconfiguring the engine") for engine in self.engines: if hasattr(engine, "close_sync"): engine.close_sync() @@ -119,6 +254,16 @@ def reconfigure_from_plan(self, plan: Any, config: Any): scheduler_type=scheduler_type, hetero_config=hetero, ) + old_scheduler = self.scheduler + if old_scheduler is not None and old_scheduler is not scheduler: + try: + scheduler.import_learned_state( + old_scheduler.export_learned_state() + ) + except Exception as exc: + logger.warning( + "Failed to carry scheduler state across reconfigure: %s", exc + ) self.scheduler = scheduler self.engines = [] @@ -147,6 +292,9 @@ def reconfigure_from_plan(self, plan: Any, config: Any): self.num_instances, ) + self._setup_metrics_polling(hetero) + self._reattach_shared_state() + # ---------------------------------------------------------------- # ---------------------------------------------------------------- @@ -165,6 +313,17 @@ def wait_for_ready(self, timeout: float = 300.0): f"All {self.num_instances} heterogeneous instances are ready: " f"TP layout={self.tp_list}" ) + # Give the metrics poller a moment to capture its first snapshots, + # then calibrate KV capacities against vLLM's own profiling. + if self._metrics_poller is not None: + import time as _time + + deadline = _time.time() + 15.0 + while _time.time() < deadline: + if len(self.metrics_snapshots()) >= self.num_instances: + break + _time.sleep(0.5) + self._calibrate_capacities_from_metrics() def wait_until_idle(self, timeout: float = 3600.0, poll_interval: float = 0.5): """Block until no rollout requests are active before reconfiguration.""" @@ -190,10 +349,20 @@ def wait_until_idle(self, timeout: float = 3600.0, poll_interval: float = 0.5): async def close(self): """Close.""" + self._stop_metrics_poller() + self._close_shared_state() await self._cmlfq_backend.close() for engine in self.engines: await engine.close() + def close_sync(self): + """Synchronous close used when the engine is replaced at rebind.""" + self._stop_metrics_poller() + self._close_shared_state() + for engine in self.engines: + if hasattr(engine, "close_sync"): + engine.close_sync() + # ---------------------------------------------------------------- # ---------------------------------------------------------------- @@ -234,10 +403,15 @@ async def generate( ) if input_tokens <= 0: - input_tokens = max(1, len(prompt) // 3) + # Cheap chars-based fallback only: every bundled workflow now + # passes its exact token count, so this path serves unknown + # future callers. ~4 chars/token is a better prior than the + # legacy //3 (which overestimated English ~40%). + input_tokens = max(1, int(len(prompt) / 4)) with self._lock: + route_scheduler = self.scheduler requested_cmlfq_route = bool( request_id and hasattr(self.scheduler, "get_request_route") ) @@ -253,6 +427,7 @@ async def generate( prompt_id=prompt_id, n_samples=n_samples, epoch=epoch, + max_new_tokens=max_new_tokens, ) @@ -262,12 +437,23 @@ async def generate( input_tokens=input_tokens, n_samples=n_samples, epoch=epoch, + max_new_tokens=max_new_tokens, ) with self._lock: - if result.instance_index < 0: - idx = self._rr_counter % max(1, self.num_instances) + scheduled_by_scheduler = result.instance_index >= 0 + if not scheduled_by_scheduler: + if hasattr(route_scheduler, "finish_request"): + raise RuntimeError(f"C-MLFQ routing failed: {result.reason}") + with route_scheduler._lock: + ready = route_scheduler._selectable([ + h for h in route_scheduler._instances if h.is_ready + ]) + if not ready: + raise RuntimeError(f"No ready rollout instance: {result.reason}") + idx = ready[self._rr_counter % len(ready)].index self._rr_counter += 1 + result.is_fallback = True logger.warning( f"Scheduling failed ({result.reason}), falling back to instance {idx}" ) @@ -275,6 +461,27 @@ async def generate( idx = result.instance_index engine = self.engines[idx] instance_config = dict(self.instance_configs[idx]) + # The scheduler never accounted a fallback route (its schedule() + # call failed), so debit it here. Otherwise the completion path + # would dec a request that was never inc'd and silently steal + # tokens from whichever request is still in flight on idx. + if not scheduled_by_scheduler and not cmlfq_managed: + with_signal = getattr( + self.scheduler, "_record_route", None + ) + if callable(with_signal): + with self.scheduler._lock: + result.reserved_tokens = with_signal( + self.scheduler.get_instance_handle(idx), + input_tokens, + category=result.category or "any", + prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) + else: + handle = self.scheduler.get_instance_handle(idx) + if handle is not None: + handle.inc_active() output_tokens = 0 try: @@ -305,20 +512,24 @@ async def generate( if not cmlfq_managed: if result.request_id and hasattr( - self.scheduler, "finish_request" + route_scheduler, "finish_request" ): with self._lock: - self.scheduler.finish_request( + route_scheduler.finish_request( result.request_id, output_tokens, ) else: with self._lock: - self.scheduler.on_request_done( + completion_kwargs = {} + if result.reserved_tokens is not None: + completion_kwargs["reserved_tokens"] = result.reserved_tokens + route_scheduler.on_request_done( instance_index=idx, prompt_id=prompt_id, final_bucket=result.category, output_tokens=output_tokens, + **completion_kwargs, ) def _configure_cost_runtime(self): @@ -417,10 +628,12 @@ async def _wait_for_scout( n_samples: int, epoch: int, timeout: float = 60.0, + max_new_tokens: int = 0, ) -> SchedulingResult: """Wait for scout.""" loop = asyncio.get_event_loop() future = loop.create_future() + self._pending_futures.setdefault(prompt_id, []).append(future) from RL_Framework.infra.scheduling.la_mlfq import WaitingRequest @@ -432,6 +645,7 @@ async def _wait_for_scout( epoch=epoch, sample_index=-1, future=future, + max_new_tokens=max_new_tokens, ) self.scheduler.scout_manager.add_waiting(wr) @@ -449,6 +663,12 @@ async def _wait_for_scout( f"Error while waiting for scout: prompt={prompt_id}, error={e}, " f"using default routing" ) + finally: + futures = self._pending_futures.get(prompt_id, []) + if future in futures: + futures.remove(future) + if not futures: + self._pending_futures.pop(prompt_id, None) return self.scheduler.schedule( @@ -456,6 +676,7 @@ async def _wait_for_scout( prompt_id="", n_samples=1, epoch=epoch, + max_new_tokens=max_new_tokens, ) async def generate_batch( @@ -495,7 +716,7 @@ async def generate_batch( valid_results = [] for i, result in enumerate(results): - if isinstance(result, Exception): + if isinstance(result, BaseException): logger.warning(f"Prompt {i} Generation failed: {result}") continue valid_results.append(result) @@ -631,8 +852,18 @@ def get_cluster_info(self) -> dict[str, Any]: # ---------------------------------------------------------------- @classmethod - def from_config(cls, config) -> "HeterogeneousRolloutEngine": - """From config.""" + def from_config( + cls, + config, + carry_scheduler_state_from: "BaseScheduler | None" = None, + ) -> "HeterogeneousRolloutEngine": + """From config. + + ``carry_scheduler_state_from`` transfers learned scheduler state + (output-length EMA, history tables) across a rebind. Weight sync + rebinds rebuild the engine every sync step; without the carry the + EMA would reset to its prior every step and never converge. + """ hetero = config.heterogeneous_rollout @@ -641,6 +872,20 @@ def from_config(cls, config) -> "HeterogeneousRolloutEngine": scheduler_type=scheduler_type, hetero_config=hetero, ) + if carry_scheduler_state_from is not None: + try: + scheduler.import_learned_state( + carry_scheduler_state_from.export_learned_state() + ) + logger.info( + "Carried learned scheduler state (%s -> %s) across rebind", + type(carry_scheduler_state_from).__name__, + type(scheduler).__name__, + ) + except Exception as exc: + logger.warning( + "Failed to carry scheduler state across rebind: %s", exc + ) engine = cls( @@ -687,6 +932,8 @@ def from_config(cls, config) -> "HeterogeneousRolloutEngine": f"Created heterogeneous engine from configuration: {engine.num_instances} instances, " f"TP layout={engine.tp_list}, scheduler={scheduler_type}" ) + engine._setup_metrics_polling(hetero) + engine._setup_shared_state(hetero) return engine # ---------------------------------------------------------------- diff --git a/examples/validate_rollout_scheduling.py b/examples/validate_rollout_scheduling.py new file mode 100644 index 0000000..d7268fb --- /dev/null +++ b/examples/validate_rollout_scheduling.py @@ -0,0 +1,137 @@ +"""Live HTTP validation; no fake lengths, injected loads or manual draining. + +Deployment (model, instance IDs, TP sizes and endpoints) is supplied at runtime. +The output records both route decisions and real vLLM cache-counter deltas. +This is a correctness smoke, not a throughput benchmark or training run. +""" +import argparse +import asyncio +import copy +import json +from pathlib import Path +from urllib.parse import urlparse + +import httpx +import yaml + +from RL_Framework.config import AsyncRLConfig +from RL_Framework.engine.heterogeneous_engine import HeterogeneousRolloutEngine +from RL_Framework.infra.scheduling.metrics_feed import parse_prometheus_metrics + + +async def validate(args): + raw = yaml.safe_load(Path(args.config).read_text(encoding="utf-8")) + raw["model_path"] = args.model + instances = [] + for endpoint in args.endpoint: + iid, tp, address = endpoint.split(",", 2) + url = urlparse(address) + if url.scheme != "http" or not url.hostname or not url.port: + raise ValueError("endpoint must be instance-id,tp,http://host:port") + instances.append(dict(instance_id=iid, tp=int(tp), host=url.hostname, port=url.port)) + if len(instances) < 2: + raise ValueError("At least two independent vLLM instances are required") + raw["heterogeneous_rollout"]["instances"] = instances + if args.shared_load_dir: + raw["heterogeneous_rollout"]["scheduling"]["shared_load_dir"] = args.shared_load_dir + config = AsyncRLConfig.from_dict(copy.deepcopy(raw)) + engine = HeterogeneousRolloutEngine.from_config(config) + scheduler = engine.scheduler + report = {"deployment": instances, "model": args.model, "phases": []} + tasks = [] + async with httpx.AsyncClient(timeout=30) as client: + async def counters(): + values = {} + for cfg in engine.instance_configs: + url = f"http://{cfg['host']}:{cfg['port']}/metrics" + response = await client.get(url) + response.raise_for_status() + all_values = parse_prometheus_metrics(response.text) + values[cfg["instance_id"]] = { + k: v for k, v in all_values.items() + if k in ("vllm:prefix_cache_hits_total", "vllm:prefix_cache_queries_total") + } + return values + + lengths = {} + async def generate(prompt, pid, budget=32): + if prompt not in lengths: + response = await client.post(engine.engines[0].base_url + "/tokenize", json={ + "model": args.model, "prompt": prompt, "add_special_tokens": False, + }) + response.raise_for_status() + payload = response.json() + lengths[prompt] = payload.get("count", len(payload.get("tokens", []))) + if lengths[prompt] <= 0: + raise AssertionError("/tokenize did not return a positive token count") + result = await engine.generate( + prompt, input_tokens=lengths[prompt], prompt_id=pid, + max_new_tokens=budget, temperature=0.0, + ) + if not result.get("tokens"): + raise AssertionError("Generation returned no token logprobs") + return {"input_tokens": lengths[prompt], "output_tokens": len(result["tokens"]), + "route": result["_schedule_info"]} + + try: + await asyncio.to_thread(engine.wait_for_ready, 180) + snapshots = engine.metrics_snapshots() + report["snapshots"] = {key: vars(value) for key, value in snapshots.items()} + for handle in scheduler._instances: + snap = engine._metrics_poller.get(handle.instance_id) + if snap is None or snap.kv_capacity_tokens <= 0: + raise AssertionError(f"No fresh measured capacity for {handle.instance_id}") + if scheduler.kv_capacity_of(handle) != snap.kv_capacity_tokens: + raise AssertionError("Measured capacity is not used by scheduler") + + prompt = ("Systems share compute and memory. " * 96) + "Summarize this in one sentence." + for enabled in (False, True): + scheduler._prefix_affinity_enabled = enabled + scheduler._affinity.clear() + before = await counters() + rows = [await generate(prompt, "repeated-prompt") for _ in range(4)] + after = await counters() + if enabled: + targets = {row["route"]["instance_id"] for row in rows} + if len(targets) != 1 or any( + row["route"]["reason"] != "prefix_affinity" for row in rows[1:] + ): + raise AssertionError("Repeated prompt did not follow affinity") + report["phases"].append({"affinity": enabled, "requests": rows, + "cache_before": before, "cache_after": after}) + + tasks = [asyncio.create_task(generate( + ("Read this. " * (30 + 100 * i)) + "Summarize briefly.", f"mixed-{i}", 64 + )) for i in range(4)] + report["mixed_requests"] = await asyncio.gather(*tasks) + for handle in scheduler._instances: + if handle.active_requests != 0 or handle.active_tokens != 0: + raise AssertionError(f"Leaked accounting: {handle}") + old_state = engine._shared_state + engine.reconfigure_from_plan(None, config) + if old_state is not None and engine._shared_state is not old_state: + raise AssertionError("Shared owner was not reused after a drained reconfigure") + report["after_reconfigure"] = await generate("Say hello.", "post-reconfigure", 8) + report["passed"] = True + finally: + for task in tasks: + if not task.done(): + task.cancel() + if tasks: + await asyncio.gather(*tasks, return_exceptions=True) + await engine.close() + if args.output: + output = Path(args.output) + output.parent.mkdir(parents=True, exist_ok=True) + output.write_text(json.dumps(report, indent=2), encoding="utf-8") + print(json.dumps(report, indent=2)) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--config", default="configs/scheduling_adaptive.yaml") + parser.add_argument("--model", required=True) + parser.add_argument("--endpoint", action="append", required=True) + parser.add_argument("--shared-load-dir", default="") + parser.add_argument("--output") + asyncio.run(validate(parser.parse_args())) diff --git a/infra/scheduling/base.py b/infra/scheduling/base.py index eddf839..16cb837 100644 --- a/infra/scheduling/base.py +++ b/infra/scheduling/base.py @@ -1,9 +1,11 @@ """Support code for Base.""" import logging +import math import threading +import time from abc import ABC, abstractmethod -from collections import defaultdict +from collections import defaultdict, deque, OrderedDict from dataclasses import dataclass, field from enum import Enum from typing import Any, Optional @@ -11,6 +13,93 @@ logger = logging.getLogger(__name__) +DEFAULT_LOAD_METRIC = "requests" +VALID_LOAD_METRICS = ("requests", "tokens", "kv_tokens") +DEFAULT_EXPECTED_OUTPUT_TOKENS = 512 +OUTPUT_EMA_ALPHA = 0.2 +DEFAULT_KV_ACTIVATION_RESERVE_BYTES = 4 * 1024**3 + + +@dataclass +class MetricsFeedbackConfig: + """Closed-loop /metrics feedback settings for schedulers. + + ``enabled`` defaults off so existing configs keep the pure open-loop + behavior. When on, engine-layer poller snapshots drive (a) an additive + bias correction of the local load estimate and (b) occupancy admission + with hysteresis; preemption events penalize an instance for a TTL + window instead of blacklisting it forever (the counter is cumulative). + """ + + enabled: bool = False + admission_enter: float = 0.90 + admission_exit: float = 0.75 + preemption_penalty_ttl_s: float = 60.0 + poll_interval_s: float = 3.0 + request_timeout_s: float = 1.0 + staleness_ttl_s: float = 10.0 + bias_alpha: float = 0.3 + + def __post_init__(self): + if not 0 <= self.admission_exit < self.admission_enter <= 1: + raise ValueError("metrics admission requires 0 <= exit < enter <= 1") + if not 0 < self.bias_alpha <= 1: + raise ValueError("metrics_bias_alpha must be in (0, 1]") + for name in ("poll_interval_s", "request_timeout_s", "staleness_ttl_s"): + if not math.isfinite(getattr(self, name)) or getattr(self, name) <= 0: + raise ValueError(f"{name} must be finite and positive") + + @classmethod + def from_scheduling(cls, sched: Any) -> "MetricsFeedbackConfig": + get = lambda name, default: getattr(sched, name, default) if sched is not None else default + return cls( + enabled=bool(get("enable_metrics_feedback", False)), + admission_enter=float(get("metrics_admission_enter", 0.90)), + admission_exit=float(get("metrics_admission_exit", 0.75)), + preemption_penalty_ttl_s=float(get("metrics_preemption_penalty_ttl_s", 60.0)), + poll_interval_s=float(get("metrics_poll_interval_s", 3.0)), + request_timeout_s=float(get("metrics_request_timeout_s", 1.0)), + staleness_ttl_s=float(get("metrics_staleness_ttl_s", 10.0)), + bias_alpha=float(get("metrics_bias_alpha", 0.3)), + ) + + +def estimate_kv_capacity_tokens( + *, + tp_degree: int, + mem_capacity_bytes: float, + gpu_memory_utilization: float, + weights_bytes: float, + n_layers: int, + n_kv_heads: int, + n_heads: int, + d_model: int, + head_dim: int = 0, + dtype_bytes: int = 2, + activation_reserve_bytes: float = DEFAULT_KV_ACTIVATION_RESERVE_BYTES, +) -> int: + """Estimate how many tokens of KV cache a TP bucket can hold. + + Per GPU: ``mem_capacity * gpu_memory_utilization`` is the vLLM budget, + minus the sharded weights and a fixed activation/workspace reserve. + KV cost per token per GPU is ``2 (K+V) * n_layers * kv_heads_per_gpu * + head_dim * dtype_bytes``; GQA replicates KV heads when the TP degree + exceeds the number of KV heads. + """ + tp = max(1, int(tp_degree)) + dim = int(head_dim) or max(1, int(d_model) // max(1, int(n_heads))) + kv_heads_per_gpu = max(1, -(-max(1, int(n_kv_heads)) // tp)) + kv_bytes_per_token = 2 * max(1, int(n_layers)) * kv_heads_per_gpu * dim * max(1, int(dtype_bytes)) + budget = ( + float(mem_capacity_bytes) * min(1.0, max(0.05, float(gpu_memory_utilization))) + - float(weights_bytes) / tp + - float(activation_reserve_bytes) + ) + if kv_bytes_per_token <= 0 or budget <= 0: + return 0 + return int(budget // kv_bytes_per_token) + + # --------------------------------------------------------------------------- # --------------------------------------------------------------------------- @@ -44,6 +133,7 @@ class SchedulingResult: pending: bool = False prompt_id: str = "" request_id: str = "" + reserved_tokens: int | None = None @dataclass @@ -95,12 +185,71 @@ class InstanceHandle: tp_degree: int is_ready: bool = True active_requests: int = 0 + active_tokens: int = 0 + kv_capacity_tokens: int = 0 + _token_estimates: deque = field(default_factory=deque) - def inc_active(self): - self.active_requests += 1 + def inc_active( + self, + prompt_tokens: int = 0, + expected_output_tokens: int = 0, + ): + """Account a routed request. - def dec_active(self): + ``active_requests`` counts in-flight requests. ``active_tokens`` + tracks the estimated in-flight token load (prompt + expected + output). Each estimated request pushes its estimate onto a FIFO + queue so the exact same amount is subtracted at completion. + """ + self.active_requests += 1 + estimate = max(0, int(prompt_tokens)) + max(0, int(expected_output_tokens)) + self.active_tokens += estimate + self._token_estimates.append(estimate) + + def dec_active(self, reserved_tokens: int | None = None) -> int: + """Retire a completed request. + + Subtracts only what was added at ``inc_active`` time: the oldest + pending token estimate, if any. Keeps the invariant + ``active_tokens == sum(_token_estimates)`` so the counter returns + exactly to zero once all in-flight requests drain. Returns the + token estimate actually retired (0 when the queue was empty), so + callers can mirror the exact delta into a shared cross-rank state. + """ + if reserved_tokens is not None: + # Completion order is not submission order. Retire this request's + # reservation, not a still-running long request at the FIFO head. + self._token_estimates.remove(reserved_tokens) + self.active_tokens -= reserved_tokens + self.active_requests = max(0, self.active_requests - 1) + return reserved_tokens self.active_requests = max(0, self.active_requests - 1) + retired = 0 + if self._token_estimates: + estimate = self._token_estimates.popleft() + retired = estimate + self.active_tokens = max(0, self.active_tokens - estimate) + return retired + + def load(self, metric: str = DEFAULT_LOAD_METRIC, kv_capacity_tokens: int = 0) -> float: + """Current load under the requested metric. + + ``kv_tokens`` returns the occupancy ratio (0..1+) of estimated + in-flight tokens against the bucket's KV capacity. Callers must + NOT compare ratios against raw token counts (unit mismatch); + schedulers filter uncapacitated instances out of ``kv_tokens`` + candidate sets via ``is_kv_capacitated``. + """ + if metric == "kv_tokens": + return self.active_tokens / max(1, int(kv_capacity_tokens)) + if metric == "tokens": + return float(self.active_tokens) + return float(self.active_requests) + + @property + def is_kv_capacitated(self) -> bool: + """Whether a KV capacity has been assigned to this handle.""" + return self.kv_capacity_tokens > 0 # --------------------------------------------------------------------------- @@ -110,12 +259,337 @@ def dec_active(self): class BaseScheduler(ABC): """Base scheduler implementation.""" - def __init__(self, name: str = "base"): + def __init__( + self, + name: str = "base", + load_metric: str = DEFAULT_LOAD_METRIC, + kv_capacity_tokens_by_tp: dict[int, int] | None = None, + feedback: MetricsFeedbackConfig | None = None, + prefix_affinity: bool = False, + prefix_affinity_max_entries: int = 8192, + ): self.name = name + if load_metric not in VALID_LOAD_METRICS: + logger.warning( + "[%s] unknown load_metric '%s', falling back to '%s'", + name, load_metric, DEFAULT_LOAD_METRIC, + ) + load_metric = DEFAULT_LOAD_METRIC + self._load_metric = load_metric + self._kv_capacity_by_tp: dict[int, int] = { + int(tp): int(cap) + for tp, cap in (kv_capacity_tokens_by_tp or {}).items() + if int(cap) > 0 + } + if load_metric == "kv_tokens" and not self._kv_capacity_by_tp: + logger.warning( + "[%s] load_metric='kv_tokens' without kv_capacity_tokens_by_tp; " + "falling back to unnormalized token counts", + name, + ) + self._feedback = feedback or MetricsFeedbackConfig() + self._metrics_feed: Any = None + self._bias_ema: dict[str, float] = {} + self._bias_seen_ts: dict[str, float] = {} + self._admission_blocked: dict[str, bool] = {} + self._shared_state: Any = None + self._capacity_seen: dict[str, int] = {} + # Prefix affinity: prompt_id -> instance_id, bounded LRU. RL + # rollouts repeat prompts (n_samples, multi-turn replays, shared + # system prompts); keeping them on one instance lets vLLM's prefix + # cache serve the repeated prefill instead of recomputing it on a + # cold instance. + self._prefix_affinity_enabled = bool(prefix_affinity) + self._prefix_affinity_max_entries = max(64, int(prefix_affinity_max_entries)) + self._affinity: OrderedDict[str, str] = OrderedDict() self._instances: list[InstanceHandle] = [] self._instances_by_tp: dict[int, list[InstanceHandle]] = defaultdict(list) self._lock = threading.Lock() self._stats = SchedulerStats() + self._tie_rr_counter = 0 + # Per-category EMA of observed output lengths, used to estimate the + # output-token share of in-flight requests at routing time. + self._output_ema: dict[str, float] = defaultdict( + lambda: float(DEFAULT_EXPECTED_OUTPUT_TOKENS) + ) + + # ---------------------------------------------------------------- + + # ---------------------------------------------------------------- + + def kv_capacity_of(self, handle: InstanceHandle) -> int: + """Configured KV capacity (tokens) for a handle's TP bucket.""" + if self._metrics_feed is not None and self._feedback.enabled: + snapshot = self._metrics_feed.get(handle.instance_id) + capacity = getattr(snapshot, "kv_capacity_tokens", -1) + if capacity > 0: + return int(capacity) + return handle.kv_capacity_tokens or self._kv_capacity_by_tp.get(handle.tp_degree, 0) + + def attach_metrics_feed(self, feed: Any) -> None: + """Attach an engine-layer metrics poller (``get(instance_id)``).""" + self._metrics_feed = feed + + def attach_shared_state(self, state: Any) -> None: + """Attach an engine-layer cross-rank load publisher/aggregator. + + The engine (which survives scheduler rebinding) owns the state; + attaching only wires this scheduler's accounting publishes and + aggregated ranking reads. Callers must attach with the scheduler's + local counters at zero (fresh scheduler) — the engine resets the + rank's published totals before attaching. + """ + self._shared_state = state + + def load_of(self, handle: InstanceHandle) -> float: + """Load of a handle under the configured metric. + + With a shared state attached the ranking value is the + cluster-wide aggregate (all live ranks' in-flight accounting, + ours included) instead of this rank's partial view; without fresh + aggregate data it degrades to the local estimate. When /metrics + feedback is enabled, an additive bias (EMA of observed - base) + then anchors whichever base is used to the real occupancy. + """ + capacity = self.kv_capacity_of(handle) + if self._capacity_seen.get(handle.instance_id, capacity) != capacity: + self._bias_ema.pop(handle.instance_id, None) + self._bias_seen_ts.pop(handle.instance_id, None) + self._capacity_seen[handle.instance_id] = capacity + if self._load_metric == "kv_tokens" and capacity <= 0 and any( + self.kv_capacity_of(h) > 0 for h in self._instances + ): + return float("inf") + base = handle.load( + self._load_metric, + kv_capacity_tokens=capacity, + ) + base = self._apply_shared_aggregate(handle, base) + return self._apply_metrics_bias(handle, base) + + def _apply_shared_aggregate( + self, handle: InstanceHandle, local: float + ) -> float: + if self._shared_state is None: + return local + try: + totals = self._shared_state.totals() + except Exception as exc: + logger.warning( + "[%s] shared aggregate read failed, using local view: %s", + self.name, exc, + ) + return local + entry = totals.get(handle.instance_id) if totals else None + if not entry: + return local + if self._load_metric == "requests": + return float(entry.get("requests", 0)) + tokens = float(entry.get("tokens", 0)) + if self._load_metric == "kv_tokens": + capacity = self.kv_capacity_of(handle) + return tokens / max(1, capacity) + return tokens + + def _observed_load(self, handle: InstanceHandle, snapshot: Any) -> float | None: + """Map a metrics snapshot onto the current load metric's unit.""" + if self._load_metric == "kv_tokens": + if self.kv_capacity_of(handle) <= 0: + return None + usage = float(getattr(snapshot, "gpu_cache_usage", -1.0)) + return usage if 0.0 <= usage <= 1.0 else None + if self._load_metric == "tokens": + tokens = float(getattr(snapshot, "kv_cache_tokens", -1.0)) + if tokens < 0: + capacity = getattr(snapshot, "kv_capacity_tokens", -1) + usage = getattr(snapshot, "gpu_cache_usage", -1) + if capacity > 0 and 0 <= usage <= 1: + tokens = capacity * usage + return tokens if tokens >= 0 else None + running = int(getattr(snapshot, "running", -1)) + waiting = int(getattr(snapshot, "waiting", -1)) + return float(running + max(0, waiting)) if running >= 0 else None + + def _apply_metrics_bias(self, handle: InstanceHandle, local: float) -> float: + if self._metrics_feed is None or not self._feedback.enabled: + return local + snapshot = self._metrics_feed.get(handle.instance_id) + if snapshot is None: + self._bias_ema.pop(handle.instance_id, None) + self._bias_seen_ts.pop(handle.instance_id, None) + return local + observed = self._observed_load(handle, snapshot) + if observed is None: + return local + key = handle.instance_id + updated_at = float(getattr(snapshot, "updated_at", 0.0)) + if updated_at > self._bias_seen_ts.get(key, 0.0): + self._bias_seen_ts[key] = updated_at + alpha = self._feedback.bias_alpha + previous = self._bias_ema.get(key, 0.0) + self._bias_ema[key] = alpha * (observed - local) + (1.0 - alpha) * previous + return max(0.0, local + self._bias_ema.get(key, 0.0)) + + def _admission_ok(self, handle: InstanceHandle) -> bool: + """Occupancy admission with hysteresis and preemption TTL. + + Blocking requires a FRESH snapshot; a stale or missing feed keeps + the last decision sticky so a transient scrape failure cannot flap + admission state. + """ + if self._metrics_feed is None or not self._feedback.enabled: + return True + key = handle.instance_id + blocked = self._admission_blocked.get(key, False) + snapshot = self._metrics_feed.get(key) + if snapshot is None: + self._admission_blocked.pop(key, None) + return True + now = time.time() + last_preemption_at = float(getattr(snapshot, "last_preemption_at", 0.0)) + if ( + last_preemption_at > 0.0 + and now - last_preemption_at < self._feedback.preemption_penalty_ttl_s + ): + blocked = True + else: + occupancy = float(getattr(snapshot, "gpu_cache_usage", -1.0)) + if 0.0 <= occupancy <= 1.0: + if not blocked and occupancy > self._feedback.admission_enter: + blocked = True + elif blocked and occupancy < self._feedback.admission_exit: + blocked = False + self._admission_blocked[key] = blocked + return not blocked + + def _apply_admission(self, handles: list[InstanceHandle]) -> list[InstanceHandle]: + if self._metrics_feed is None or not self._feedback.enabled: + return handles + admitted = [h for h in handles if self._admission_ok(h)] + if admitted: + return admitted + # A preferred TP subset being blocked is not a global overload. + comparable = self._instances + if self._load_metric == "kv_tokens" and any(self.kv_capacity_of(h) > 0 for h in comparable): + comparable = [h for h in comparable if self.kv_capacity_of(h) > 0] + if any(h.is_ready and self._admission_ok(h) for h in comparable): + return [] + # All instances overloaded (or degraded): keep every candidate so + # routing still works instead of hard-failing; the least-loaded + # pick then applies graceful overload shedding. + return handles + + def _selectable( + self, + handles: list[InstanceHandle], + metric: str | None = None, + ) -> list[InstanceHandle]: + """Filter handles comparable under the metric. + + Under ``kv_tokens`` every compared value must be a ratio; an + instance whose TP bucket has no configured capacity would + otherwise contribute a raw token count into a ratio-sorted + candidate set, starving the big buckets. If NO instance has a + capacity we fall back to comparing raw token counts for all + (consistent units), matching the legacy degradation. + """ + if (metric or self._load_metric) != "kv_tokens": + return self._apply_admission(list(handles)) + capacitated = [h for h in handles if self.kv_capacity_of(h) > 0] + if any(self.kv_capacity_of(h) > 0 for h in self._instances): + if len(capacitated) < len(handles): + missing = [ + h.instance_id + for h in handles + if self.kv_capacity_of(h) <= 0 + ] + logger.warning( + "[%s] kv_tokens: excluding instances without KV capacity " + "from selection: %s", + self.name, + ", ".join(missing), + ) + return self._apply_admission(capacitated) + return self._apply_admission(list(handles)) + + def _expected_output_tokens( + self, category: str = "any", max_new_tokens: int = 0 + ) -> int: + """Expected generation length for a category (EMA, floor at 1). + + A small request budget caps a large historical mean to avoid + overestimating output. This does not fix underestimation of tails. + """ + expected = max(1, int(self._output_ema[category])) + if max_new_tokens > 0: + expected = min(expected, max(1, int(max_new_tokens))) + return expected + + def _update_output_ema(self, category: str, output_tokens: int) -> None: + """Fold an observed completion length into the category EMA.""" + if output_tokens <= 0: + return + prev = self._output_ema[category or "any"] + self._output_ema[category or "any"] = ( + OUTPUT_EMA_ALPHA * float(output_tokens) + + (1.0 - OUTPUT_EMA_ALPHA) * prev + ) + + def _record_route( + self, + handle: InstanceHandle, + input_tokens: int, + category: str = "any", + prompt_id: str = "", + max_new_tokens: int = 0, + ) -> int: + """Account a routed request (callers must hold ``self._lock``).""" + expected = self._expected_output_tokens(category, max_new_tokens) + handle.inc_active( + prompt_tokens=max(0, int(input_tokens)), + expected_output_tokens=expected, + ) + if prompt_id and self._prefix_affinity_enabled: + self._affinity_record(prompt_id, handle) + estimate = max(0, int(input_tokens)) + expected + if self._shared_state is not None: + try: + self._shared_state.add( + handle.instance_id, delta_requests=1, delta_tokens=estimate + ) + except Exception as exc: + logger.warning( + "[%s] shared-state publish failed on route: %s", + self.name, exc, + ) + return estimate + + def _complete_route( + self, + instance_index: int, + output_tokens: int = 0, + category: str = "any", + reserved_tokens: int | None = None, + ) -> None: + """Retire a completed request and update the output EMA. + + Callers must hold ``self._lock``. + """ + handle = self.get_instance_handle(instance_index) + retired = handle.dec_active(reserved_tokens) if handle is not None else 0 + if self._shared_state is not None and handle is not None: + try: + self._shared_state.add( + handle.instance_id, + delta_requests=-1, + delta_tokens=-retired, + ) + except Exception as exc: + logger.warning( + "[%s] shared-state publish failed on completion: %s", + self.name, exc, + ) + self._update_output_ema(category, output_tokens) # ---------------------------------------------------------------- @@ -153,6 +627,7 @@ def schedule( prompt_id: str = "", n_samples: int = 1, epoch: int = -1, + max_new_tokens: int = 0, ) -> SchedulingResult: """Schedule.""" ... @@ -184,6 +659,125 @@ def on_epoch_end(self, epoch: int): # ---------------------------------------------------------------- + def _min_load_rotating(self, candidates: list[InstanceHandle]) -> InstanceHandle: + """Least-load selection with rotation among equal-load endpoints. + + A plain ``min`` always picks the first registered endpoint when + requests complete between scheduling calls, starving later + endpoints under low-concurrency multi-turn workloads. Rotation + position is stored after the selected handle so endpoints are not + permanently skipped when requests enter/leave the candidate set. + """ + scores = {id(h): self.load_of(h) for h in candidates} + min_load = min(scores.values()) + least_loaded = {key for key, score in scores.items() if score == min_load} + start = self._tie_rr_counter % len(self._instances) + for offset in range(len(self._instances)): + position = (start + offset) % len(self._instances) + handle = self._instances[position] + if id(handle) in least_loaded: + self._tie_rr_counter = (position + 1) % len(self._instances) + return handle + # Candidates not present in _instances (stale handles): fall back + # to a deterministic first-min pick. + return min(candidates, key=lambda h: self.load_of(h)) + + # ---------------------------------------------------------------- + + # ---------------------------------------------------------------- + + def _affinity_pick( + self, + prompt_id: str, + candidates: list[InstanceHandle], + ) -> InstanceHandle | None: + """Sticky-instance pick for ``prompt_id`` among valid candidates. + + Callers pass the post-filter candidate set (readiness, queue, + capacity normalization, admission): affinity never overrides + those safety filters, it only reorders preference among the + survivors. Returns ``None`` when affinity is off, the prompt is + unknown, or its sticky instance did not survive filtering. + Callers must hold ``self._lock``. + """ + if not self._prefix_affinity_enabled or not prompt_id: + return None + instance_id = self._affinity.get(prompt_id) + if instance_id is None: + return None + self._affinity.move_to_end(prompt_id) + for handle in candidates: + if handle.instance_id == instance_id: + # Global-overload fallback may have reintroduced blocked or + # full endpoints. Affinity must not override that condition. + if not self._admission_ok(handle): + return None + limit = getattr(self, "_max_queue_length", 0) + if limit > 0 and handle.active_requests >= limit: + return None + if self._load_metric == "kv_tokens" and self.load_of(handle) >= 1: + return None + return handle + return None + + def _affinity_record(self, prompt_id: str, handle: InstanceHandle) -> None: + """Bind prompt_id to the instance that just served it (LRU). + + Callers must hold ``self._lock``. + """ + if not self._prefix_affinity_enabled or not prompt_id: + return + self._affinity[prompt_id] = handle.instance_id + self._affinity.move_to_end(prompt_id) + while len(self._affinity) > self._prefix_affinity_max_entries: + self._affinity.popitem(last=False) + + def export_learned_state(self) -> dict: + """Scheduler state that should survive a rebind/reconfigure. + + Weight-sync rebinding rebuilds schedulers every training step + (sync_interval defaults to 1); without carrying this state the + output-length EMA would reset to its prior every step and never + learn. Subclasses with more learned state (e.g. LA-MLFQ history) + should extend this dict. + """ + return { + "output_ema": { + cat: float(v) for cat, v in self._output_ema.items() + }, + "prefix_affinity": dict(self._affinity), + } + + def import_learned_state(self, state: dict) -> None: + """Restore state exported by ``export_learned_state``. + + Only categories with valid values are applied; unknown/shorter + states are tolerated so a scheduler-type change at rebind time + degrades gracefully instead of crashing. + """ + if not isinstance(state, dict): + return + ema = state.get("output_ema") + if isinstance(ema, dict): + for cat, v in ema.items(): + try: + v = float(v) + except (TypeError, ValueError): + continue + if v > 0: + self._output_ema[str(cat)] = v + affinity = state.get("prefix_affinity") + if ( + self._prefix_affinity_enabled + and isinstance(affinity, dict) + and affinity + ): + for pid, iid in affinity.items(): + if isinstance(pid, str) and isinstance(iid, str): + self._affinity[pid] = iid + while len(self._affinity) > self._prefix_affinity_max_entries: + self._affinity.popitem(last=False) + def get_stats(self) -> SchedulerStats: return self._stats @@ -214,7 +808,8 @@ def print_stats(self): for h in self._instances: print( f" {h.instance_id}: TP={h.tp_degree}, " - f"active={h.active_requests}, ready={h.is_ready}" + f"active={h.active_requests}, active_tokens={h.active_tokens}, " + f"ready={h.is_ready}" ) print("=" * 60) diff --git a/infra/scheduling/cmlfq_scheduler.py b/infra/scheduling/cmlfq_scheduler.py index 2b582cf..a70c987 100644 --- a/infra/scheduling/cmlfq_scheduler.py +++ b/infra/scheduling/cmlfq_scheduler.py @@ -176,6 +176,7 @@ def schedule( prompt_id: str = "", n_samples: int = 1, epoch: int = -1, + max_new_tokens: int = 0, ) -> SchedulingResult: """Schedule.""" with self._lock: diff --git a/infra/scheduling/la_mlfq.py b/infra/scheduling/la_mlfq.py index a1e8ef1..d00ee0c 100644 --- a/infra/scheduling/la_mlfq.py +++ b/infra/scheduling/la_mlfq.py @@ -12,9 +12,11 @@ from typing import Any, Callable, Optional from RL_Framework.infra.scheduling.base import ( + DEFAULT_LOAD_METRIC, BaseScheduler, InstanceHandle, LoadBalanceStrategy, + MetricsFeedbackConfig, RoutingRule, SchedulerStats, SchedulingResult, @@ -56,6 +58,7 @@ class WaitingRequest: sample_index: int created_at: float = field(default_factory=time.time) future: Optional[asyncio.Future] = None + max_new_tokens: int = 0 @dataclass @@ -428,8 +431,20 @@ def __init__( load_balance_strategy: str = "least_connections", max_queue_length: int = 100, enable_fallback: bool = True, + load_metric: str = DEFAULT_LOAD_METRIC, + kv_capacity_tokens_by_tp: dict[int, int] | None = None, + feedback: MetricsFeedbackConfig | None = None, + prefix_affinity: bool = False, + prefix_affinity_max_entries: int = 8192, ): - super().__init__(name="LA-MLFQ") + super().__init__( + name="LA-MLFQ", + load_metric=load_metric, + kv_capacity_tokens_by_tp=kv_capacity_tokens_by_tp, + feedback=feedback, + prefix_affinity=prefix_affinity, + prefix_affinity_max_entries=prefix_affinity_max_entries, + ) self._buckets = buckets or dict(self.DEFAULT_BUCKETS) @@ -511,6 +526,7 @@ def schedule( prompt_id: str = "", n_samples: int = 1, epoch: int = -1, + max_new_tokens: int = 0, ) -> SchedulingResult: """Schedule.""" if epoch >= 0: @@ -521,7 +537,7 @@ def schedule( if not prompt_id: - return self._default_length_route(input_tokens, prompt_id) + return self._default_length_route(input_tokens, prompt_id, max_new_tokens) historical_bucket = self.history_table.lookup( @@ -534,7 +550,7 @@ def schedule( ) result = self._route_to_bucket( historical_bucket, input_tokens, prompt_id, - reason="history_hit", + reason="history_hit", max_new_tokens=max_new_tokens, ) if result.instance_index >= 0: return result @@ -547,7 +563,7 @@ def schedule( result = self._route_to_bucket( "short", input_tokens, prompt_id, - reason="scout", + reason="scout", max_new_tokens=max_new_tokens, ) if result.instance_index >= 0: self.scout_manager.register_scout( @@ -569,7 +585,7 @@ def schedule( f"[LA-MLFQ] Short Bucket has no available instance, " f"prompt={prompt_id} using default routing" ) - return self._default_length_route(input_tokens, prompt_id) + return self._default_length_route(input_tokens, prompt_id, max_new_tokens) else: @@ -580,7 +596,7 @@ def schedule( if target_bucket: return self._route_to_bucket( target_bucket, input_tokens, prompt_id, - reason="scout_follow", + reason="scout_follow", max_new_tokens=max_new_tokens, ) @@ -597,7 +613,7 @@ def schedule( ) - return self._default_length_route(input_tokens, prompt_id) + return self._default_length_route(input_tokens, prompt_id, max_new_tokens) def _route_to_bucket( self, @@ -605,24 +621,62 @@ def _route_to_bucket( input_tokens: int, prompt_id: str = "", reason: str = "", + max_new_tokens: int = 0, ) -> SchedulingResult: """Route to bucket.""" rule = self._bucket_rules.get(bucket) if rule is None: - return self._default_length_route(input_tokens, prompt_id) + return self._default_length_route(input_tokens, prompt_id, max_new_tokens) with self._lock: self._stats.category_counts[bucket] += 1 + # Instance-level prefix affinity outranks bucket routing: a + # cache hit on the sticky instance beats the bucket's TP + # preference. Safety filters are applied to the candidate set + # before the sticky lookup. + affinity_candidates = self._selectable( + [ + h for h in self._instances + if self._prefix_affinity_enabled and h.is_ready + and h.tp_degree in rule.preferred_tp_degrees and ( + self._max_queue_length <= 0 + or h.active_requests < self._max_queue_length + ) + ] + ) + selected = self._affinity_pick(prompt_id, affinity_candidates) + if selected is not None: + self._stats.preferred_routes += 1 + self._stats.category_tp_counts[bucket][selected.tp_degree] += 1 + reserved = self._record_route( + selected, input_tokens, category=bucket, prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) + return SchedulingResult( + instance_index=selected.index, + reserved_tokens=reserved, + tp_degree=selected.tp_degree, + category=bucket, + is_fallback=False, + reason="prefix_affinity", + prompt_id=prompt_id, + ) + + selected = self._try_select(rule.preferred_tp_degrees) if selected is not None: self._stats.preferred_routes += 1 self._stats.category_tp_counts[bucket][selected.tp_degree] += 1 - selected.inc_active() + reserved = self._record_route( + selected, input_tokens, category=bucket, prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) return SchedulingResult( instance_index=selected.index, + reserved_tokens=reserved, tp_degree=selected.tp_degree, category=bucket, is_fallback=False, @@ -636,9 +690,13 @@ def _route_to_bucket( if selected is not None: self._stats.fallback_routes += 1 self._stats.category_tp_counts[bucket][selected.tp_degree] += 1 - selected.inc_active() + reserved = self._record_route( + selected, input_tokens, category=bucket, prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) return SchedulingResult( instance_index=selected.index, + reserved_tokens=reserved, tp_degree=selected.tp_degree, category=bucket, is_fallback=True, @@ -647,14 +705,18 @@ def _route_to_bucket( ) - all_ready = [h for h in self._instances if h.is_ready] + all_ready = self._selectable([h for h in self._instances if h.is_ready]) if all_ready: - selected = min(all_ready, key=lambda h: h.active_requests) + selected = self._min_load_rotating(all_ready) self._stats.fallback_routes += 1 self._stats.category_tp_counts[bucket][selected.tp_degree] += 1 - selected.inc_active() + reserved = self._record_route( + selected, input_tokens, category=bucket, prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) return SchedulingResult( instance_index=selected.index, + reserved_tokens=reserved, tp_degree=selected.tp_degree, category=bucket, is_fallback=True, @@ -675,11 +737,14 @@ def _route_to_bucket( ) def _default_length_route( - self, input_tokens: int, prompt_id: str = "" + self, input_tokens: int, prompt_id: str = "", max_new_tokens: int = 0 ) -> SchedulingResult: """Default length route.""" bucket = self._categorize_to_bucket(input_tokens) - return self._route_to_bucket(bucket, input_tokens, prompt_id, reason="default") + return self._route_to_bucket( + bucket, input_tokens, prompt_id, + reason="default", max_new_tokens=max_new_tokens, + ) def _try_select(self, tp_preferences: list[int]) -> Optional[InstanceHandle]: """Try select.""" @@ -690,6 +755,7 @@ def _try_select(self, tp_preferences: list[int]) -> Optional[InstanceHandle]: ] if not candidates: continue + candidates = self._selectable(candidates) if self._max_queue_length > 0: candidates = [ h for h in candidates @@ -702,7 +768,7 @@ def _try_select(self, tp_preferences: list[int]) -> Optional[InstanceHandle]: self._rr_index[tp_degree] += 1 return candidates[idx] else: - return min(candidates, key=lambda h: h.active_requests) + return self._min_load_rotating(candidates) return None # ---------------------------------------------------------------- @@ -715,13 +781,17 @@ def on_request_done( prompt_id: str = "", final_bucket: str = "", output_tokens: int = 0, + reserved_tokens: int | None = None, ): """On request done.""" - handle = self.get_instance_handle(instance_index) - if handle: - with self._lock: - handle.dec_active() + with self._lock: + self._complete_route( + instance_index, + output_tokens=output_tokens, + category=final_bucket or "any", + reserved_tokens=reserved_tokens, + ) if prompt_id and final_bucket: @@ -779,9 +849,12 @@ def _process_released_waiting( ): """Process released waiting.""" for wr in waiting_requests: + if wr.future is not None and wr.future.done(): + continue result = self._route_to_bucket( target_bucket, wr.input_tokens, wr.prompt_id, reason="scout_released", + max_new_tokens=wr.max_new_tokens, ) if wr.future is not None and not wr.future.done(): @@ -808,7 +881,9 @@ def on_epoch_end(self, epoch: int): waiting = self.scout_manager.force_release(pid) for wr in waiting: - result = self._default_length_route(wr.input_tokens, wr.prompt_id) + if wr.future is not None and wr.future.done(): + continue + result = self._default_length_route(wr.input_tokens, wr.prompt_id, wr.max_new_tokens) if wr.future is not None and not wr.future.done(): wr.future.set_result(result) self.scout_manager.reset() @@ -836,6 +911,20 @@ def print_stats(self): def reset_stats(self): super().reset_stats() + def export_learned_state(self) -> dict: + state = super().export_learned_state() + # Carry the prompt->bucket history table by reference swap; the + # table is internally locked, and the old scheduler is discarded + # right after the rebind anyway. + state["la_mlfq_history_table"] = self.history_table + return state + + def import_learned_state(self, state: dict) -> None: + super().import_learned_state(state) + table = (state or {}).get("la_mlfq_history_table") + if table is not None and hasattr(table, "lookup") and hasattr(table, "update"): + self.history_table = table + # ---------------------------------------------------------------- # ---------------------------------------------------------------- @@ -864,4 +953,11 @@ def from_config(cls, hetero_config: Any) -> "LAMLFQScheduler": load_balance_strategy=sched.load_balance_strategy, max_queue_length=sched.max_queue_length, enable_fallback=sched.enable_fallback, + load_metric=getattr(sched, "load_metric", DEFAULT_LOAD_METRIC), + kv_capacity_tokens_by_tp=getattr(sched, "kv_capacity_tokens_by_tp", None), + feedback=MetricsFeedbackConfig.from_scheduling(sched), + prefix_affinity=bool(getattr(sched, "prefix_affinity", False)), + prefix_affinity_max_entries=int( + getattr(sched, "prefix_affinity_max_entries", 8192) + ), ) diff --git a/infra/scheduling/length_aware.py b/infra/scheduling/length_aware.py index b0ee680..c45fbec 100644 --- a/infra/scheduling/length_aware.py +++ b/infra/scheduling/length_aware.py @@ -7,9 +7,11 @@ from typing import Any from RL_Framework.infra.scheduling.base import ( + DEFAULT_LOAD_METRIC, BaseScheduler, InstanceHandle, LoadBalanceStrategy, + MetricsFeedbackConfig, RoutingRule, SchedulerStats, SchedulingResult, @@ -35,8 +37,20 @@ def __init__( load_balance_strategy: str = "least_connections", max_queue_length: int = 100, enable_fallback: bool = True, + load_metric: str = DEFAULT_LOAD_METRIC, + kv_capacity_tokens_by_tp: dict[int, int] | None = None, + feedback: MetricsFeedbackConfig | None = None, + prefix_affinity: bool = False, + prefix_affinity_max_entries: int = 8192, ): - super().__init__(name="LengthAware") + super().__init__( + name="LengthAware", + load_metric=load_metric, + kv_capacity_tokens_by_tp=kv_capacity_tokens_by_tp, + feedback=feedback, + prefix_affinity=prefix_affinity, + prefix_affinity_max_entries=prefix_affinity_max_entries, + ) self._rules: dict[str, RoutingRule] = dict(self.DEFAULT_ROUTING_RULES) @@ -102,6 +116,7 @@ def schedule( prompt_id: str = "", n_samples: int = 1, epoch: int = -1, + max_new_tokens: int = 0, ) -> SchedulingResult: """Schedule.""" category = self.categorize(input_tokens) @@ -112,13 +127,50 @@ def schedule( self._stats.category_counts[category] += 1 + # Prefix affinity outranks the length-bucket preference: a + # cache hit on the sticky instance beats a nominally better + # TP fit on a cold one. Filters (readiness/queue/capacity/ + # admission) are applied to the candidate set first. + affinity_candidates = self._selectable( + [ + h for h in self._instances + if self._prefix_affinity_enabled and h.is_ready + and h.tp_degree in rule.preferred_tp_degrees and ( + self._max_queue_length <= 0 + or h.active_requests < self._max_queue_length + ) + ] + ) + selected = self._affinity_pick(prompt_id, affinity_candidates) + if selected is not None: + self._stats.preferred_routes += 1 + self._stats.category_tp_counts[category][selected.tp_degree] += 1 + reserved = self._record_route( + selected, input_tokens, category=category, prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) + return SchedulingResult( + instance_index=selected.index, + reserved_tokens=reserved, + tp_degree=selected.tp_degree, + category=category, + is_fallback=False, + reason="prefix_affinity", + prompt_id=prompt_id, + ) + + selected = self._try_select(rule.preferred_tp_degrees) if selected is not None: self._stats.preferred_routes += 1 self._stats.category_tp_counts[category][selected.tp_degree] += 1 - selected.inc_active() + reserved = self._record_route( + selected, input_tokens, category=category, prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) return SchedulingResult( instance_index=selected.index, + reserved_tokens=reserved, tp_degree=selected.tp_degree, category=category, is_fallback=False, @@ -131,9 +183,13 @@ def schedule( if selected is not None: self._stats.fallback_routes += 1 self._stats.category_tp_counts[category][selected.tp_degree] += 1 - selected.inc_active() + reserved = self._record_route( + selected, input_tokens, category=category, prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) return SchedulingResult( instance_index=selected.index, + reserved_tokens=reserved, tp_degree=selected.tp_degree, category=category, is_fallback=True, @@ -141,14 +197,18 @@ def schedule( ) - all_ready = [h for h in self._instances if h.is_ready] + all_ready = self._selectable([h for h in self._instances if h.is_ready]) if all_ready: - selected = min(all_ready, key=lambda h: h.active_requests) + selected = self._min_load_rotating(all_ready) self._stats.fallback_routes += 1 self._stats.category_tp_counts[category][selected.tp_degree] += 1 - selected.inc_active() + reserved = self._record_route( + selected, input_tokens, category=category, prompt_id=prompt_id, + max_new_tokens=max_new_tokens, + ) return SchedulingResult( instance_index=selected.index, + reserved_tokens=reserved, tp_degree=selected.tp_degree, category=category, is_fallback=True, @@ -173,12 +233,16 @@ def on_request_done( prompt_id: str = "", final_bucket: str = "", output_tokens: int = 0, + reserved_tokens: int | None = None, ): """On request done.""" - handle = self.get_instance_handle(instance_index) - if handle: - with self._lock: - handle.dec_active() + with self._lock: + self._complete_route( + instance_index, + output_tokens=output_tokens, + category=final_bucket or "any", + reserved_tokens=reserved_tokens, + ) def _try_select(self, tp_preferences: list[int]) -> InstanceHandle | None: """Try select.""" @@ -190,6 +254,7 @@ def _try_select(self, tp_preferences: list[int]) -> InstanceHandle | None: if not candidates: continue + candidates = self._selectable(candidates) if self._max_queue_length > 0: candidates = [ @@ -205,8 +270,8 @@ def _try_select(self, tp_preferences: list[int]) -> InstanceHandle | None: self._rr_index[tp_degree] += 1 return candidates[idx] else: - # least_connections - return min(candidates, key=lambda h: h.active_requests) + # least_connections with rotation among equal-load endpoints + return self._min_load_rotating(candidates) return None @@ -247,6 +312,13 @@ def from_config(cls, hetero_config: Any) -> "LengthAwareScheduler": load_balance_strategy=sched.load_balance_strategy, max_queue_length=sched.max_queue_length, enable_fallback=sched.enable_fallback, + load_metric=getattr(sched, "load_metric", DEFAULT_LOAD_METRIC), + kv_capacity_tokens_by_tp=getattr(sched, "kv_capacity_tokens_by_tp", None), + feedback=MetricsFeedbackConfig.from_scheduling(sched), + prefix_affinity=bool(getattr(sched, "prefix_affinity", False)), + prefix_affinity_max_entries=int( + getattr(sched, "prefix_affinity_max_entries", 8192) + ), ) diff --git a/infra/scheduling/load_balance.py b/infra/scheduling/load_balance.py index 6d9fb9d..6098ef6 100644 --- a/infra/scheduling/load_balance.py +++ b/infra/scheduling/load_balance.py @@ -7,9 +7,11 @@ from typing import Any, Optional from RL_Framework.infra.scheduling.base import ( + DEFAULT_LOAD_METRIC, BaseScheduler, InstanceHandle, LoadBalanceStrategy, + MetricsFeedbackConfig, SchedulingResult, ) @@ -24,8 +26,20 @@ def __init__( load_balance_strategy: str = "least_connections", max_queue_length: int = 100, weights: dict[int, float] | None = None, + load_metric: str = DEFAULT_LOAD_METRIC, + kv_capacity_tokens_by_tp: dict[int, int] | None = None, + feedback: MetricsFeedbackConfig | None = None, + prefix_affinity: bool = False, + prefix_affinity_max_entries: int = 8192, ): - super().__init__(name="LoadBalance") + super().__init__( + name="LoadBalance", + load_metric=load_metric, + kv_capacity_tokens_by_tp=kv_capacity_tokens_by_tp, + feedback=feedback, + prefix_affinity=prefix_affinity, + prefix_affinity_max_entries=prefix_affinity_max_entries, + ) if load_balance_strategy == "round_robin": self._strategy = LoadBalanceStrategy.ROUND_ROBIN @@ -53,15 +67,17 @@ def schedule( prompt_id: str = "", n_samples: int = 1, epoch: int = -1, + max_new_tokens: int = 0, ) -> SchedulingResult: """Schedule.""" with self._lock: self._stats.total_requests += 1 + ready = self._selectable([h for h in self._instances if h.is_ready]) candidates = [ - h for h in self._instances - if h.is_ready and ( + h for h in ready + if ( self._max_queue_length <= 0 or h.active_requests < self._max_queue_length ) @@ -69,7 +85,7 @@ def schedule( if not candidates: - candidates = [h for h in self._instances if h.is_ready] + candidates = ready if not candidates: self._stats.failed_routes += 1 @@ -83,17 +99,26 @@ def schedule( ) - selected = self._select(candidates) + selected = self._affinity_pick(prompt_id, candidates) + if selected is not None: + self._stats.category_counts["any"] += 1 + self._stats.category_tp_counts["any"][selected.tp_degree] += 1 + reason = "prefix_affinity" + else: + selected = self._select(candidates) + self._stats.category_counts["any"] += 1 + self._stats.category_tp_counts["any"][selected.tp_degree] += 1 + reason = "" self._stats.preferred_routes += 1 - self._stats.category_counts["any"] += 1 - self._stats.category_tp_counts["any"][selected.tp_degree] += 1 - selected.inc_active() + reserved = self._record_route(selected, input_tokens, category="any", prompt_id=prompt_id, max_new_tokens=max_new_tokens) return SchedulingResult( instance_index=selected.index, + reserved_tokens=reserved, tp_degree=selected.tp_degree, category="any", is_fallback=False, + reason=reason, prompt_id=prompt_id, ) @@ -108,33 +133,11 @@ def _select(self, candidates: list[InstanceHandle]) -> InstanceHandle: def weighted_load(h: InstanceHandle) -> float: w = self._weights.get(h.tp_degree, 1.0) - return h.active_requests / max(w, 0.01) + return self.load_of(h) / max(w, 0.01) return min(candidates, key=weighted_load) else: - min_active = min(h.active_requests for h in candidates) - least_loaded = [ - h for h in candidates if h.active_requests == min_active - ] - # A plain ``min`` always picks the first registered endpoint when - # requests complete between scheduling calls. That starves later - # endpoints (and can leave an entire rollout node idle) for - # low-concurrency, multi-turn workloads such as R2E-Gym. Rotate - # among equal-load endpoints while preserving least-connections - # as the primary selection criterion. - least_loaded_ids = {id(h) for h in least_loaded} - start = self._rr_counter % len(self._instances) - for offset in range(len(self._instances)): - position = (start + offset) % len(self._instances) - handle = self._instances[position] - if id(handle) in least_loaded_ids: - # Store the position after the selected handle. Unlike - # taking counter modulo the changing tie-set size, this - # cannot permanently skip endpoints when active requests - # enter and leave the candidate set. - self._rr_counter = (position + 1) % len(self._instances) - return handle - raise RuntimeError("least-connections tie set is inconsistent") + return self._min_load_rotating(candidates) def on_request_done( self, @@ -142,12 +145,16 @@ def on_request_done( prompt_id: str = "", final_bucket: str = "", output_tokens: int = 0, + reserved_tokens: int | None = None, ): """On request done.""" - handle = self.get_instance_handle(instance_index) - if handle: - with self._lock: - handle.dec_active() + with self._lock: + self._complete_route( + instance_index, + output_tokens=output_tokens, + category=final_bucket or "any", + reserved_tokens=reserved_tokens, + ) # ---------------------------------------------------------------- @@ -160,4 +167,11 @@ def from_config(cls, hetero_config: Any) -> "LoadBalanceScheduler": return cls( load_balance_strategy=sched.load_balance_strategy, max_queue_length=sched.max_queue_length, + load_metric=getattr(sched, "load_metric", DEFAULT_LOAD_METRIC), + kv_capacity_tokens_by_tp=getattr(sched, "kv_capacity_tokens_by_tp", None), + feedback=MetricsFeedbackConfig.from_scheduling(sched), + prefix_affinity=bool(getattr(sched, "prefix_affinity", False)), + prefix_affinity_max_entries=int( + getattr(sched, "prefix_affinity_max_entries", 8192) + ), ) diff --git a/infra/scheduling/metrics_feed.py b/infra/scheduling/metrics_feed.py new file mode 100644 index 0000000..72df2a7 --- /dev/null +++ b/infra/scheduling/metrics_feed.py @@ -0,0 +1,321 @@ +"""Background /metrics feed from vLLM instances for closed-loop scheduling. + +The poller lives at the engine layer (it must survive scheduler rebinding +during weight sync) and stores the latest Prometheus snapshot per instance. +Schedulers read snapshots through the ``get(instance_id)`` interface and use +them for additive bias correction and admission control; any feed outage +degrades transparently back to the local open-loop estimate. +""" + +from __future__ import annotations + +import logging +import math +import re +import random +import threading +import time +import urllib.request +from dataclasses import dataclass + +logger = logging.getLogger(__name__) + + +@dataclass +class InstanceMetrics: + """Latest observed metrics for one vLLM instance.""" + + instance_id: str + updated_at: float + running: int = -1 + waiting: int = -1 + gpu_cache_usage: float = -1.0 # 0..1 occupancy of the GPU KV cache + kv_cache_tokens: float = -1.0 + preemptions_total: float = -1.0 # cumulative counter, never resets + last_preemption_at: float = 0.0 + consecutive_failures: int = 0 + kv_capacity_tokens: int = -1 # vLLM-profiled capacity (config_info) + + +_METRIC_ALIASES: dict[str, tuple[str, ...]] = { + "running": ( + "vllm:num_requests_running", + "vllm_num_requests_running", + ), + "waiting": ( + "vllm:num_requests_waiting", + "vllm_num_requests_waiting", + ), + "gpu_cache_usage": ( + "vllm:gpu_cache_usage_perc", + "vllm_gpu_cache_usage_perc", + "vllm:kv_cache_usage_perc", + "vllm_kv_cache_usage_perc", + ), + "kv_cache_tokens": ( + "vllm:kv_cache_tokens", + "vllm_kv_cache_tokens", + "vllm:gpu_cache_usage_tokens", + "vllm_gpu_cache_usage_tokens", + ), + "preemptions_total": ( + "vllm:num_preemptions_total", + "vllm_num_preemptions_total", + "vllm:preemption_total", + "vllm_preemption_total", + "vllm:preemptions_total", + "vllm_preemptions_total", + ), +} + +# cache_config_info carries vLLM's own profiled KV capacity (ground truth); +# the label value lives in the exposition line itself, so it is parsed +# separately from the numeric gauges. +_KV_SIZE_LABEL_KEYS = ( + 'kv_cache_size_tokens="', + 'kv_cache_size_tokens=\'', +) + + +def parse_prometheus_metrics(text: str) -> dict[str, float]: + """Lenient Prometheus text exposition parsing (name -> last value). + + Metric names differ across vLLM versions and may carry ``{label}`` + suffixes; both are tolerated. Lines we cannot understand are skipped + instead of raising so a partial scrape still yields the core gauges. + """ + values: dict[str, float] = {} + for line in text.splitlines(): + line = line.strip() + if not line or line.startswith("#"): + continue + match = re.match(r'^([\w:]+)(?:\{.*\})?\s+([^\s]+)', line) + if match is None: + continue + name, raw_value = match.groups() + try: + value = float(raw_value) + if not math.isfinite(value): + continue + if name in values: + if "cache_usage_perc" in name: + value = max(values[name], value) + else: + value += values[name] + values[name] = value + except ValueError: + continue + return values + + +def extract_kv_capacity_tokens(text: str) -> int: + """Pull vLLM's own profiled KV capacity out of cache_config_info. + + The gauge line embeds ``kv_cache_size_tokens=""`` as a label, which + the numeric parser cannot see. -1 when the exposition does not carry + it (older vLLM versions). + """ + for line in text.splitlines(): + if line.lstrip().startswith("#") or "cache_config_info" not in line: + continue + for key in _KV_SIZE_LABEL_KEYS: + start = line.find(key) + if start < 0: + continue + start += len(key) + end = line.find(key[-1], start) + if end <= start: + continue + try: + value = int(float(line[start:end])) + except (ValueError, OverflowError): + continue + if value > 0: + return value + return -1 + + +def extract_core_metrics(values: dict[str, float]) -> dict[str, float]: + """Map raw exposition values onto the core scheduling metrics.""" + normalized = {key.replace(":", "_"): value for key, value in values.items()} + extracted: dict[str, float] = {} + for field, aliases in _METRIC_ALIASES.items(): + for alias in aliases: + if alias in values: + extracted[field] = values[alias] + break + alias = alias.replace(":", "_") + if alias in normalized: + extracted[field] = normalized[alias] + break + return extracted + + +class VLLMMetricsPoller: + """Polls ``/metrics`` of every registered instance on a staggered clock. + + Each endpoint gets a random initial phase so multiple training ranks / + instances do not scrape in lockstep (which would synchronize admission + flapping). Snapshots older than ``ttl_s`` are reported as stale. + """ + + def __init__( + self, + interval_s: float = 3.0, + timeout_s: float = 1.0, + ttl_s: float = 10.0, + ): + self.interval_s = max(0.2, float(interval_s)) + self.timeout_s = max(0.1, float(timeout_s)) + self.ttl_s = max(self.interval_s * 2.0, float(ttl_s)) + self._endpoints: dict[str, str] = {} + self._next_due: dict[str, float] = {} + self._snapshots: dict[str, InstanceMetrics] = {} + self._lock = threading.Lock() + self._stop_event = threading.Event() + self._thread: threading.Thread | None = None + self._failures: dict[str, int] = {} + + # ---------------------------------------------------------------- + + # ---------------------------------------------------------------- + + def set_endpoints(self, urls: dict[str, str]) -> None: + """Replace the endpoint table; fresh snapshots keep their TTL.""" + now = time.time() + with self._lock: + for instance_id, old_url in self._endpoints.items(): + if urls.get(instance_id) != old_url: + self._snapshots.pop(instance_id, None) + self._next_due.pop(instance_id, None) + self._endpoints = dict(urls) + for instance_id in urls: + if instance_id not in self._next_due: + self._next_due[instance_id] = now + random.uniform( + 0.0, self.interval_s + ) + + def start(self) -> None: + if self._thread is not None and self._thread.is_alive(): + return + self._stop_event.clear() + self._thread = threading.Thread( + target=self._loop, + name="vllm-metrics-poller", + daemon=True, + ) + self._thread.start() + logger.info( + "Started vLLM metrics poller: interval=%.1fs timeout=%.1fs ttl=%.1fs endpoints=%d", + self.interval_s, + self.timeout_s, + self.ttl_s, + len(self._endpoints), + ) + + def stop(self) -> None: + self._stop_event.set() + thread = self._thread + if thread is not None and thread.is_alive(): + thread.join(timeout=max(2.0, self.timeout_s * 2.0)) + self._thread = None + + def get(self, instance_id: str, ttl_s: float | None = None) -> InstanceMetrics | None: + """Latest snapshot if still fresh, else ``None``.""" + ttl = self.ttl_s if ttl_s is None else ttl_s + with self._lock: + snapshot = self._snapshots.get(instance_id) + if snapshot is None: + return None + if time.time() - snapshot.updated_at > ttl: + return None + return snapshot + + def snapshots(self) -> dict[str, InstanceMetrics]: + with self._lock: + return dict(self._snapshots) + + # ---------------------------------------------------------------- + + # ---------------------------------------------------------------- + + def _loop(self) -> None: + while not self._stop_event.is_set(): + now = time.time() + with self._lock: + due = [ + (instance_id, url) + for instance_id, url in self._endpoints.items() + if now >= self._next_due.get(instance_id, 0.0) + ] + for instance_id, _ in due: + self._next_due[instance_id] = now + self.interval_s + for instance_id, url in due: + if self._stop_event.is_set(): + break + self._poll_one(instance_id, url) + self._stop_event.wait(0.2) + + def _poll_one(self, instance_id: str, url: str) -> None: + with self._lock: + registered = instance_id in self._endpoints + try: + request = urllib.request.Request(url, method="GET") + with urllib.request.urlopen(request, timeout=self.timeout_s) as response: + payload = response.read().decode("utf-8", errors="replace") + core = extract_core_metrics(parse_prometheus_metrics(payload)) + if not core: + raise ValueError("/metrics contains no recognized vLLM load gauges") + except Exception as exc: + self._record_failure(instance_id, exc) + return + + now = time.time() + with self._lock: + if self._stop_event.is_set() or (registered and self._endpoints.get(instance_id) != url): + return + previous = self._snapshots.get(instance_id) + last_preemption_at = 0.0 + preemptions_total = core.get("preemptions_total", -1.0) + if ( + previous is not None + and preemptions_total >= 0 + and previous.preemptions_total >= 0 + and preemptions_total > previous.preemptions_total + ): + # Preemption counters are cumulative; a preemption happened + # between the two polls. Pin the event time to now (worst + # case: up to one interval of uncertainty). + last_preemption_at = now + elif previous is not None and not ( + 0 <= preemptions_total < previous.preemptions_total + ): + last_preemption_at = previous.last_preemption_at + self._failures[instance_id] = 0 + self._snapshots[instance_id] = InstanceMetrics( + instance_id=instance_id, + updated_at=now, + running=int(core.get("running", -1)), + waiting=int(core.get("waiting", -1)), + gpu_cache_usage=float(core.get("gpu_cache_usage", -1.0)), + kv_cache_tokens=float(core.get("kv_cache_tokens", -1.0)), + preemptions_total=preemptions_total, + last_preemption_at=last_preemption_at, + consecutive_failures=0, + kv_capacity_tokens=extract_kv_capacity_tokens(payload), + ) + + def _record_failure(self, instance_id: str, exc: Exception) -> None: + with self._lock: + previous = self._snapshots.get(instance_id) + failures = self._failures.get(instance_id, 0) + 1 + self._failures[instance_id] = failures + if previous is not None: + previous.consecutive_failures = failures + if failures % 10 == 1: + logger.warning( + "Metrics poll failed for %s (failures=%d): %s", + instance_id, + failures, + exc, + ) diff --git a/infra/scheduling/shared_token_state.py b/infra/scheduling/shared_token_state.py new file mode 100644 index 0000000..efb6219 --- /dev/null +++ b/infra/scheduling/shared_token_state.py @@ -0,0 +1,210 @@ +"""Cross-rank shared load publication for token-aware scheduling. + +Generalizes the C-MLFQ ``SharedCMLFQLoadState`` pattern to every +scheduler: each training rank publishes its per-instance in-flight +accounting (request count + estimated token load) to a shared filesystem +using atomic renames plus a heartbeat TTL. Schedulers aggregate all live +ranks' files so instance ranking reflects cluster-wide load instead of a +single rank's partial view — without this, N independent least-connections +schedulers all see the same "empty" instance and herd traffic onto it. +""" + +from __future__ import annotations + +import atexit +import json +import logging +import os +import socket +import threading +import time +import uuid +from collections import defaultdict +from pathlib import Path + +logger = logging.getLogger(__name__) + + +class SharedTokenLoadState: + """Publish per-rank token/request accounting; aggregate across ranks. + + Each process is the only writer of its own snapshot file. Aggregate + reads sum every live file (heartbeat within TTL) so a crashed rank's + counts disappear instead of poisoning the total. A short-lived cache + (invalidated by our own writes) keeps per-schedule aggregate reads + cheap while other ranks' updates appear within ``cache_ttl_s``. + """ + + def __init__( + self, + directory: str, + ttl_s: float = 30.0, + heartbeat_interval_s: float = 10.0, + writer_id: str = "", + cache_ttl_s: float = 1.0, + ): + self.directory = Path(directory) + self.directory.mkdir(parents=True, exist_ok=True) + rank = os.environ.get("RANK", "0") + self.writer_id = writer_id or f"rank_{rank}" + if Path(self.writer_id).name != self.writer_id: + raise ValueError("writer_id must be a file name, not a path") + self._owner = uuid.uuid4().hex + self.ttl_s = max(heartbeat_interval_s * 2.0, ttl_s) + self.heartbeat_interval_s = max(1.0, heartbeat_interval_s) + self.cache_ttl_s = max(0.0, float(cache_ttl_s)) + self._path = self.directory / f"{self.writer_id}.json" + self._counts: dict[str, dict[str, int]] = {} + self._lock = threading.Lock() + self._stop_event = threading.Event() + self._cache: dict[str, dict[str, int]] | None = None + self._cache_at: float = 0.0 + self._publish() + self._heartbeat_thread = threading.Thread( + target=self._heartbeat_loop, + name=f"shared-token-load-{self.writer_id}", + daemon=True, + ) + self._heartbeat_thread.start() + atexit.register(self.close) + + # ---------------------------------------------------------------- + + # ---------------------------------------------------------------- + + def add( + self, + instance_id: str, + delta_requests: int = 0, + delta_tokens: int = 0, + ) -> None: + """Apply a delta to this rank's published accounting.""" + with self._lock: + if self._stop_event.is_set(): + return + entry = self._counts.setdefault( + instance_id, {"requests": 0, "tokens": 0} + ) + entry["requests"] = max(0, entry["requests"] + int(delta_requests)) + entry["tokens"] = max(0, entry["tokens"] + int(delta_tokens)) + self._publish_locked() + self._cache = None + + def reset(self) -> None: + """Zero this rank's accounting (e.g. a fresh scheduler attached).""" + with self._lock: + if self._stop_event.is_set(): + return + self._counts.clear() + self._publish_locked() + self._cache = None + + def totals(self) -> dict[str, dict[str, int]]: + """Cluster-wide sums over live rank files (cached). + + Returns ``{instance_id: {"requests": int, "tokens": int}}``. + """ + now = time.monotonic() + with self._lock: + if ( + self._cache is not None + and self.cache_ttl_s > 0.0 + and now - self._cache_at < min(self.cache_ttl_s, self.ttl_s) + ): + return self._with_local(self._cache) + wall = time.time() + aggregated: dict[str, dict[str, int]] = defaultdict( + lambda: {"requests": 0, "tokens": 0} + ) + try: + paths = list(self.directory.glob("*.json")) + except OSError: + paths = [] + for path in paths: + if path == self._path: + continue # own contribution comes from live memory below + try: + payload = json.loads(path.read_text(encoding="utf-8")) + if wall - float(payload.get("updated_at", 0.0)) > self.ttl_s: + continue + counts = payload.get("counts", {}) + if not isinstance(counts, dict): + continue + for instance_id, entry in counts.items(): + if not isinstance(entry, dict): + continue + total = aggregated[str(instance_id)] + total["requests"] += max(0, int(entry.get("requests", 0))) + total["tokens"] += max(0, int(entry.get("tokens", 0))) + except (OSError, ValueError, TypeError, OverflowError): + continue + result = { + instance_id: dict(entry) for instance_id, entry in aggregated.items() + } + with self._lock: + self._cache = result + self._cache_at = time.monotonic() + return self._with_local(result) + + def _with_local(self, peers): + result = {iid: dict(entry) for iid, entry in peers.items()} + for iid, entry in self._counts.items(): + total = result.setdefault(iid, {"requests": 0, "tokens": 0}) + for key in ("requests", "tokens"): + total[key] += entry[key] + return result + + def close(self) -> None: + """Stop the heartbeat and remove our file so peers drop our counts.""" + if self._stop_event.is_set(): + return + self._stop_event.set() + self._heartbeat_thread.join(timeout=self.heartbeat_interval_s + 1) + with self._lock: + self._counts.clear() + try: + payload = json.loads(self._path.read_text(encoding="utf-8")) + if payload.get("owner") == self._owner: + self._path.unlink(missing_ok=True) + except (OSError, ValueError): + pass + self._cache = None + atexit.unregister(self.close) + + # ---------------------------------------------------------------- + + # ---------------------------------------------------------------- + + def _heartbeat_loop(self) -> None: + while not self._stop_event.wait(self.heartbeat_interval_s): + try: + self._publish() + except OSError as exc: + logger.warning("Shared load heartbeat failed: %s", exc) + + def _publish(self) -> None: + with self._lock: + self._publish_locked() + + def _publish_locked(self) -> None: + if self._stop_event.is_set(): + return + payload = { + "owner": self._owner, + "writer_id": self.writer_id, + "rank": int(os.environ.get("RANK", "0")), + "pid": os.getpid(), + "hostname": socket.gethostname(), + "updated_at": time.time(), + "counts": { + iid: dict(entry) for iid, entry in self._counts.items() + }, + } + tmp_path = self._path.with_name( + f".{self._path.name}.{os.getpid()}.{threading.get_ident()}.tmp" + ) + tmp_path.write_text( + json.dumps(payload, ensure_ascii=True, sort_keys=True), + encoding="utf-8", + ) + os.replace(tmp_path, self._path) diff --git a/tests/test_scheduling_audit.py b/tests/test_scheduling_audit.py new file mode 100644 index 0000000..4e34532 --- /dev/null +++ b/tests/test_scheduling_audit.py @@ -0,0 +1,205 @@ +"""Behavioral regressions found by auditing the actual rollout call chain.""" + +import asyncio +import json +import multiprocessing +import time +from contextlib import contextmanager +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from threading import Thread + +import pytest + +from RL_Framework.engine.heterogeneous_engine import HeterogeneousRolloutEngine +from RL_Framework.infra.scheduling.base import MetricsFeedbackConfig +from RL_Framework.infra.scheduling.cmlfq_scheduler import CMLFQScheduler +from RL_Framework.infra.scheduling.la_mlfq import LAMLFQScheduler, WaitingRequest +from RL_Framework.infra.scheduling.length_aware import LengthAwareScheduler +from RL_Framework.infra.scheduling.load_balance import LoadBalanceScheduler +from RL_Framework.infra.scheduling.metrics_feed import InstanceMetrics, VLLMMetricsPoller +from RL_Framework.infra.scheduling.shared_token_state import SharedTokenLoadState + + +@pytest.mark.parametrize("scheduler_cls", [LoadBalanceScheduler, LengthAwareScheduler, LAMLFQScheduler]) +def test_short_completion_does_not_retire_long_request(scheduler_cls): + async def run(): + scheduler = scheduler_cls(load_metric="tokens") + engine = HeterogeneousRolloutEngine("m", scheduler) + engine.add_instance("a", "localhost", 1, 1) + entered, release = asyncio.Event(), asyncio.Event() + + class Endpoint: + async def generate(self, prompt, **kwargs): + if prompt == "long": + entered.set() + await release.wait() + return {"text": "x", "tokens": [1], "logprobs": []} + + engine.engines = [Endpoint()] + long = asyncio.create_task(engine.generate("long", input_tokens=9000, max_new_tokens=64)) + await entered.wait() + initial = scheduler.get_instance_handle(0).active_tokens + await engine.generate("short", input_tokens=20, max_new_tokens=8) + assert scheduler.get_instance_handle(0).active_tokens == initial + with pytest.raises(TimeoutError): + engine.wait_until_idle(timeout=0) + release.set() + await long + assert scheduler.get_instance_handle(0).active_tokens == 0 + asyncio.run(run()) + + +def test_cmlfq_direct_generation_accepts_budget_and_retires(): + async def run(): + scheduler = CMLFQScheduler() + engine = HeterogeneousRolloutEngine("m", scheduler) + engine.add_instance("a", "localhost", 1, 1) + class Endpoint: + async def generate(self, **kwargs): + return {"text": "x", "tokens": [1]} + engine.engines = [Endpoint()] + await engine.generate("p", input_tokens=10, max_new_tokens=5) + assert scheduler.get_instance_handle(0).active_requests == 0 + asyncio.run(run()) + + +@pytest.mark.parametrize("scheduler_cls", [LengthAwareScheduler, LAMLFQScheduler]) +def test_preferred_bucket_cannot_resurrect_blocked_or_unknown_capacity(scheduler_cls): + scheduler = scheduler_cls(load_metric="kv_tokens", kv_capacity_tokens_by_tp={2: 10000}) + scheduler.register_instance(0, "a", 1) + scheduler.register_instance(1, "b", 2) + assert scheduler.schedule(100).instance_index == 1 + scheduler._kv_capacity_by_tp[1] = 10000 + scheduler._feedback = MetricsFeedbackConfig(enabled=True) + class Feed: + def get(self, iid): + return InstanceMetrics(iid, time.time(), gpu_cache_usage=0.99 if iid == "a" else 0.1) + scheduler.attach_metrics_feed(Feed()) + assert scheduler.schedule(100).instance_index == 1 + + +def test_capacity_is_per_instance_not_largest_tp_peer(): + scheduler = LoadBalanceScheduler(load_metric="kv_tokens", feedback=MetricsFeedbackConfig(enabled=True)) + for index in range(2): + scheduler.register_instance(index, str(index), 1) + scheduler.get_instance_handle(index).inc_active(100, 0) + class Feed: + def get(self, iid): + return InstanceMetrics(iid, time.time(), kv_capacity_tokens=1000 if iid == "0" else 10000) + scheduler.attach_metrics_feed(Feed()) + assert scheduler.load_of(scheduler.get_instance_handle(0)) == 0.1 + assert scheduler.load_of(scheduler.get_instance_handle(1)) == 0.01 + + +def test_cancelled_scout_follower_is_not_scheduled(): + async def run(): + scheduler = LAMLFQScheduler(load_metric="tokens") + scheduler.register_instance(0, "a", 1) + future = asyncio.get_running_loop().create_future() + future.cancel() + waiter = WaitingRequest("p", 10, 2, 0, 1, future=future, max_new_tokens=8) + scheduler._process_released_waiting([waiter], "short") + assert scheduler.get_instance_handle(0).active_requests == 0 + future2 = asyncio.get_running_loop().create_future() + waiter.future = future2 + scheduler._process_released_waiting([waiter], "short") + assert future2.result().reserved_tokens == 18 + asyncio.run(run()) + + +def _rank_writer(directory, commands, replies): + state = SharedTokenLoadState(directory, writer_id="rank_1", cache_ttl_s=0) + try: + for command in iter(commands.get, "stop"): + state.add("a", *command) + replies.put("published") + finally: + state.close() + + +def test_real_process_updates_are_visible_with_cache_disabled(tmp_path): + ctx = multiprocessing.get_context("spawn") + commands, replies = ctx.Queue(), ctx.Queue() + writer = ctx.Process(target=_rank_writer, args=(str(tmp_path), commands, replies)) + reader = SharedTokenLoadState(str(tmp_path), writer_id="rank_0", cache_ttl_s=0) + writer.start() + try: + assert reader.totals() == {} # prime empty cache + commands.put((1, 9000)) + assert replies.get(timeout=20) == "published" + assert reader.totals()["a"] == {"requests": 1, "tokens": 9000} + commands.put((-1, -9000)) + replies.get(timeout=20) + assert reader.totals()["a"]["tokens"] == 0 + finally: + commands.put("stop") + writer.join(timeout=20) + if writer.is_alive(): + writer.terminate() + writer.join() + reader.close() + assert writer.exitcode == 0 + + +def test_closed_writer_cannot_resurrect_or_delete_replacement(tmp_path): + old = SharedTokenLoadState(str(tmp_path), writer_id="rank_0") + new = SharedTokenLoadState(str(tmp_path), writer_id="rank_0") + try: + old.close() + old.add("a", 1, 200) + old.reset() + new.add("a", 1, 100) + assert new.totals()["a"]["tokens"] == 100 + assert json.loads((tmp_path / "rank_0.json").read_text())["owner"] == new._owner + finally: + old.close() + new.close() + + +@contextmanager +def metrics_server(): + state = {"status": 200, "counter": 4, "usage": 0.95} + class Handler(BaseHTTPRequestHandler): + def do_GET(self): + self.send_response(state["status"]) + self.end_headers() + self.wfile.write(( + f'vllm:kv_cache_usage_perc{{model_name="name with spaces"}} {state["usage"]}\n' + f'vllm:num_preemptions_total {state["counter"]}\n' + 'vllm:num_requests_running NaN\n' + 'vllm:num_requests_waiting 3\n' + ).encode()) + def log_message(self, *args): + pass + server = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + thread = Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield state, f"http://127.0.0.1:{server.server_port}/metrics" + finally: + server.shutdown() + server.server_close() + thread.join() + + +def test_http_feed_counter_reset_nan_and_stale_admission(): + with metrics_server() as (state, url): + poller = VLLMMetricsPoller() + poller._poll_one("a", url) + assert poller.get("a").running == -1 + assert poller.get("a").gpu_cache_usage == 0.95 + state["counter"] += 1 + poller._poll_one("a", url) + assert poller.get("a").last_preemption_at > 0 + state["counter"] = 0 # restarted server + poller._poll_one("a", url) + assert poller.get("a").last_preemption_at == 0 + scheduler = LoadBalanceScheduler(feedback=MetricsFeedbackConfig(enabled=True)) + scheduler.register_instance(0, "a", 1) + scheduler.attach_metrics_feed(poller) + assert not scheduler._admission_ok(scheduler.get_instance_handle(0)) + state["status"] = 500 + poller._poll_one("a", url) + poller._snapshots["a"].updated_at -= 100 + assert scheduler._admission_ok(scheduler.get_instance_handle(0)) + assert scheduler._admission_blocked == {} diff --git a/trainer/async_rl_trainer.py b/trainer/async_rl_trainer.py index 190d3bb..48d89a8 100644 --- a/trainer/async_rl_trainer.py +++ b/trainer/async_rl_trainer.py @@ -2571,7 +2571,25 @@ def _rebind_rollout_engine_from_config(self, *, reason: str) -> None: "[RuntimeElasticExecutor] warning: failed to close old " f"rollout clients during rebind: {exc}" ) - self.rollout_engine = HeterogeneousRolloutEngine.from_config(self.config) + # Stop the old engine's /metrics poller thread so rebinding + # every sync step does not leak one poller per step. + stop_poller = getattr(old_engine, "_stop_metrics_poller", None) + if callable(stop_poller): + stop_poller() + # Same for the old cross-rank shared load state: two writers + # with the same writer_id would overwrite each other's files. + close_shared = getattr(old_engine, "_close_shared_state", None) + if callable(close_shared): + close_shared() + # Preserve learned scheduler state (output-length EMA, history + # tables) across the rebind: weight sync rebuilds the engine + # every sync step and an unconditional reset would zero the + # EMA back to its prior each time. + old_scheduler = getattr(old_engine, "scheduler", None) + self.rollout_engine = HeterogeneousRolloutEngine.from_config( + self.config, + carry_scheduler_state_from=old_scheduler, + ) else: self.rollout_engine.reconfigure_from_plan(None, self.config) if hasattr(self.rollout_engine, "wait_for_ready"): diff --git a/workflow/agentic.py b/workflow/agentic.py index ca19cfc..b012ee3 100644 --- a/workflow/agentic.py +++ b/workflow/agentic.py @@ -44,6 +44,7 @@ async def run_episode( all_actions = [] current_turn = 0 + prompt_id = str(data.get("prompt_id") or data.get("id") or "") while current_turn < self.max_turns: input_ids = self.tokenizer.apply_chat_template( @@ -53,12 +54,18 @@ async def run_episode( ) prompt_str = self.tokenizer.decode(input_ids) + if not prompt_id: + import hashlib + prompt_id = hashlib.sha256(prompt_str.encode("utf-8")).hexdigest() response = await engine.generate( prompt=prompt_str, + prompt_id=prompt_id, max_new_tokens=self.max_new_tokens, temperature=self.temperature, n=1, + # apply_chat_template already produced the exact ids. + input_tokens=len(input_ids), ) output_text = response["text"] diff --git a/workflow/code_agent.py b/workflow/code_agent.py index bd368a7..810ce97 100644 --- a/workflow/code_agent.py +++ b/workflow/code_agent.py @@ -3,6 +3,7 @@ from __future__ import annotations import re +import asyncio from typing import Any, Callable import torch @@ -134,10 +135,16 @@ async def run_episode( generate_kwargs = { + "prompt_id": prompt_id, "prompt": current_text, "max_new_tokens": self.max_new_tokens, "temperature": self.temperature, "n": 1, + # Reused segment count (approximate at BPE boundaries); skips the + # engine-side estimation encode. + "input_tokens": sum( + len(tokens) for tokens, _, _ in segments + ), } if cmlfq_request_id: generate_kwargs.update({ @@ -146,7 +153,10 @@ async def run_episode( }) try: response = await engine.generate(**generate_kwargs) - except Exception: + except BaseException: + # BaseException: CancelledError (task cancellation) must + # also release the C-MLFQ request state, otherwise the + # leaked entry wedges wait_until_idle until timeout. cancel_cmlfq = getattr( engine, "cancel_cmlfq_request", None ) @@ -174,6 +184,11 @@ async def run_episode( try: pass_rate, metadata = await self.executor.execute(code, test_cases) + except asyncio.CancelledError: + cancel = getattr(engine, "cancel_cmlfq_request", None) + if cmlfq_request_id and callable(cancel): + cancel(cmlfq_request_id) + raise except Exception as e: pass_rate = 0.0 metadata = {"error": str(e)} diff --git a/workflow/dapo_math.py b/workflow/dapo_math.py index e772c77..bf4bbfc 100644 --- a/workflow/dapo_math.py +++ b/workflow/dapo_math.py @@ -145,11 +145,15 @@ async def run_episode(self, engine: Any, data: dict[str, Any], version: int = 0) generation_prompt, prompt_len = self._fit_generation_prompt(current_text) max_tokens = self._generation_budget(prompt_len, used, turn) kwargs = { + "prompt_id": prompt_id, "prompt": generation_prompt, "max_new_tokens": max_tokens, "temperature": self.temperature, "top_p": self.top_p, "n": 1, + # Exact token count already computed by the fitter; + # passing it skips the engine-side estimation encode. + "input_tokens": prompt_len, } if cmlfq_request_id: kwargs.update({"request_id": cmlfq_request_id, "prompt_id": prompt_id}) @@ -195,7 +199,10 @@ async def run_episode(self, engine: Any, data: dict[str, Any], version: int = 0) if cmlfq_request_id and callable(route_tool_return): route_tool_return(cmlfq_request_id, tool_event, generated_tokens) current_text += output_text + "\n" + feedback + "\n" - except Exception: + except BaseException: + # BaseException: CancelledError (task cancellation) must also + # release the C-MLFQ request state, otherwise the leaked entry + # wedges wait_until_idle until timeout. cancel_cmlfq = getattr(engine, "cancel_cmlfq_request", None) if cmlfq_request_id and callable(cancel_cmlfq): cancel_cmlfq(cmlfq_request_id) @@ -266,6 +273,7 @@ async def eval_one(i: int) -> dict[str, Any]: top_p=1.0, n=1, prompt_id=self._prompt_id(row), + input_tokens=prompt_len, ) metrics = evaluate_math_completion(response.get("text", ""), self._extract_answer(row)) return {"ok": True, "reward": float(metrics["reward"]), "accurate": float(metrics["accuracy"])} diff --git a/workflow/r2e_gym.py b/workflow/r2e_gym.py index d2f0a5c..7d5d21f 100644 --- a/workflow/r2e_gym.py +++ b/workflow/r2e_gym.py @@ -208,11 +208,14 @@ async def run_episode( reserve_feedback_tokens=128 if turn < self.max_turns - 1 else 0, ) generate_kwargs = { + "prompt_id": prompt_id, "prompt": generation_prompt, "max_new_tokens": max_tokens_this_turn, "temperature": self.temperature, "top_p": self.top_p, "n": 1, + # Exact count from the fitter; skips engine estimation. + "input_tokens": generation_prompt_tokens, # Issue generation has an explicit closing delimiter. "stop": [] if self.patch_harness else ["[/ISSUE]"], "include_stop_str_in_output": True, @@ -281,7 +284,10 @@ async def run_episode( route_tool_return(cmlfq_request_id, tool_event, generated_tokens) current_text += output_text + "\n" + feedback + "\n" - except Exception: + except BaseException: + # BaseException: CancelledError (task cancellation) must also + # release the C-MLFQ request state, otherwise the leaked entry + # wedges wait_until_idle until timeout. cancel_cmlfq = getattr(engine, "cancel_cmlfq_request", None) if cmlfq_request_id and callable(cancel_cmlfq): cancel_cmlfq(cmlfq_request_id) @@ -413,6 +419,7 @@ async def eval_one(index: int) -> dict[str, Any]: "top_p": 1.0, "n": 1, "prompt_id": prompt_id, + "input_tokens": prompt_len, "stop": [] if self.patch_harness else ["[/ISSUE]"], "include_stop_str_in_output": True, } @@ -470,6 +477,13 @@ async def eval_one(index: int) -> dict[str, Any]: "metrics": metrics, "turn_metrics": turn_metrics, } + except asyncio.CancelledError: + # Task cancellation must propagate, not be converted + # into a (false) evaluation result row. + cancel_cmlfq = getattr(engine, "cancel_cmlfq_request", None) + if request_id and callable(cancel_cmlfq): + cancel_cmlfq(request_id) + raise except Exception as exc: cancel_cmlfq = getattr(engine, "cancel_cmlfq_request", None) if request_id and callable(cancel_cmlfq): diff --git a/workflow/rlvr.py b/workflow/rlvr.py index 399aba8..004dbb8 100644 --- a/workflow/rlvr.py +++ b/workflow/rlvr.py @@ -54,6 +54,7 @@ async def run_episode( response = await engine.generate( prompt=prompt_str, + prompt_id=str(data.get("prompt_id") or data.get("id") or question), max_new_tokens=self.max_new_tokens, temperature=self.temperature, n=1, diff --git a/workflow/search_r1.py b/workflow/search_r1.py index c3479aa..b0d8ec8 100644 --- a/workflow/search_r1.py +++ b/workflow/search_r1.py @@ -3,6 +3,7 @@ from __future__ import annotations import re +import asyncio from typing import Any, Callable import torch @@ -122,10 +123,18 @@ async def run_episode( generate_kwargs = { + "prompt_id": prompt_id, "prompt": current_text, "max_new_tokens": self.max_new_tokens, "temperature": self.temperature, "n": 1, + # Reused segment count (BPE boundaries may differ from a + # full-text encode); avoids another tokenization pass. + # Includes prompt + + # prior outputs + tool results); skips engine estimation. + "input_tokens": sum( + len(tokens) for tokens, _, _ in segments + ), } if cmlfq_request_id: generate_kwargs.update({ @@ -134,7 +143,10 @@ async def run_episode( }) try: response = await engine.generate(**generate_kwargs) - except Exception: + except BaseException: + # BaseException: CancelledError (task cancellation) must + # also release the C-MLFQ request state, otherwise the + # leaked entry wedges wait_until_idle until timeout. cancel_cmlfq = getattr( engine, "cancel_cmlfq_request", None ) @@ -159,6 +171,11 @@ async def run_episode( try: search_result = await self.search_tool.search(tool_query) tool_status = "success" + except asyncio.CancelledError: + cancel = getattr(engine, "cancel_cmlfq_request", None) + if cmlfq_request_id and callable(cancel): + cancel(cmlfq_request_id) + raise except Exception as exc: search_result = f"Error: {exc}" tool_status = "failure"