[Spec] Dissolve EagleDraftInputV2Mixin so spec-info dataclasses hold data only (#29220)

This commit is contained in:
Liangsheng Yin
2026-06-24 18:08:30 -07:00
committed by GitHub
parent 2b64fc7a2c
commit c7734e6871
15 changed files with 122 additions and 134 deletions
@@ -149,7 +149,7 @@ def _make_eagle_draft_extend_v2_input(case, batch, *, device: str):
def _set_draft_extend_v2_prefix_lens(batch, case, *, device: str):
# Production sets seq_lens = prefix + extend before init_forward_metadata
# (eagle_info_v2.py bumps seq_lens by num_draft_tokens). Match that here.
# (the draft-extend path bumps seq_lens by num_draft_tokens). Match that here.
seq_lens = tuple(p + e for p, e in zip(case.prefix_lens, case.input_lens))
batch.seq_lens = torch.tensor(seq_lens, dtype=torch.int32, device=device)
batch.seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int32, device="cpu")
@@ -821,7 +821,7 @@ def _set_draft_extend_v2_prefix_lens(
device: str,
) -> None:
# Production sets seq_lens = prefix + extend before init_forward_metadata
# (eagle_info_v2.py bumps seq_lens by num_draft_tokens). Match that here.
# (the draft-extend path bumps seq_lens by num_draft_tokens). Match that here.
seq_lens = tuple(p + e for p, e in zip(case.prefix_lens, case.input_lens))
batch.seq_lens = torch.tensor(seq_lens, dtype=torch.int32, device=device)
batch.seq_lens_cpu = torch.tensor(seq_lens, dtype=torch.int32, device="cpu")