From f482da879e51c25a404342e39d43e98579f0289b Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Wed, 21 Jan 2026 11:25:30 -0600 Subject: [PATCH 1/5] Track off-policyness across RL steps --- .../data_parallel_inference_coordinator.py | 1 + .../core/inference/engines/dynamic_engine.py | 6 ++ megatron/core/inference/headers.py | 1 + megatron/core/inference/inference_client.py | 6 ++ megatron/core/inference/inference_request.py | 76 +++++++++++++++++++ .../endpoints/chat_completions.py | 9 +++ 6 files changed, 99 insertions(+) diff --git a/megatron/core/inference/data_parallel_inference_coordinator.py b/megatron/core/inference/data_parallel_inference_coordinator.py index 93e30f6aa25..d5cf1605d90 100644 --- a/megatron/core/inference/data_parallel_inference_coordinator.py +++ b/megatron/core/inference/data_parallel_inference_coordinator.py @@ -245,6 +245,7 @@ def start(self): Headers.UNPAUSE, Headers.SUSPEND, Headers.RESUME, + Headers.INCREMENT_STALENESS, Headers.STOP, ]: # control signals for the engine diff --git a/megatron/core/inference/engines/dynamic_engine.py b/megatron/core/inference/engines/dynamic_engine.py index eb60453c4b0..30f3885fe3f 100644 --- a/megatron/core/inference/engines/dynamic_engine.py +++ b/megatron/core/inference/engines/dynamic_engine.py @@ -1617,6 +1617,12 @@ def schedule_requests(self) -> int: self.suspend_signal = True elif header == Headers.RESUME: self.suspend_signal = False + elif header == Headers.INCREMENT_STALENESS: + waiting = set(self.waiting_request_ids) + for request_id, entry in self.requests.items(): + entry.record.increment_staleness( + policy_only=request_id in waiting, + ) elif header == Headers.STOP: self.received_stop = True else: diff --git a/megatron/core/inference/headers.py b/megatron/core/inference/headers.py index a22d1328679..2551bc54f53 100644 --- a/megatron/core/inference/headers.py +++ b/megatron/core/inference/headers.py @@ -17,6 +17,7 @@ class Headers(Enum): UNPAUSE = auto() SUSPEND = auto() RESUME = auto() + INCREMENT_STALENESS = auto() STOP = auto() STOP_ACK = auto() diff --git a/megatron/core/inference/inference_client.py b/megatron/core/inference/inference_client.py index a927a393b8c..5a4ed24bebb 100644 --- a/megatron/core/inference/inference_client.py +++ b/megatron/core/inference/inference_client.py @@ -213,6 +213,11 @@ def unpause_engines(self) -> None: self.running.set() self._send_signal_to_engines(Headers.UNPAUSE) + def increment_staleness(self): + """Sends a signal to increment staleness on all in-flight requests.""" + assert self.paused.is_set(), "Can only increment staleness while engines are paused." + self._send_signal_to_engines(Headers.INCREMENT_STALENESS) + def suspend_engines(self): """Sends a signal to pause all inference engines.""" self._send_signal_to_engines(Headers.PAUSE) @@ -220,6 +225,7 @@ def suspend_engines(self): def resume_engines(self): """Sends a signal to unpause all inference engines.""" + self.paused.clear() self._send_signal_to_engines(Headers.RESUME) self._send_signal_to_engines(Headers.UNPAUSE) diff --git a/megatron/core/inference/inference_request.py b/megatron/core/inference/inference_request.py index dc59fa32f27..0c1443185d1 100644 --- a/megatron/core/inference/inference_request.py +++ b/megatron/core/inference/inference_request.py @@ -287,6 +287,8 @@ class DynamicInferenceRequest(InferenceRequest): prompt_tokens: Optional[torch.Tensor] = None # remaining prompt tokens are used for chunked prefill remaining_prompt_tokens: Optional[torch.Tensor] = None + policy_staleness: Optional[torch.Tensor] = None + kv_cache_staleness: Optional[torch.Tensor] = None latency: Optional[float] = None # routing_indices stores MoE routing decisions for all tokens generated so far. # Shape: [total_tokens, num_layers, topk] - accumulated across all generation steps @@ -491,6 +493,55 @@ def request_id(self) -> int: """ return self.requests[0].request_id + @staticmethod + def _update_staleness_tensor( + tensor: Optional[torch.Tensor], total_tokens: int, increment: bool = True, + ) -> torch.Tensor: + """Update a per-token staleness tensor, extending with zeros if needed. + + Args: + tensor: Existing staleness tensor, or None to create a new one. + total_tokens: Expected length of the tensor after update. + increment: If True, increment all values by 1 (including new positions). + """ + if tensor is None: + tensor = torch.zeros(total_tokens, dtype=torch.int32, device='cpu') + elif len(tensor) < total_tokens: + tensor = torch.cat( + ( + tensor, + torch.zeros( + total_tokens - len(tensor), + dtype=tensor.dtype, + device=tensor.device, + ), + ), + dim=0, + ) + if increment: + tensor = tensor + 1 + return tensor + + def increment_staleness(self, policy_only: bool = False): + """Increment per-token staleness counters in-place. + + Each call indicates that a training step has occurred since these tokens + were generated. Tokens not yet tracked are initialized to 1. + + Args: + policy_only: If True, only increment policy_staleness. Use this for + evicted requests that have no KV cache to age. + """ + request = self[-1] + total_tokens = len(request.prompt_tokens) + len(request.generated_tokens) + request.policy_staleness = self._update_staleness_tensor( + request.policy_staleness, total_tokens, increment=True, + ) + if not policy_only: + request.kv_cache_staleness = self._update_staleness_tensor( + request.kv_cache_staleness, total_tokens, increment=True, + ) + def checkpoint(self, tokenizer: MegatronTokenizer | None = None): """Maintain reference to previous request, and then append a new request that concatenates the previous prompt and generations. @@ -501,6 +552,18 @@ def checkpoint(self, tokenizer: MegatronTokenizer | None = None): old_request = self[-1] + total_tokens = len(old_request.prompt_tokens) + len(old_request.generated_tokens) + + # Carry forward policy_staleness without incrementing. + policy_staleness = self._update_staleness_tensor( + old_request.policy_staleness, total_tokens, increment=False, + ) if old_request.policy_staleness is not None else None + + # Reset kv_cache_staleness to 0. + kv_cache_staleness = self._update_staleness_tensor( + None, total_tokens, increment=False, + ) if old_request.kv_cache_staleness is not None else None + # New prompt (concatenate prompt + generated tokens). new_prompt_tokens = torch.cat( ( @@ -530,6 +593,8 @@ def checkpoint(self, tokenizer: MegatronTokenizer | None = None): request_id=old_request.request_id, prompt_tokens=new_prompt_tokens, sampling_params=new_sampling_params, + policy_staleness=policy_staleness, + kv_cache_staleness=kv_cache_staleness, ) # Preserve event_add_engine from old request if it exists, otherwise set it. # This ensures TTFT calculation works correctly for evicted/resumed requests. @@ -566,6 +631,15 @@ def merge_lists(key): except TypeError as e: # generally means r.generated_text is None generated_text = None + # Ensure staleness tensors are always materialized (zeros if never incremented). + total_tokens = len(prompt_tokens) + len(generated_tokens) + policy_staleness = self._update_staleness_tensor( + self.requests[-1].policy_staleness, total_tokens, increment=False, + ) + kv_cache_staleness = self._update_staleness_tensor( + self.requests[-1].kv_cache_staleness, total_tokens, increment=False, + ) + # Merged request. request = DynamicInferenceRequest( request_id=self.requests[0].request_id, @@ -579,6 +653,8 @@ def merge_lists(key): generated_log_probs=merge_lists("generated_log_probs"), generated_top_n_logprobs=merge_lists("generated_top_n_logprobs"), sampling_params=self.requests[0].sampling_params, + policy_staleness=policy_staleness, + kv_cache_staleness=kv_cache_staleness, ttft=self.requests[0].ttft, tpot=merge_lists("tpot"), status=self.requests[-1].status, diff --git a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/chat_completions.py b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/chat_completions.py index 8d7372bcaba..a4cb61fb962 100644 --- a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/chat_completions.py +++ b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/chat_completions.py @@ -201,6 +201,15 @@ async def chat_completions(): "tool_calls" if metadata.get("tool_calls", []) else "stop" ), # Original code hardcoded this. } + if result.get("policy_staleness") is not None: + choice_data["policy_staleness"] = result["policy_staleness"] + if result.get("kv_cache_staleness") is not None: + choice_data["kv_cache_staleness"] = result["kv_cache_staleness"] + events = result.get("events") + if events is not None: + num_evictions = sum(1 for e in events if e.get("type") == "EVICT") + if num_evictions > 0: + choice_data["num_evictions"] = num_evictions if current_app.config['verbose']: logging.info(result) if result["routing_indices"] is not None: From 925a4a3e59407d27e63c3718f5dceef97121a446 Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Thu, 19 Feb 2026 16:03:15 -0600 Subject: [PATCH 2/5] Add unit tests --- .../inference/engines/test_dynamic_engine.py | 192 ++++++++++++++++++ 1 file changed, 192 insertions(+) diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index be33b9257ad..ca6325cf896 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -1847,3 +1847,195 @@ def test_suspend_resume_cycle(self, kv_cache_management_mode, static_kv_memory_p f"Tensor address must be stable when static_kv_memory_pointers is set. " f"Before: {addr_before:#x}, After: {addr_after:#x}" ) + + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @pytest.mark.parametrize("use_checkpoint", [False, True], ids=["persist", "recompute"]) + @torch.inference_mode() + def test_staleness_tracking(self, use_checkpoint): + """End-to-end staleness tracking through generation with checkpoint() cycles. + + Exercises increment_staleness(), checkpoint(), and merge() on live + engine request records. The ``use_checkpoint`` flag mirrors what the + engine does during suspend in RECOMPUTE mode (calls record.checkpoint() + on every active record). + + Timeline (prompt_len=8, num_tokens_to_generate=8): + + generate 3 tokens + training step 1 -> increment_staleness (init both staleness to 1) + [checkpoint if recompute: policy carried forward, kv_cache reset to 0] + generate 3 tokens + training step 2 -> increment_staleness + [checkpoint if recompute: policy carried forward, kv_cache reset to 0] + generate final 2 tokens -> requests finish + merge and validate + + Both paths produce identical policy_staleness. kv_cache_staleness + differs: persist keeps cumulative counts, recompute resets to 0 at + each checkpoint. + """ + PROMPT_LEN = 8 + NUM_TOKENS = 8 + + test_config = DynamicEngineTestConfig( + num_requests=0, + min_prompt_length=PROMPT_LEN, + max_prompt_length=PROMPT_LEN, + num_tokens_to_generate=NUM_TOKENS, + ) + env = self._build_test_env(test_config) + engine = env.engine + + # Add requests with termination_id=-1 to disable early stopping. + for i in range(2): + prompt_tokens = torch.randint( + 0, + test_config.vocab_size - 1, + (PROMPT_LEN,), + dtype=torch.int64, + device=torch.cuda.current_device(), + ) + engine._add_request( + DynamicInferenceRequest( + request_id=i, + prompt_tokens=prompt_tokens, + sampling_params=SamplingParams( + num_tokens_to_generate=NUM_TOKENS, termination_id=-1 + ), + ) + ) + + # -- Generate 3 tokens -- + for _ in range(3): + engine.step_modern() + + assert len(engine.requests) == 2 + for entry in engine.requests.values(): + assert len(entry.record[-1].generated_tokens) == 3 + assert entry.record[-1].policy_staleness is None + assert entry.record[-1].kv_cache_staleness is None + + # -- Training step 1: first increment initializes both staleness to all 1s -- + for entry in engine.requests.values(): + entry.record.increment_staleness() + + for entry in engine.requests.values(): + ps = entry.record[-1].policy_staleness + ks = entry.record[-1].kv_cache_staleness + assert ps.shape == ks.shape == (PROMPT_LEN + 3,) + assert ps.dtype == ks.dtype == torch.int32 + assert (ps == 1).all() + assert (ks == 1).all() + + # -- Checkpoint (mirrors what engine.suspend does in RECOMPUTE mode) -- + # policy_staleness is carried forward; kv_cache_staleness resets to 0. + if use_checkpoint: + for entry in engine.requests.values(): + old_req = entry.record[-1] + event_add_engine = old_req.event_add_engine + entry.record.checkpoint() + # Carry forward event_add_engine so the engine can compute TTFT + # for the first post-checkpoint token without crashing. + entry.record[-1].event_add_engine = event_add_engine + + for entry in engine.requests.values(): + ps = entry.record[-1].policy_staleness + ks = entry.record[-1].kv_cache_staleness + assert ps.shape == (PROMPT_LEN + 3,) + assert (ps == 1).all() + assert ks.shape == (PROMPT_LEN + 3,) + if use_checkpoint: + assert (ks == 0).all() + else: + assert (ks == 1).all() + + # -- Generate 3 more tokens -- + for _ in range(3): + engine.step_modern() + + assert len(engine.requests) == 2 + + # -- Training step 2: old tokens +1, new tokens init to 1 -- + for entry in engine.requests.values(): + entry.record.increment_staleness() + + for entry in engine.requests.values(): + ps = entry.record[-1].policy_staleness + ks = entry.record[-1].kv_cache_staleness + assert ps.shape == ks.shape == (PROMPT_LEN + 6,) + # policy_staleness is the same in both paths. + assert (ps[: PROMPT_LEN + 3] == 2).all() + assert (ps[PROMPT_LEN + 3 :] == 1).all() + # kv_cache_staleness differs: recompute had a reset before this increment. + if use_checkpoint: + assert (ks == 1).all() # 0+1 for old tokens, 0+1 for new tokens + else: + assert (ks[: PROMPT_LEN + 3] == 2).all() + assert (ks[PROMPT_LEN + 3 :] == 1).all() + + # -- Checkpoint -- + if use_checkpoint: + for entry in engine.requests.values(): + old_req = entry.record[-1] + event_add_engine = old_req.event_add_engine + entry.record.checkpoint() + entry.record[-1].event_add_engine = event_add_engine + + for entry in engine.requests.values(): + ps = entry.record[-1].policy_staleness + ks = entry.record[-1].kv_cache_staleness + assert ps.shape == ks.shape == (PROMPT_LEN + 6,) + assert (ps[: PROMPT_LEN + 3] == 2).all() + assert (ps[PROMPT_LEN + 3 :] == 1).all() + assert (ks == 0).all() # reset again + + # -- Generate remaining 2 tokens, collect finished records -- + finished_records = [] + while engine.has_unfinished_requests(): + result = engine.step_modern() + finished_records.extend(result["finished_request_records"]) + + assert len(finished_records) == 2 + + # -- Validate merged results -- + for record in finished_records: + merged = record.merge() + + # policy_staleness is identical in both paths. + # merge() materializes staleness to cover all tokens, including + # the final 2 generated after the last increment (staleness 0). + assert merged.policy_staleness is not None + assert merged.policy_staleness.shape == (PROMPT_LEN + NUM_TOKENS,) + assert (merged.policy_staleness[: PROMPT_LEN + 3] == 2).all() + assert (merged.policy_staleness[PROMPT_LEN + 3 : PROMPT_LEN + 6] == 1).all() + assert (merged.policy_staleness[PROMPT_LEN + 6 :] == 0).all() + + # kv_cache_staleness differs between persist and recompute. + assert merged.kv_cache_staleness is not None + assert merged.kv_cache_staleness.shape == (PROMPT_LEN + NUM_TOKENS,) + if use_checkpoint: + assert (merged.kv_cache_staleness == 0).all() # last checkpoint reset it + else: + assert (merged.kv_cache_staleness[: PROMPT_LEN + 3] == 2).all() + assert (merged.kv_cache_staleness[PROMPT_LEN + 3 : PROMPT_LEN + 6] == 1).all() + assert (merged.kv_cache_staleness[PROMPT_LEN + 6 :] == 0).all() + + # Original prompt preserved, all 8 tokens generated. + assert len(merged.prompt_tokens) == PROMPT_LEN + assert len(merged.generated_tokens) == NUM_TOKENS + + # -- Verify evicted requests skip kv_cache increment -- + # Eviction always calls checkpoint() (regardless of KV mode), which + # resets kv_cache_staleness. A subsequent increment with policy_only=True + # (what the engine does for waiting-queue requests) should only bump + # policy_staleness. + record = finished_records[0] + record.checkpoint() + pre_ps = record[-1].policy_staleness.clone() + + record.increment_staleness(policy_only=True) + + assert (record[-1].policy_staleness == pre_ps + 1).all() + assert (record[-1].kv_cache_staleness == 0).all() From a9c2d6e34aa71bcfe6952c4c7a57041d77568a60 Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Fri, 20 Feb 2026 02:28:46 -0600 Subject: [PATCH 3/5] Bundle in requested bugfix --- .../dynamic_text_gen_server/flask_server.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/flask_server.py b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/flask_server.py index 46818ffae31..73b9684ad48 100644 --- a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/flask_server.py +++ b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/flask_server.py @@ -85,7 +85,7 @@ def health_check(): logger.info(f"Using parsers: {parsers}") loop.set_default_executor(ThreadPoolExecutor(max_workers=8192)) - await serve(AsyncioWSGIMiddleware(app), config) + await serve(AsyncioWSGIMiddleware(app, max_body_size=config.wsgi_max_body_size), config) @trace_async_exceptions From 1d227b2ff5212fa882696118b2e9a189506582de Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Fri, 20 Feb 2026 02:29:38 -0600 Subject: [PATCH 4/5] lint --- .../core/inference/engines/dynamic_engine.py | 4 +-- megatron/core/inference/inference_request.py | 32 +++++++++++-------- 2 files changed, 19 insertions(+), 17 deletions(-) diff --git a/megatron/core/inference/engines/dynamic_engine.py b/megatron/core/inference/engines/dynamic_engine.py index 30f3885fe3f..8fe9f027cdc 100644 --- a/megatron/core/inference/engines/dynamic_engine.py +++ b/megatron/core/inference/engines/dynamic_engine.py @@ -1620,9 +1620,7 @@ def schedule_requests(self) -> int: elif header == Headers.INCREMENT_STALENESS: waiting = set(self.waiting_request_ids) for request_id, entry in self.requests.items(): - entry.record.increment_staleness( - policy_only=request_id in waiting, - ) + entry.record.increment_staleness(policy_only=request_id in waiting) elif header == Headers.STOP: self.received_stop = True else: diff --git a/megatron/core/inference/inference_request.py b/megatron/core/inference/inference_request.py index 0c1443185d1..328a77aeed0 100644 --- a/megatron/core/inference/inference_request.py +++ b/megatron/core/inference/inference_request.py @@ -495,7 +495,7 @@ def request_id(self) -> int: @staticmethod def _update_staleness_tensor( - tensor: Optional[torch.Tensor], total_tokens: int, increment: bool = True, + tensor: Optional[torch.Tensor], total_tokens: int, increment: bool = True ) -> torch.Tensor: """Update a per-token staleness tensor, extending with zeros if needed. @@ -511,9 +511,7 @@ def _update_staleness_tensor( ( tensor, torch.zeros( - total_tokens - len(tensor), - dtype=tensor.dtype, - device=tensor.device, + total_tokens - len(tensor), dtype=tensor.dtype, device=tensor.device ), ), dim=0, @@ -535,11 +533,11 @@ def increment_staleness(self, policy_only: bool = False): request = self[-1] total_tokens = len(request.prompt_tokens) + len(request.generated_tokens) request.policy_staleness = self._update_staleness_tensor( - request.policy_staleness, total_tokens, increment=True, + request.policy_staleness, total_tokens, increment=True ) if not policy_only: request.kv_cache_staleness = self._update_staleness_tensor( - request.kv_cache_staleness, total_tokens, increment=True, + request.kv_cache_staleness, total_tokens, increment=True ) def checkpoint(self, tokenizer: MegatronTokenizer | None = None): @@ -555,14 +553,20 @@ def checkpoint(self, tokenizer: MegatronTokenizer | None = None): total_tokens = len(old_request.prompt_tokens) + len(old_request.generated_tokens) # Carry forward policy_staleness without incrementing. - policy_staleness = self._update_staleness_tensor( - old_request.policy_staleness, total_tokens, increment=False, - ) if old_request.policy_staleness is not None else None + policy_staleness = ( + self._update_staleness_tensor( + old_request.policy_staleness, total_tokens, increment=False + ) + if old_request.policy_staleness is not None + else None + ) # Reset kv_cache_staleness to 0. - kv_cache_staleness = self._update_staleness_tensor( - None, total_tokens, increment=False, - ) if old_request.kv_cache_staleness is not None else None + kv_cache_staleness = ( + self._update_staleness_tensor(None, total_tokens, increment=False) + if old_request.kv_cache_staleness is not None + else None + ) # New prompt (concatenate prompt + generated tokens). new_prompt_tokens = torch.cat( @@ -634,10 +638,10 @@ def merge_lists(key): # Ensure staleness tensors are always materialized (zeros if never incremented). total_tokens = len(prompt_tokens) + len(generated_tokens) policy_staleness = self._update_staleness_tensor( - self.requests[-1].policy_staleness, total_tokens, increment=False, + self.requests[-1].policy_staleness, total_tokens, increment=False ) kv_cache_staleness = self._update_staleness_tensor( - self.requests[-1].kv_cache_staleness, total_tokens, increment=False, + self.requests[-1].kv_cache_staleness, total_tokens, increment=False ) # Merged request. From 6ba8d6d525b3b25ca999b1427e5d14c030a340df Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Fri, 20 Feb 2026 02:48:44 -0600 Subject: [PATCH 5/5] Cleanup --- .../inference/engines/test_dynamic_engine.py | 87 ++++--------------- 1 file changed, 16 insertions(+), 71 deletions(-) diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index ca6325cf896..d71ccccd49a 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -1854,27 +1854,8 @@ def test_suspend_resume_cycle(self, kv_cache_management_mode, static_kv_memory_p @pytest.mark.parametrize("use_checkpoint", [False, True], ids=["persist", "recompute"]) @torch.inference_mode() def test_staleness_tracking(self, use_checkpoint): - """End-to-end staleness tracking through generation with checkpoint() cycles. - - Exercises increment_staleness(), checkpoint(), and merge() on live - engine request records. The ``use_checkpoint`` flag mirrors what the - engine does during suspend in RECOMPUTE mode (calls record.checkpoint() - on every active record). - - Timeline (prompt_len=8, num_tokens_to_generate=8): - - generate 3 tokens - training step 1 -> increment_staleness (init both staleness to 1) - [checkpoint if recompute: policy carried forward, kv_cache reset to 0] - generate 3 tokens - training step 2 -> increment_staleness - [checkpoint if recompute: policy carried forward, kv_cache reset to 0] - generate final 2 tokens -> requests finish - merge and validate - - Both paths produce identical policy_staleness. kv_cache_staleness - differs: persist keeps cumulative counts, recompute resets to 0 at - each checkpoint. + """Test that staleness is correct tracked. + The use_checkpoint parameter simulates the behavior of different kv_cache_management_mode. """ PROMPT_LEN = 8 NUM_TOKENS = 8 @@ -1888,7 +1869,6 @@ def test_staleness_tracking(self, use_checkpoint): env = self._build_test_env(test_config) engine = env.engine - # Add requests with termination_id=-1 to disable early stopping. for i in range(2): prompt_tokens = torch.randint( 0, @@ -1907,17 +1887,15 @@ def test_staleness_tracking(self, use_checkpoint): ) ) - # -- Generate 3 tokens -- for _ in range(3): engine.step_modern() - assert len(engine.requests) == 2 for entry in engine.requests.values(): assert len(entry.record[-1].generated_tokens) == 3 assert entry.record[-1].policy_staleness is None assert entry.record[-1].kv_cache_staleness is None - # -- Training step 1: first increment initializes both staleness to all 1s -- + # Increment staleness. for entry in engine.requests.values(): entry.record.increment_staleness() @@ -1925,39 +1903,29 @@ def test_staleness_tracking(self, use_checkpoint): ps = entry.record[-1].policy_staleness ks = entry.record[-1].kv_cache_staleness assert ps.shape == ks.shape == (PROMPT_LEN + 3,) - assert ps.dtype == ks.dtype == torch.int32 assert (ps == 1).all() assert (ks == 1).all() - # -- Checkpoint (mirrors what engine.suspend does in RECOMPUTE mode) -- - # policy_staleness is carried forward; kv_cache_staleness resets to 0. + # Simulate RECOMPUTE if use_checkpoint: for entry in engine.requests.values(): old_req = entry.record[-1] event_add_engine = old_req.event_add_engine entry.record.checkpoint() - # Carry forward event_add_engine so the engine can compute TTFT - # for the first post-checkpoint token without crashing. + # Prevent TTFT crash due to missing _add_request in test. entry.record[-1].event_add_engine = event_add_engine - for entry in engine.requests.values(): - ps = entry.record[-1].policy_staleness - ks = entry.record[-1].kv_cache_staleness - assert ps.shape == (PROMPT_LEN + 3,) - assert (ps == 1).all() - assert ks.shape == (PROMPT_LEN + 3,) - if use_checkpoint: + for entry in engine.requests.values(): + ps = entry.record[-1].policy_staleness + ks = entry.record[-1].kv_cache_staleness + assert ps.shape == ks.shape == (PROMPT_LEN + 3,) + assert (ps == 1).all() assert (ks == 0).all() - else: - assert (ks == 1).all() - # -- Generate 3 more tokens -- for _ in range(3): engine.step_modern() - assert len(engine.requests) == 2 - - # -- Training step 2: old tokens +1, new tokens init to 1 -- + # Increment staleness. for entry in engine.requests.values(): entry.record.increment_staleness() @@ -1965,17 +1933,14 @@ def test_staleness_tracking(self, use_checkpoint): ps = entry.record[-1].policy_staleness ks = entry.record[-1].kv_cache_staleness assert ps.shape == ks.shape == (PROMPT_LEN + 6,) - # policy_staleness is the same in both paths. assert (ps[: PROMPT_LEN + 3] == 2).all() assert (ps[PROMPT_LEN + 3 :] == 1).all() - # kv_cache_staleness differs: recompute had a reset before this increment. if use_checkpoint: - assert (ks == 1).all() # 0+1 for old tokens, 0+1 for new tokens + assert (ks == 1).all() else: assert (ks[: PROMPT_LEN + 3] == 2).all() assert (ks[PROMPT_LEN + 3 :] == 1).all() - # -- Checkpoint -- if use_checkpoint: for entry in engine.requests.values(): old_req = entry.record[-1] @@ -1984,58 +1949,38 @@ def test_staleness_tracking(self, use_checkpoint): entry.record[-1].event_add_engine = event_add_engine for entry in engine.requests.values(): - ps = entry.record[-1].policy_staleness ks = entry.record[-1].kv_cache_staleness - assert ps.shape == ks.shape == (PROMPT_LEN + 6,) - assert (ps[: PROMPT_LEN + 3] == 2).all() - assert (ps[PROMPT_LEN + 3 :] == 1).all() - assert (ks == 0).all() # reset again + assert (ks == 0).all() - # -- Generate remaining 2 tokens, collect finished records -- finished_records = [] while engine.has_unfinished_requests(): result = engine.step_modern() finished_records.extend(result["finished_request_records"]) - assert len(finished_records) == 2 - - # -- Validate merged results -- for record in finished_records: merged = record.merge() - # policy_staleness is identical in both paths. - # merge() materializes staleness to cover all tokens, including - # the final 2 generated after the last increment (staleness 0). assert merged.policy_staleness is not None assert merged.policy_staleness.shape == (PROMPT_LEN + NUM_TOKENS,) assert (merged.policy_staleness[: PROMPT_LEN + 3] == 2).all() assert (merged.policy_staleness[PROMPT_LEN + 3 : PROMPT_LEN + 6] == 1).all() assert (merged.policy_staleness[PROMPT_LEN + 6 :] == 0).all() - # kv_cache_staleness differs between persist and recompute. assert merged.kv_cache_staleness is not None assert merged.kv_cache_staleness.shape == (PROMPT_LEN + NUM_TOKENS,) if use_checkpoint: - assert (merged.kv_cache_staleness == 0).all() # last checkpoint reset it + assert (merged.kv_cache_staleness == 0).all() else: assert (merged.kv_cache_staleness[: PROMPT_LEN + 3] == 2).all() assert (merged.kv_cache_staleness[PROMPT_LEN + 3 : PROMPT_LEN + 6] == 1).all() assert (merged.kv_cache_staleness[PROMPT_LEN + 6 :] == 0).all() - # Original prompt preserved, all 8 tokens generated. - assert len(merged.prompt_tokens) == PROMPT_LEN - assert len(merged.generated_tokens) == NUM_TOKENS - - # -- Verify evicted requests skip kv_cache increment -- - # Eviction always calls checkpoint() (regardless of KV mode), which - # resets kv_cache_staleness. A subsequent increment with policy_only=True - # (what the engine does for waiting-queue requests) should only bump - # policy_staleness. + # Verify evicted requests don't have their policy staleness incremented. record = finished_records[0] record.checkpoint() pre_ps = record[-1].policy_staleness.clone() - record.increment_staleness(policy_only=True) + record.increment_staleness(policy_only=True) # This mimics the coordinator's action. assert (record[-1].policy_staleness == pre_ps + 1).all() assert (record[-1].kv_cache_staleness == 0).all()