compile _resolve_spec_extras gather kernels (#26129)

This commit is contained in:
Liangsheng Yin
2026-05-23 02:34:41 -07:00
committed by GitHub
parent 83a18e687d
commit 1e59ed7443
+37 -6
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
import os
from typing import TYPE_CHECKING, Union
from typing import TYPE_CHECKING, Optional, Union
import torch
@@ -33,6 +33,25 @@ def _assert_nonneg_and_invalidate(
buf[indices] = -1
@torch.compile(dynamic=True)
def _gather_spec_extras(
indices: torch.Tensor,
topk_p_buf: torch.Tensor,
topk_index_buf: torch.Tensor,
output_tokens_buf: torch.Tensor,
hidden_states_buf: Optional[torch.Tensor],
):
"""Compiled gather of spec extras. `hidden_states_buf` is None when the
build does not capture hidden states."""
topk_p = topk_p_buf[indices]
topk_index = topk_index_buf[indices]
bonus_tokens = output_tokens_buf[indices]
hidden_states = (
hidden_states_buf[indices] if hidden_states_buf is not None else None
)
return topk_p, topk_index, bonus_tokens, hidden_states
def _resolve_future_token_ids_native(input_ids, future_token_ids_map):
input_ids[:] = torch.where(
input_ids < 0,
@@ -135,15 +154,27 @@ class FutureMap:
# FIXME: indices = batch.req_pool_indices, pinned 2 iters via
# record_batch_in_overlap; record_stream here is redundant.
indices.record_stream(torch.get_device_module(self.device).current_stream())
draft_input.topk_p = self.topk_p_buf[indices]
draft_input.topk_index = self.topk_index_buf[indices]
draft_input.bonus_tokens = self.output_tokens_buf[indices]
hidden_states_buf = (
self.hidden_states_buf if spec_need_hidden_states() else None
)
(
draft_input.topk_p,
draft_input.topk_index,
draft_input.bonus_tokens,
hidden_states,
) = _gather_spec_extras(
indices,
self.topk_p_buf,
self.topk_index_buf,
self.output_tokens_buf,
hidden_states_buf,
)
if hidden_states is not None:
draft_input.hidden_states = hidden_states
if _DEBUG_ASSERT:
_assert_nonneg_and_invalidate(
draft_input.bonus_tokens, self.output_tokens_buf, indices
)
if spec_need_hidden_states():
draft_input.hidden_states = self.hidden_states_buf[indices]
def set_input_ids_sentinel(
self, batch: ScheduleBatch, future_indices: torch.Tensor