diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index b2255c3d7..132719818 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -34,10 +34,8 @@ from sglang.srt.speculative.dflash_utils import ( from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_utils import assign_req_to_token_pool_func from sglang.srt.speculative.triton_ops.cache_locs import assign_extend_cache_locs_func -from sglang.srt.speculative.triton_ops.dflash_accept_bonus import ( +from sglang.srt.speculative.triton_ops.dflash import ( _compute_dflash_accept_bonus_triton_unchecked, -) -from sglang.srt.speculative.triton_ops.dflash_prepare_block import ( _prepare_dflash_draft_block_unchecked, ) from sglang.srt.utils import get_available_gpu_memory, is_cuda, is_hip, is_npu diff --git a/python/sglang/srt/speculative/triton_ops/dflash_accept_bonus.py b/python/sglang/srt/speculative/triton_ops/dflash.py similarity index 58% rename from python/sglang/srt/speculative/triton_ops/dflash_accept_bonus.py rename to python/sglang/srt/speculative/triton_ops/dflash.py index 7633fe3be..d435ce282 100644 --- a/python/sglang/srt/speculative/triton_ops/dflash_accept_bonus.py +++ b/python/sglang/srt/speculative/triton_ops/dflash.py @@ -138,3 +138,109 @@ def _compute_dflash_accept_bonus_triton_unchecked( BLOCK_SIZE=block, num_warps=num_warps, ) + + +@triton.jit +def _prepare_dflash_draft_block_contig_kernel( + bonus_tokens_ptr, + prefix_lens_ptr, + req_pool_indices_ptr, + req_to_token_ptr, + block_ids_out_ptr, + positions_out_ptr, + cache_loc_out_ptr, + bonus_tokens_stride, + prefix_lens_stride, + req_pool_indices_stride, + req_to_token_row_stride, + block_ids_row_stride, + positions_row_stride, + cache_loc_row_stride, + req_to_token_width, + block_size, + mask_token_id, + BLOCK_SIZE: tl.constexpr, +): + row = tl.program_id(0) + cols = tl.arange(0, BLOCK_SIZE) + row_mask = cols < block_size + + prefix_len = tl.load(prefix_lens_ptr + row * prefix_lens_stride) + req_idx = tl.load(req_pool_indices_ptr + row * req_pool_indices_stride) + bonus_token = tl.load(bonus_tokens_ptr + row * bonus_tokens_stride) + + logical_pos = prefix_len.to(tl.int64) + cols + valid = row_mask & (logical_pos < req_to_token_width) + req_row_ptr = req_to_token_ptr + req_idx * req_to_token_row_stride + slot_ids = tl.load(req_row_ptr + logical_pos, mask=valid, other=0) + + block_ids = tl.full((BLOCK_SIZE,), mask_token_id, tl.int64) + block_ids = tl.where(cols == 0, bonus_token.to(tl.int64), block_ids) + tl.store( + block_ids_out_ptr + row * block_ids_row_stride + cols, block_ids, mask=row_mask + ) + tl.store( + positions_out_ptr + row * positions_row_stride + cols, + logical_pos, + mask=row_mask, + ) + tl.store( + cache_loc_out_ptr + row * cache_loc_row_stride + cols, + slot_ids.to(tl.int64), + mask=row_mask, + ) + + +def _prepare_dflash_draft_block_unchecked( + bonus_tokens: torch.Tensor, + prefix_lens: torch.Tensor, + req_pool_indices: torch.Tensor, + req_to_token: torch.Tensor, + block_ids_out: torch.Tensor, + positions_out: torch.Tensor, + cache_loc_out: torch.Tensor, + mask_token_id: int, +) -> None: + batch_size = int(bonus_tokens.numel()) + if batch_size == 0: + return + + if req_to_token.ndim != 2 or req_to_token.stride(1) != 1: + raise ValueError("DFLASH Triton prepare_block requires row-major req_to_token.") + if not _is_row_major_contiguous_2d(block_ids_out): + raise ValueError( + "DFLASH Triton prepare_block requires contiguous block_ids_out." + ) + if not _is_row_major_contiguous_2d(positions_out): + raise ValueError( + "DFLASH Triton prepare_block requires contiguous positions_out." + ) + if not _is_row_major_contiguous_2d(cache_loc_out): + raise ValueError( + "DFLASH Triton prepare_block requires contiguous cache_loc_out." + ) + + block_size = int(block_ids_out.shape[1]) + block = triton.next_power_of_2(block_size) + num_warps = _pick_num_warps(block) + _prepare_dflash_draft_block_contig_kernel[(batch_size,)]( + bonus_tokens, + prefix_lens, + req_pool_indices, + req_to_token, + block_ids_out, + positions_out, + cache_loc_out, + bonus_tokens.stride(0), + prefix_lens.stride(0), + req_pool_indices.stride(0), + req_to_token.stride(0), + block_ids_out.stride(0), + positions_out.stride(0), + cache_loc_out.stride(0), + int(req_to_token.shape[1]), + block_size, + int(mask_token_id), + BLOCK_SIZE=block, + num_warps=num_warps, + ) diff --git a/python/sglang/srt/speculative/triton_ops/dflash_prepare_block.py b/python/sglang/srt/speculative/triton_ops/dflash_prepare_block.py deleted file mode 100644 index d14b0dd76..000000000 --- a/python/sglang/srt/speculative/triton_ops/dflash_prepare_block.py +++ /dev/null @@ -1,123 +0,0 @@ -import torch -import triton -import triton.language as tl - - -@triton.jit -def _prepare_dflash_draft_block_contig_kernel( - bonus_tokens_ptr, - prefix_lens_ptr, - req_pool_indices_ptr, - req_to_token_ptr, - block_ids_out_ptr, - positions_out_ptr, - cache_loc_out_ptr, - bonus_tokens_stride, - prefix_lens_stride, - req_pool_indices_stride, - req_to_token_row_stride, - block_ids_row_stride, - positions_row_stride, - cache_loc_row_stride, - req_to_token_width, - block_size, - mask_token_id, - BLOCK_SIZE: tl.constexpr, -): - row = tl.program_id(0) - cols = tl.arange(0, BLOCK_SIZE) - row_mask = cols < block_size - - prefix_len = tl.load(prefix_lens_ptr + row * prefix_lens_stride) - req_idx = tl.load(req_pool_indices_ptr + row * req_pool_indices_stride) - bonus_token = tl.load(bonus_tokens_ptr + row * bonus_tokens_stride) - - logical_pos = prefix_len.to(tl.int64) + cols - valid = row_mask & (logical_pos < req_to_token_width) - req_row_ptr = req_to_token_ptr + req_idx * req_to_token_row_stride - slot_ids = tl.load(req_row_ptr + logical_pos, mask=valid, other=0) - - block_ids = tl.full((BLOCK_SIZE,), mask_token_id, tl.int64) - block_ids = tl.where(cols == 0, bonus_token.to(tl.int64), block_ids) - tl.store( - block_ids_out_ptr + row * block_ids_row_stride + cols, block_ids, mask=row_mask - ) - tl.store( - positions_out_ptr + row * positions_row_stride + cols, - logical_pos, - mask=row_mask, - ) - tl.store( - cache_loc_out_ptr + row * cache_loc_row_stride + cols, - slot_ids.to(tl.int64), - mask=row_mask, - ) - - -def _pick_num_warps(block_size: int) -> int: - if block_size <= 16: - return 1 - if block_size <= 32: - return 2 - if block_size <= 64: - return 4 - return 8 - - -def _is_row_major_contiguous_2d(x: torch.Tensor) -> bool: - return x.ndim == 2 and x.is_contiguous() - - -def _prepare_dflash_draft_block_unchecked( - bonus_tokens: torch.Tensor, - prefix_lens: torch.Tensor, - req_pool_indices: torch.Tensor, - req_to_token: torch.Tensor, - block_ids_out: torch.Tensor, - positions_out: torch.Tensor, - cache_loc_out: torch.Tensor, - mask_token_id: int, -) -> None: - batch_size = int(bonus_tokens.numel()) - if batch_size == 0: - return - - if req_to_token.ndim != 2 or req_to_token.stride(1) != 1: - raise ValueError("DFLASH Triton prepare_block requires row-major req_to_token.") - if not _is_row_major_contiguous_2d(block_ids_out): - raise ValueError( - "DFLASH Triton prepare_block requires contiguous block_ids_out." - ) - if not _is_row_major_contiguous_2d(positions_out): - raise ValueError( - "DFLASH Triton prepare_block requires contiguous positions_out." - ) - if not _is_row_major_contiguous_2d(cache_loc_out): - raise ValueError( - "DFLASH Triton prepare_block requires contiguous cache_loc_out." - ) - - block_size = int(block_ids_out.shape[1]) - block = triton.next_power_of_2(block_size) - num_warps = _pick_num_warps(block) - _prepare_dflash_draft_block_contig_kernel[(batch_size,)]( - bonus_tokens, - prefix_lens, - req_pool_indices, - req_to_token, - block_ids_out, - positions_out, - cache_loc_out, - bonus_tokens.stride(0), - prefix_lens.stride(0), - req_pool_indices.stride(0), - req_to_token.stride(0), - block_ids_out.stride(0), - positions_out.stride(0), - cache_loc_out.stride(0), - int(req_to_token.shape[1]), - block_size, - int(mask_token_id), - BLOCK_SIZE=block, - num_warps=num_warps, - )