diff --git a/batchgen/batchgen_worker.py b/batchgen/batchgen_worker.py index c1a0d820..8eddab40 100644 --- a/batchgen/batchgen_worker.py +++ b/batchgen/batchgen_worker.py @@ -537,12 +537,6 @@ def __init__(self, args: BatchGenWorkerArgs): # SyncCoordinator instantiated lazily on first use so we don't # touch torch.distributed before it's initialized. self._sync_coordinator: Optional[SyncCoordinator] = None - # Phase 5.1b of worker decouple (issue #175): dual-path gate for the - # KVCacheManager stats tier. NATIVE=1 routes the 3 read-only stat - # helpers through `batchgen.worker.kv_manager.KVCacheManager`. - # COMPARE=1 runs both paths and asserts equal results. - self._kv_stats_native = os.environ.get("BATCHGEN_WORKER_KV_STATS_NATIVE", "0") == "1" - self._kv_stats_compare = os.environ.get("BATCHGEN_WORKER_KV_STATS_COMPARE", "0") == "1" self._kv_cache_manager: Optional[KVCacheManager] = None self._glm5_moe_cuda_graph_manager = None self._glm5_layer_cuda_graph_manager = None @@ -3013,11 +3007,10 @@ def _destroy_gpu_paged_kv_cache(self, *, empty_cuda_cache: bool = False) -> None f"Rank {self.rank}: Reset GPU allocation state for {reset_count} sequences" ) - # Phase 5.1b of worker decouple (issue #175): the 3 read-only KV stat - # helpers below route through `KVCacheManager` when - # BATCHGEN_WORKER_KV_STATS_NATIVE=1; COMPARE=1 runs both paths and - # asserts equal results. Phase 5.1c deletes the legacy bodies after - # parity validation. + # Thin delegations to `batchgen.worker.kv_manager.KVCacheManager`. + # The worker owns the canonical state; `_make_kv_cache_manager` builds + # a lazy `TorchKVStatsBackend` adapter that reads from + # `host_paged_kv_worker_view` + `gpu_paged_kv_cache_manager`. def _make_kv_cache_manager(self) -> KVCacheManager: """Lazy-construct: KV managers may not be bound at __init__ time.""" @@ -3054,19 +3047,7 @@ def _make_kv_utilization_request(self) -> KVUtilizationRequest: ) def _get_host_kv_free_pages(self) -> int: - if self._kv_stats_compare: - legacy = self._legacy_get_host_kv_free_pages() - native = self._make_kv_cache_manager().get_host_free_pages() - assert legacy == native, f"kv_stats compare mismatch: get_host_kv_free_pages legacy={legacy} native={native}" - return native if self._kv_stats_native else legacy - if self._kv_stats_native: - return self._make_kv_cache_manager().get_host_free_pages() - return self._legacy_get_host_kv_free_pages() - - def _legacy_get_host_kv_free_pages(self) -> int: - """Get current free pages from host KV cache.""" - stats = self.host_paged_kv_worker_view.get_stats() - return stats.num_free_pages + return self._make_kv_cache_manager().get_host_free_pages() def _get_or_create_gloo_group(self): """Get or create a Gloo process group for CPU tensor migrations. @@ -3088,81 +3069,10 @@ def _get_or_create_gloo_group(self): return self._gloo_migration_group def _get_host_kv_utilization(self) -> Dict[str, int]: - if self._kv_stats_compare: - legacy = self._legacy_get_host_kv_utilization() - native_obj = self._make_kv_cache_manager().get_host_utilization(self._make_kv_utilization_request()) - native = _dataclasses.asdict(native_obj) - assert legacy == native, f"kv_stats compare mismatch: get_host_kv_utilization legacy={legacy} native={native}" - return native if self._kv_stats_native else legacy - if self._kv_stats_native: - native_obj = self._make_kv_cache_manager().get_host_utilization(self._make_kv_utilization_request()) - return _dataclasses.asdict(native_obj) - return self._legacy_get_host_kv_utilization() - - def _legacy_get_host_kv_utilization(self) -> Dict[str, int]: - """Get host KV stats counting sequences with KV in host memory. - - Valid sequences = PREFILLED, ON_HOLD, and IN_DECODE (all have KV in host). - - PREFILLED: KV stored in host after prefill - - ON_HOLD: KV retained in host when evicted from GPU - - IN_DECODE: KV streams to host after each attention layer - - Free pages = Total - used by valid sequences. - - IMPORTANT: Host KV is shared per-node, so we count sequences from ALL ranks - on this node, not just this rank. - - Returns: - Dict with: rank, node_id, num_free_pages, num_total_pages, num_used_pages, free_percent - """ - stats = self.host_paged_kv_worker_view.get_stats() - - # Count pages used by sequences with KV in host on THIS NODE (all ranks on node) - # Host KV is shared across all GPUs on a node - node_id = self.rank // NUM_GPUS_PER_NODE - node_rank_start = node_id * NUM_GPUS_PER_NODE - node_rank_end = min(node_rank_start + NUM_GPUS_PER_NODE, self.world_size) - - # CRITICAL FIX: IN_DECODE sequences also have KV in host (streams after each layer) - valid_statuses = {SequenceStatus.PREFILLED, SequenceStatus.ON_HOLD, SequenceStatus.IN_DECODE} - - # Count sequences per status for detailed logging - status_counts = {status: [] for status in valid_statuses} - for rank_on_node in range(node_rank_start, node_rank_end): - for status in valid_statuses: - seqs = self.global_batch.get_sequences_for_rank_with_status(rank_on_node, status) - status_counts[status].extend(seqs) - - valid_sequences = [] - for seqs in status_counts.values(): - valid_sequences.extend(seqs) - - # Use C++ ground truth for page counts — shared memory atomic counters - # are accurate per-node, unlike per-sequence host_pages_allocated which - # is stale on non-owner ranks between metadata syncs. - used_pages = stats.num_used_pages - free_pages = stats.num_free_pages - free_percent = int((free_pages / stats.num_total_pages) * 100) if stats.num_total_pages > 0 else 100 - - if self.local_rank == 0: - logging.debug( - f"[HOST_KV_UTIL] C++ stats: used={used_pages}, free={free_pages}, " - f"total={stats.num_total_pages}, {len(valid_sequences)} valid seqs" - ) - - return { - 'rank': self.rank, - 'node_id': self.rank // NUM_GPUS_PER_NODE, - 'num_free_pages': free_pages, - 'num_total_pages': stats.num_total_pages, - 'num_used_pages': used_pages, - 'free_percent': free_percent, - # Include sequence counts for global aggregation - 'num_in_decode': len(status_counts[SequenceStatus.IN_DECODE]), - 'num_onhold': len(status_counts[SequenceStatus.ON_HOLD]), - 'num_prefilled': len(status_counts[SequenceStatus.PREFILLED]), - 'num_valid_sequences': len(valid_sequences), - } + native_obj = self._make_kv_cache_manager().get_host_utilization( + self._make_kv_utilization_request() + ) + return _dataclasses.asdict(native_obj) def _gather_host_kv_stats_by_node(self, worker_view: Optional[object]) -> List[Dict[str, int]]: """Gather one host-KV pool stat record per node. @@ -3933,21 +3843,7 @@ def _rebalance_host_kv(self) -> None: ) def _get_gpu_kv_free_pages(self) -> int: - if self._kv_stats_compare: - legacy = self._legacy_get_gpu_kv_free_pages() - native = self._make_kv_cache_manager().get_gpu_free_pages() - assert legacy == native, f"kv_stats compare mismatch: get_gpu_kv_free_pages legacy={legacy} native={native}" - return native if self._kv_stats_native else legacy - if self._kv_stats_native: - return self._make_kv_cache_manager().get_gpu_free_pages() - return self._legacy_get_gpu_kv_free_pages() - - def _legacy_get_gpu_kv_free_pages(self) -> int: - """Get current free pages from GPU KV cache.""" - manager = self.gpu_paged_kv_cache_manager - if manager is None: - return 0 - return manager.get_stats().num_free_pages + return self._make_kv_cache_manager().get_gpu_free_pages() # ============ Main Entry Point ============