@@ -62,6 +62,7 @@ class TestDraftDpSyncMetadata(CustomTestCase):
|
|||||||
num_tokens_for_logprob_per_req=1,
|
num_tokens_for_logprob_per_req=1,
|
||||||
)
|
)
|
||||||
proposer.draft_model_runner = SimpleNamespace(device="cpu")
|
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))
|
forward_batch = SimpleNamespace(input_ids=torch.arange(6))
|
||||||
batch = SimpleNamespace(
|
batch = SimpleNamespace(
|
||||||
@@ -81,8 +82,13 @@ class TestDraftDpSyncMetadata(CustomTestCase):
|
|||||||
[1, 3, 0, 2],
|
[1, 3, 0, 2],
|
||||||
)
|
)
|
||||||
self.assertEqual(forward_batch.global_num_tokens_cpu, [6, 18, 0, 12])
|
self.assertEqual(forward_batch.global_num_tokens_cpu, [6, 18, 0, 12])
|
||||||
# Metadata fill sets only the invariant GLOBAL count; the LOCAL
|
# The LOCAL count reuses the proposer's persistent buffer; the GLOBAL
|
||||||
# num_token_non_padded is derived later when the draft forward localizes.
|
# 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.item(), 6)
|
||||||
self.assertEqual(forward_batch.global_num_token_non_padded.dtype, torch.int32)
|
self.assertEqual(forward_batch.global_num_token_non_padded.dtype, torch.int32)
|
||||||
self.assertEqual(forward_batch.global_num_token_non_padded_cpu, 6)
|
self.assertEqual(forward_batch.global_num_token_non_padded_cpu, 6)
|
||||||
|
|||||||
Reference in New Issue
Block a user