class CandidateSampler:
"""The shared DFlash2/LiLiCorr candidate walk and realized proposal cache."""
def __init__(
self, max_num_reqs: int, num_steps: int, top_k: int, device: torch.device
):
self.num_steps = num_steps
self.top_k = top_k
self.scores = torch.empty(
max_num_reqs, num_steps, top_k, dtype=torch.float32, device=device
)
self.cached_candidate_ids = torch.zeros(
self.scores.shape, dtype=torch.int64, device=device
)
def sample(
self,
candidate_ids: torch.Tensor,
scores: torch.Tensor,
num_reqs: int,
sample_pos: torch.Tensor,
idx_mapping: torch.Tensor,
temperature: torch.Tensor,
seeds: torch.Tensor,
draft_tokens: torch.Tensor,
draft_logits: torch.Tensor | None,
use_fp64: bool,
) -> None:
candidate_ids = candidate_ids.contiguous()
sample_pos = sample_pos.contiguous()
idx_mapping = idx_mapping.contiguous()
block_k = triton.next_power_of_2(self.top_k)
_selector_walk_kernel[(num_reqs,)](
scores.contiguous(),
candidate_ids.contiguous(),
sample_pos,
idx_mapping,
temperature,
seeds,
draft_tokens,
self.scores,
num_steps=self.num_steps,
top_k=self.top_k,
BLOCK_K=block_k,
SAMPLE_PROBABILISTIC=draft_logits is not None,
USE_FP64=use_fp64,
num_warps=1,
)
if draft_logits is not None:
self._cache_draft_logits(
candidate_ids, num_reqs * self.num_steps, idx_mapping, draft_logits
)
def _cache_draft_logits(
self,
candidate_ids: torch.Tensor,
num_sample: int,
idx_mapping: torch.Tensor,
draft_logits: torch.Tensor,
) -> None:
block_k = triton.next_power_of_2(self.top_k)
_cache_draft_logits_kernel[(num_sample,)](
draft_logits,
self.cached_candidate_ids,
candidate_ids,
self.scores,
idx_mapping,
draft_logits.stride(0),
draft_logits.stride(1),
num_steps=self.num_steps,
top_k=self.top_k,
BLOCK_K=block_k,
num_warps=1,
)