[AMD][DSV4] perf: use full 1024-thread block for indexer top-k on ROCm (#36004)
This commit is contained in:
@@ -28,7 +28,16 @@ limitations under the License.
|
|||||||
namespace {
|
namespace {
|
||||||
|
|
||||||
constexpr uint32_t kMaxTopK = 1024;
|
constexpr uint32_t kMaxTopK = 1024;
|
||||||
|
#ifdef USE_ROCM
|
||||||
|
// CDNA3/CDNA4: this kernel is one block per row and is latency-bound on its
|
||||||
|
// O(c4_len) histogram and emit passes. A full 1024-thread block (16 wavefronts
|
||||||
|
// of 64 lanes) instead of 512 doubles the per-block scan parallelism, which is
|
||||||
|
// ~1.6x faster at 128k context (c4_len = 32768) and never slower at short
|
||||||
|
// context. The selected index set is unchanged. CUDA keeps 512.
|
||||||
|
constexpr uint32_t kBlockSize = 1024;
|
||||||
|
#else
|
||||||
constexpr uint32_t kBlockSize = 512;
|
constexpr uint32_t kBlockSize = 512;
|
||||||
|
#endif
|
||||||
|
|
||||||
#ifdef SGL_TOPK_DYNAMIC_SMEM_BYTES
|
#ifdef SGL_TOPK_DYNAMIC_SMEM_BYTES
|
||||||
constexpr size_t kSMEM = static_cast<size_t>(SGL_TOPK_DYNAMIC_SMEM_BYTES);
|
constexpr size_t kSMEM = static_cast<size_t>(SGL_TOPK_DYNAMIC_SMEM_BYTES);
|
||||||
|
|||||||
@@ -249,5 +249,48 @@ def test_topk_transform_ragged_kernel(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(
|
||||||
|
torch.version.hip is None,
|
||||||
|
reason="deepseek_v4_topk_transform_512 is only built on ROCm",
|
||||||
|
)
|
||||||
|
@pytest.mark.parametrize("bs", [1, 48])
|
||||||
|
@pytest.mark.parametrize("c4_len", [2048, 8192, 32768])
|
||||||
|
@torch.inference_mode()
|
||||||
|
def test_deepseek_v4_topk_transform(bs: int, c4_len: int) -> None:
|
||||||
|
# c4_len 32768 is the 128k-context decode shape, i.e. the longest scan the
|
||||||
|
# kernel runs and the one most sensitive to the block size it launches with.
|
||||||
|
from sgl_kernel import deepseek_v4_topk_transform_512
|
||||||
|
|
||||||
|
torch.manual_seed(42)
|
||||||
|
topk, page_size = 1024, 64
|
||||||
|
|
||||||
|
scores = torch.randn(bs, c4_len, dtype=torch.float32, device="cuda")
|
||||||
|
seq_lens = torch.full((bs,), c4_len, dtype=torch.int32, device="cuda")
|
||||||
|
# Identity page table, so emitted paged slots equal raw token positions and
|
||||||
|
# can be compared against torch.topk indices directly.
|
||||||
|
num_pages = (c4_len + page_size - 1) // page_size
|
||||||
|
page_table = (
|
||||||
|
torch.arange(num_pages, dtype=torch.int32, device="cuda")
|
||||||
|
.unsqueeze(0)
|
||||||
|
.expand(bs, -1)
|
||||||
|
.contiguous()
|
||||||
|
)
|
||||||
|
page_indices = torch.full((bs, topk), -1, dtype=torch.int32, device="cuda")
|
||||||
|
|
||||||
|
deepseek_v4_topk_transform_512(
|
||||||
|
scores, seq_lens, page_table, page_indices, page_size
|
||||||
|
)
|
||||||
|
|
||||||
|
indices_ref = torch.topk(scores, topk, dim=-1, sorted=False).indices
|
||||||
|
assert_equal(
|
||||||
|
scores,
|
||||||
|
torch.sort(indices_ref, dim=-1).values,
|
||||||
|
torch.sort(page_indices, dim=-1).values,
|
||||||
|
bs,
|
||||||
|
topk,
|
||||||
|
c4_len,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__]))
|
sys.exit(pytest.main([__file__]))
|
||||||
|
|||||||
Reference in New Issue
Block a user