From c6d9d73fd6eb2cf9d3d5f426146c44b5863c7265 Mon Sep 17 00:00:00 2001 From: Michael <13900043+michaelzhang-ai@users.noreply.github.com> Date: Mon, 15 Jun 2026 19:14:09 -0700 Subject: [PATCH] [Spec][test] fix(kv_canary): assert draft-extend-v2 oracle tokens in token_oracle test (#28325) --- .../kv_canary/test_self_unit_token_oracle.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/test/registered/kv_canary/test_self_unit_token_oracle.py b/test/registered/kv_canary/test_self_unit_token_oracle.py index 467cd8cc4..67255fc93 100644 --- a/test/registered/kv_canary/test_self_unit_token_oracle.py +++ b/test/registered/kv_canary/test_self_unit_token_oracle.py @@ -46,8 +46,18 @@ class TestTokenOracleManager(CustomTestCase): expected_inputs_out=expected_inputs, ) + # draft-extend-v2 is not an extend mode, so fill_expected_inputs records + # deterministic oracle tokens (not input_ids). The req row [3, 7] expands + # to num_tokens_per_req=4 draft tokens each -> one row per draft token. + expected_generalized_req_ids = torch.tensor( + [3, 3, 3, 3, 7, 7, 7, 7], dtype=torch.int64, device=self.device + ) + expected_tokens = manager.oracle.expected_tokens( + generalized_req_ids=expected_generalized_req_ids, + positions=forward_batch.positions.to(torch.int64), + ) self.assertTrue( - torch.equal(expected_inputs.tokens[:8], forward_batch.input_ids) + torch.equal(expected_inputs.tokens[:8], expected_tokens.to(torch.int64)) ) self.assertTrue( torch.equal(expected_inputs.positions[:8], forward_batch.positions)