[CI] Fix the DSpark dp-tier unit test fixture after #34919 (#38314)

This commit is contained in:
Shuwen Wang
2026-09-07 03:16:00 -07:00
committed by GitHub
parent 8ae962021d
commit 755f97c622
@@ -62,6 +62,7 @@ class TestDraftDpSyncMetadata(CustomTestCase):
num_tokens_for_logprob_per_req=1,
)
proposer.draft_model_runner = SimpleNamespace(device="cpu")
proposer._num_token_non_padded = torch.empty((1,), dtype=torch.int32)
forward_batch = SimpleNamespace(input_ids=torch.arange(6))
batch = SimpleNamespace(
@@ -81,8 +82,13 @@ class TestDraftDpSyncMetadata(CustomTestCase):
[1, 3, 0, 2],
)
self.assertEqual(forward_batch.global_num_tokens_cpu, [6, 18, 0, 12])
# Metadata fill sets only the invariant GLOBAL count; the LOCAL
# num_token_non_padded is derived later when the draft forward localizes.
# The LOCAL count reuses the proposer's persistent buffer; the GLOBAL
# count is the invariant every DP rank agrees on.
self.assertIs(
forward_batch.num_token_non_padded, proposer._num_token_non_padded
)
self.assertEqual(forward_batch.num_token_non_padded.item(), 6)
self.assertEqual(forward_batch.num_token_non_padded_cpu, 6)
self.assertEqual(forward_batch.global_num_token_non_padded.item(), 6)
self.assertEqual(forward_batch.global_num_token_non_padded.dtype, torch.int32)
self.assertEqual(forward_batch.global_num_token_non_padded_cpu, 6)