[spec decoding] fix mrope_positions in draft extend (#28464)
This commit is contained in:
@@ -781,6 +781,13 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
and getattr(ret.spec_info, "positions", None) is not None
|
and getattr(ret.spec_info, "positions", None) is not None
|
||||||
):
|
):
|
||||||
ret.compute_spec_mrope_positions(model_runner, batch)
|
ret.compute_spec_mrope_positions(model_runner, batch)
|
||||||
|
elif ret.forward_mode.is_draft_extend_v2():
|
||||||
|
# Draft-extend tokens are uniform text continuation; reuse the
|
||||||
|
# spec mrope path with the input-consistent `ret.positions` rather
|
||||||
|
# than the per-request rebuild (which mis-sizes mm requests).
|
||||||
|
ret.compute_spec_mrope_positions(
|
||||||
|
model_runner, batch, seq_positions=ret.positions
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
ret._compute_mrope_positions(model_runner, batch)
|
ret._compute_mrope_positions(model_runner, batch)
|
||||||
|
|
||||||
@@ -923,15 +930,18 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def compute_spec_mrope_positions(
|
def compute_spec_mrope_positions(
|
||||||
self, model_runner: ModelRunner, batch: ScheduleBatch
|
self, model_runner: ModelRunner, batch: ScheduleBatch, seq_positions=None
|
||||||
):
|
):
|
||||||
# TODO support batched deltas
|
# TODO support batched deltas
|
||||||
batch_size = self.seq_lens.shape[0]
|
batch_size = self.seq_lens.shape[0]
|
||||||
device = model_runner.device
|
device = model_runner.device
|
||||||
mm_inputs = batch.multimodal_inputs
|
mm_inputs = batch.multimodal_inputs
|
||||||
|
|
||||||
# target_verify or draft_decode
|
# target_verify / draft_decode read spec_info.positions; draft_extend
|
||||||
seq_positions = batch.spec_info.positions.view(batch_size, -1)
|
# passes its own positions (uniform num_draft_tokens per request).
|
||||||
|
if seq_positions is None:
|
||||||
|
seq_positions = batch.spec_info.positions
|
||||||
|
seq_positions = seq_positions.view(batch_size, -1)
|
||||||
# Split text-only and mixed batches here because SpecV2 text-only batches can avoid an extra D2H.
|
# Split text-only and mixed batches here because SpecV2 text-only batches can avoid an extra D2H.
|
||||||
if all(mm_input is None for mm_input in mm_inputs):
|
if all(mm_input is None for mm_input in mm_inputs):
|
||||||
mrope_delta_tensor = torch.zeros(
|
mrope_delta_tensor = torch.zeros(
|
||||||
|
|||||||
Reference in New Issue
Block a user