From 63296a2bb83aab94de1fa28469f306fe746cf991 Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Wed, 22 Jul 2026 16:29:21 +0800 Subject: [PATCH] feat(nvidia): add moe_align_block_size operator --- src/base/moe_align_block_size.h | 207 ++++++++ .../nvidia/ops/moe_align_block_size/kernel.cu | 87 ++++ .../ops/moe_align_block_size/kernel.cuh | 83 +++ .../nvidia/ops/moe_align_block_size/kernel.h | 37 ++ tests/test_moe_align_block_size.py | 475 ++++++++++++++++++ 5 files changed, 889 insertions(+) create mode 100644 src/base/moe_align_block_size.h create mode 100644 src/native/cuda/nvidia/ops/moe_align_block_size/kernel.cu create mode 100644 src/native/cuda/nvidia/ops/moe_align_block_size/kernel.cuh create mode 100644 src/native/cuda/nvidia/ops/moe_align_block_size/kernel.h create mode 100644 tests/test_moe_align_block_size.py diff --git a/src/base/moe_align_block_size.h b/src/base/moe_align_block_size.h new file mode 100644 index 000000000..6a23dae30 --- /dev/null +++ b/src/base/moe_align_block_size.h @@ -0,0 +1,207 @@ +#ifndef INFINI_OPS_BASE_MOE_ALIGN_BLOCK_SIZE_H_ +#define INFINI_OPS_BASE_MOE_ALIGN_BLOCK_SIZE_H_ + +#include +#include +#include +#include +#include + +#include "operator.h" + +namespace infini::ops { + +// Align routed token indices into expert-specific blocks following vLLM's +// low-level `moe_align_block_size` operator. +class MoeAlignBlockSize : public Operator { + public: + MoeAlignBlockSize(const Tensor topk_ids, const int64_t num_experts, + const int64_t block_size, Tensor sorted_token_ids, + Tensor experts_ids, Tensor num_tokens_post_pad) + : topk_ids_metadata_{topk_ids}, + expert_map_metadata_{std::nullopt}, + sorted_token_ids_metadata_{sorted_token_ids}, + experts_ids_metadata_{experts_ids}, + num_tokens_post_pad_metadata_{num_tokens_post_pad}, + numel_{topk_ids.numel()}, + num_experts_{num_experts}, + block_size_{block_size}, + sorted_token_ids_size_{sorted_token_ids.numel()}, + experts_ids_size_{experts_ids.numel()} { + Validate(topk_ids, std::nullopt, sorted_token_ids, experts_ids, + num_tokens_post_pad); + } + + MoeAlignBlockSize(const Tensor topk_ids, const Tensor expert_map, + const int64_t num_experts, const int64_t block_size, + Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad) + : topk_ids_metadata_{topk_ids}, + expert_map_metadata_{expert_map}, + sorted_token_ids_metadata_{sorted_token_ids}, + experts_ids_metadata_{experts_ids}, + num_tokens_post_pad_metadata_{num_tokens_post_pad}, + numel_{topk_ids.numel()}, + num_experts_{num_experts}, + block_size_{block_size}, + sorted_token_ids_size_{sorted_token_ids.numel()}, + experts_ids_size_{experts_ids.numel()} { + Validate(topk_ids, std::optional{expert_map}, sorted_token_ids, + experts_ids, num_tokens_post_pad); + } + + void operator()(const Tensor topk_ids, const int64_t num_experts, + const int64_t block_size, Tensor sorted_token_ids, + Tensor experts_ids, Tensor num_tokens_post_pad) const { + ValidateInvocation(topk_ids, std::nullopt, num_experts, block_size, + sorted_token_ids, experts_ids, num_tokens_post_pad); + Run(topk_ids, std::nullopt, num_experts, block_size, sorted_token_ids, + experts_ids, num_tokens_post_pad); + } + + void operator()(const Tensor topk_ids, const Tensor expert_map, + const int64_t num_experts, const int64_t block_size, + Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad) const { + ValidateInvocation(topk_ids, std::optional{expert_map}, num_experts, + block_size, sorted_token_ids, experts_ids, + num_tokens_post_pad); + Run(topk_ids, std::optional{expert_map}, num_experts, block_size, + sorted_token_ids, experts_ids, num_tokens_post_pad); + } + + protected: + void ValidateInvocation(const Tensor topk_ids, + const std::optional maybe_expert_map, + const int64_t num_experts, const int64_t block_size, + Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad) const { + assert(num_experts == num_experts_ && block_size == block_size_ && + "`MoeAlignBlockSize` attributes changed after descriptor creation"); + + assert(CallMetadataMatches(topk_ids, maybe_expert_map, sorted_token_ids, + experts_ids, num_tokens_post_pad) && + "`MoeAlignBlockSize` tensor metadata differs from its descriptor"); + } + + bool CallMetadataMatches(const Tensor topk_ids, + const std::optional maybe_expert_map, + const Tensor sorted_token_ids, + const Tensor experts_ids, + const Tensor num_tokens_post_pad) const { + const std::equal_to same_metadata; + const auto same_expert_map_metadata = + expert_map_metadata_.has_value() == maybe_expert_map.has_value() && + (!expert_map_metadata_ || + same_metadata(*expert_map_metadata_, *maybe_expert_map)); + + return same_metadata(topk_ids_metadata_, topk_ids) && + same_expert_map_metadata && + same_metadata(sorted_token_ids_metadata_, sorted_token_ids) && + same_metadata(experts_ids_metadata_, experts_ids) && + same_metadata(num_tokens_post_pad_metadata_, num_tokens_post_pad); + } + + void Validate(const Tensor topk_ids, + const std::optional maybe_expert_map, + Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad) const { + assert(topk_ids.ndim() == 2 && + "`MoeAlignBlockSize` requires 2D `topk_ids`"); + assert(topk_ids.dtype() == DataType::kInt32 && + "`MoeAlignBlockSize` currently requires int32 `topk_ids`"); + assert(topk_ids.IsContiguous() && + "`MoeAlignBlockSize` requires contiguous `topk_ids`"); + assert(num_experts_ > 0 && num_experts_ < 1024 && + "`MoeAlignBlockSize` requires `num_experts` in [1, 1023]"); + assert(block_size_ > 0 && + "`MoeAlignBlockSize` requires a positive `block_size`"); + assert(block_size_ <= std::numeric_limits::max() && + numel_ <= + static_cast(std::numeric_limits::max()) && + "`MoeAlignBlockSize` requires int32-addressable token indices"); + + const auto same_device_as_topk_ids = [&](const Tensor tensor) { + return tensor.device().type() == topk_ids.device().type() && + tensor.device().index() == topk_ids.device().index(); + }; + assert(same_device_as_topk_ids(sorted_token_ids) && + same_device_as_topk_ids(experts_ids) && + same_device_as_topk_ids(num_tokens_post_pad) && + "`MoeAlignBlockSize` requires all tensors on the same device"); + + assert(sorted_token_ids.ndim() == 1 && experts_ids.ndim() == 1 && + num_tokens_post_pad.ndim() == 1 && + num_tokens_post_pad.numel() == 1 && + "`MoeAlignBlockSize` requires 1D output tensors"); + assert(sorted_token_ids.dtype() == DataType::kInt32 && + experts_ids.dtype() == DataType::kInt32 && + num_tokens_post_pad.dtype() == DataType::kInt32 && + "`MoeAlignBlockSize` requires int32 output tensors"); + assert(sorted_token_ids.IsContiguous() && experts_ids.IsContiguous() && + num_tokens_post_pad.IsContiguous() && + "`MoeAlignBlockSize` requires contiguous output tensors"); + + const auto num_experts = static_cast(num_experts_); + const auto block_size = static_cast(block_size_); + auto required_sorted_size = numel_ + num_experts * (block_size - 1); + if (numel_ < num_experts) { + const auto small_input_size = numel_ * block_size; + required_sorted_size = small_input_size < required_sorted_size + ? small_input_size + : required_sorted_size; + } + const auto required_experts_size = + (required_sorted_size + block_size - 1) / block_size; + assert(required_sorted_size <= + static_cast(std::numeric_limits::max()) && + "`MoeAlignBlockSize` requires int32-addressable padded indices"); + assert(sorted_token_ids_size_ >= required_sorted_size && + experts_ids_size_ >= required_experts_size && + "`MoeAlignBlockSize` output tensors are too small"); + + if (maybe_expert_map) { + assert(maybe_expert_map->ndim() == 1 && + maybe_expert_map->numel() == + static_cast(num_experts_) && + "`MoeAlignBlockSize` requires `expert_map` shape " + "[`num_experts`]"); + assert(maybe_expert_map->dtype() == DataType::kInt32 && + "`MoeAlignBlockSize` currently requires int32 `expert_map`"); + assert(maybe_expert_map->IsContiguous() && + "`MoeAlignBlockSize` requires contiguous `expert_map`"); + assert(same_device_as_topk_ids(*maybe_expert_map) && + "`MoeAlignBlockSize` requires `expert_map` on the input device"); + } + } + + virtual void Run(const Tensor topk_ids, + const std::optional maybe_expert_map, + const int64_t num_experts, const int64_t block_size, + Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad) const = 0; + + Tensor topk_ids_metadata_; + + std::optional expert_map_metadata_; + + Tensor sorted_token_ids_metadata_; + + Tensor experts_ids_metadata_; + + Tensor num_tokens_post_pad_metadata_; + + Tensor::Size numel_{0}; + + int64_t num_experts_{0}; + + int64_t block_size_{0}; + + Tensor::Size sorted_token_ids_size_{0}; + + Tensor::Size experts_ids_size_{0}; +}; + +} // namespace infini::ops + +#endif // INFINI_OPS_BASE_MOE_ALIGN_BLOCK_SIZE_H_ diff --git a/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.cu b/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.cu new file mode 100644 index 000000000..58b63e76e --- /dev/null +++ b/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.cu @@ -0,0 +1,87 @@ +#include + +#include +#include +#include + +#include "native/cuda/nvidia/ops/moe_align_block_size/kernel.cuh" +#include "native/cuda/nvidia/ops/moe_align_block_size/kernel.h" + +namespace infini::ops { +namespace { + +class DeviceGuard { + public: + explicit DeviceGuard(int device_index) { + auto status = cudaGetDevice(&previous_device_); + assert(status == cudaSuccess && + "`MoeAlignBlockSize` failed to query the current CUDA device"); + + if (previous_device_ != device_index) { + status = cudaSetDevice(device_index); + assert(status == cudaSuccess && + "`MoeAlignBlockSize` failed to select the input CUDA device"); + restore_ = true; + } + } + + ~DeviceGuard() { + if (restore_) { + const auto status = cudaSetDevice(previous_device_); + assert(status == cudaSuccess && + "`MoeAlignBlockSize` failed to restore the CUDA device"); + } + } + + private: + int previous_device_{0}; + + bool restore_{false}; +}; + +} // namespace + +Operator::Operator( + const Tensor topk_ids, const int64_t num_experts, const int64_t block_size, + Tensor sorted_token_ids, Tensor experts_ids, Tensor num_tokens_post_pad) + : MoeAlignBlockSize{topk_ids, num_experts, block_size, + sorted_token_ids, experts_ids, num_tokens_post_pad}, + device_index_{topk_ids.device().index()} {} + +Operator::Operator( + const Tensor topk_ids, const Tensor expert_map, const int64_t num_experts, + const int64_t block_size, Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad) + : MoeAlignBlockSize{topk_ids, expert_map, num_experts, + block_size, sorted_token_ids, experts_ids, + num_tokens_post_pad}, + device_index_{topk_ids.device().index()} {} + +void Operator::Run( + const Tensor topk_ids, const std::optional maybe_expert_map, + const int64_t num_experts, const int64_t block_size, + Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad) const { + DeviceGuard device_guard{device_index_}; + + constexpr int32_t kThreads = 256; + const auto shared_memory_size = + static_cast(num_experts_) * sizeof(int32_t); + moe_align_block_size_detail::MoeAlignBlockSizeKernel<<< + 1, kThreads, shared_memory_size, static_cast(stream_)>>>( + reinterpret_cast(topk_ids.data()), + maybe_expert_map + ? reinterpret_cast(maybe_expert_map->data()) + : nullptr, + reinterpret_cast(sorted_token_ids.data()), + reinterpret_cast(experts_ids.data()), + reinterpret_cast(num_tokens_post_pad.data()), numel_, + static_cast(num_experts_), static_cast(block_size_), + sorted_token_ids_size_, experts_ids_size_); + + const auto status = cudaGetLastError(); + assert(status == cudaSuccess && + "`MoeAlignBlockSize` CUDA kernel launch failed"); +} + +} // namespace infini::ops diff --git a/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.cuh b/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.cuh new file mode 100644 index 000000000..467383221 --- /dev/null +++ b/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.cuh @@ -0,0 +1,83 @@ +#ifndef INFINI_OPS_NVIDIA_MOE_ALIGN_BLOCK_SIZE_KERNEL_CUH_ +#define INFINI_OPS_NVIDIA_MOE_ALIGN_BLOCK_SIZE_KERNEL_CUH_ + +#include +#include + +namespace infini::ops { +namespace moe_align_block_size_detail { + +__device__ __forceinline__ int32_t MapExpertId(int32_t expert_id, + const int32_t* expert_map, + int32_t num_experts) { + if (expert_id < 0 || expert_id >= num_experts) { + return -1; + } + if (expert_map != nullptr) { + expert_id = expert_map[expert_id]; + } + + return expert_id >= 0 && expert_id < num_experts ? expert_id : -1; +} + +__global__ void MoeAlignBlockSizeKernel( + const int32_t* __restrict__ topk_ids, + const int32_t* __restrict__ expert_map, + int32_t* __restrict__ sorted_token_ids, int32_t* __restrict__ experts_ids, + int32_t* __restrict__ num_tokens_post_pad, size_t numel, + int32_t num_experts, int32_t block_size, size_t sorted_token_ids_size, + size_t experts_ids_size) { + extern __shared__ int32_t expert_offsets[]; + + for (size_t i = threadIdx.x; i < sorted_token_ids_size; i += blockDim.x) { + sorted_token_ids[i] = static_cast(numel); + } + for (size_t i = threadIdx.x; i < experts_ids_size; i += blockDim.x) { + experts_ids[i] = -1; + } + for (int32_t i = threadIdx.x; i < num_experts; i += blockDim.x) { + expert_offsets[i] = 0; + } + if (threadIdx.x == 0) { + *num_tokens_post_pad = 0; + } + __syncthreads(); + + for (size_t i = threadIdx.x; i < numel; i += blockDim.x) { + const auto expert_id = MapExpertId(topk_ids[i], expert_map, num_experts); + if (expert_id >= 0) { + atomicAdd(&expert_offsets[expert_id], 1); + } + } + __syncthreads(); + + if (threadIdx.x == 0) { + int32_t token_offset = 0; + for (int32_t expert_id = 0; expert_id < num_experts; ++expert_id) { + const int32_t count = expert_offsets[expert_id]; + expert_offsets[expert_id] = token_offset; + const int32_t padded_count = + (count + block_size - 1) / block_size * block_size; + for (int32_t i = 0; i < padded_count; i += block_size) { + experts_ids[(token_offset + i) / block_size] = expert_id; + } + token_offset += padded_count; + } + *num_tokens_post_pad = token_offset; + } + __syncthreads(); + + if (threadIdx.x == 0) { + for (size_t i = 0; i < numel; ++i) { + const auto expert_id = MapExpertId(topk_ids[i], expert_map, num_experts); + if (expert_id >= 0) { + sorted_token_ids[expert_offsets[expert_id]++] = static_cast(i); + } + } + } +} + +} // namespace moe_align_block_size_detail +} // namespace infini::ops + +#endif // INFINI_OPS_NVIDIA_MOE_ALIGN_BLOCK_SIZE_KERNEL_CUH_ diff --git a/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.h b/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.h new file mode 100644 index 000000000..56e93b406 --- /dev/null +++ b/src/native/cuda/nvidia/ops/moe_align_block_size/kernel.h @@ -0,0 +1,37 @@ +#ifndef INFINI_OPS_NVIDIA_MOE_ALIGN_BLOCK_SIZE_KERNEL_H_ +#define INFINI_OPS_NVIDIA_MOE_ALIGN_BLOCK_SIZE_KERNEL_H_ + +#include + +#include "base/moe_align_block_size.h" + +namespace infini::ops { + +template <> +class Operator + : public MoeAlignBlockSize { + public: + Operator(const Tensor topk_ids, const int64_t num_experts, + const int64_t block_size, Tensor sorted_token_ids, + Tensor experts_ids, Tensor num_tokens_post_pad); + + Operator(const Tensor topk_ids, const Tensor expert_map, + const int64_t num_experts, const int64_t block_size, + Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad); + + using MoeAlignBlockSize::operator(); + + protected: + void Run(const Tensor topk_ids, const std::optional maybe_expert_map, + const int64_t num_experts, const int64_t block_size, + Tensor sorted_token_ids, Tensor experts_ids, + Tensor num_tokens_post_pad) const override; + + private: + int device_index_{0}; +}; + +} // namespace infini::ops + +#endif // INFINI_OPS_NVIDIA_MOE_ALIGN_BLOCK_SIZE_KERNEL_H_ diff --git a/tests/test_moe_align_block_size.py b/tests/test_moe_align_block_size.py new file mode 100644 index 000000000..955af40e2 --- /dev/null +++ b/tests/test_moe_align_block_size.py @@ -0,0 +1,475 @@ +import subprocess +import sys +import textwrap + +import infini.ops +import pytest +import torch + +from tests.utils import get_stream + + +if not hasattr(infini.ops, "MoeAlignBlockSize"): + pytest.skip( + "`MoeAlignBlockSize` is not available on this platform", + allow_module_level=True, + ) + + +@pytest.mark.parametrize( + "topk_ids_values, num_experts, block_size", + ( + (((0, 2), (0, 2), (2, 0)), 4, 2), + (((1, 1), (3, 0), (3, 3)), 5, 4), + (((4, 1), (0, 4), (2, 1), (4, 2)), 5, 8), + (((-1, 4), (2**31 - 1, 0), (2, -(2**31))), 4, 4), + ), +) +def test_moe_align_block_size( + topk_ids_values, + num_experts, + block_size, + device, + implementation_index, +): + if device != "cuda": + pytest.skip("`moe_align_block_size` requires the NVIDIA backend") + + topk_ids = torch.tensor(topk_ids_values, dtype=torch.int32, device=device) + outputs = _make_outputs(topk_ids, num_experts, block_size) + + _moe_align_block_size( + topk_ids, + None, + num_experts, + block_size, + *outputs, + implementation_index=implementation_index, + ) + + _assert_matches_reference(topk_ids, None, num_experts, block_size, *outputs) + + +def test_moe_align_block_size_default_expert_map(device, implementation_index): + if device != "cuda": + pytest.skip("`moe_align_block_size` requires the NVIDIA backend") + + topk_ids = torch.tensor(((0, 2), (1, 2), (0, 1)), dtype=torch.int32, device=device) + num_experts = 4 + block_size = 4 + outputs = _make_outputs(topk_ids, num_experts, block_size) + + infini.ops.moe_align_block_size( + topk_ids, + num_experts, + block_size, + *outputs, + stream=get_stream(topk_ids.device), + implementation_index=implementation_index, + ) + + _assert_matches_reference(topk_ids, None, num_experts, block_size, *outputs) + + +@pytest.mark.parametrize( + "topk_ids_values, expert_map_values, num_experts, block_size", + ( + ( + ((0, 1), (2, 3), (0, 2), (3, 1)), + (0, -1, 1, -1), + 4, + 4, + ), + ( + ((-1, 0), (1, 2), (3, 4), (2**31 - 1, 1)), + (0, 5, -1, 2), + 4, + 4, + ), + ), +) +def test_moe_align_block_size_expert_map( + topk_ids_values, + expert_map_values, + num_experts, + block_size, + device, + implementation_index, +): + if device != "cuda": + pytest.skip("`moe_align_block_size` requires the NVIDIA backend") + + topk_ids = torch.tensor(topk_ids_values, dtype=torch.int32, device=device) + expert_map = torch.tensor(expert_map_values, dtype=torch.int32, device=device) + outputs = _make_outputs(topk_ids, num_experts, block_size) + + _moe_align_block_size( + topk_ids, + expert_map, + num_experts, + block_size, + *outputs, + implementation_index=implementation_index, + ) + + _assert_matches_reference(topk_ids, expert_map, num_experts, block_size, *outputs) + + +def test_moe_align_block_size_large_sparse_case(device, implementation_index): + if device != "cuda": + pytest.skip("`moe_align_block_size` requires the NVIDIA backend") + + num_experts = 257 + block_size = 128 + routed_experts = torch.arange(390, dtype=torch.int32, device=device) % 3 + routed_experts[0] = 256 + topk_ids = routed_experts.reshape(130, 3) + outputs = _make_outputs(topk_ids, num_experts, block_size) + + _moe_align_block_size( + topk_ids, + None, + num_experts, + block_size, + *outputs, + implementation_index=implementation_index, + ) + + _assert_matches_reference(topk_ids, None, num_experts, block_size, *outputs) + + +def test_moe_align_block_size_is_deterministic(device, implementation_index): + if device != "cuda": + pytest.skip("`moe_align_block_size` requires the NVIDIA backend") + + num_experts = 7 + block_size = 32 + topk_ids = (torch.arange(2048, dtype=torch.int32, device=device) % 7).reshape( + 512, 4 + ) + results = [] + + for _ in range(5): + outputs = _make_outputs(topk_ids, num_experts, block_size) + _moe_align_block_size( + topk_ids, + None, + num_experts, + block_size, + *outputs, + implementation_index=implementation_index, + ) + results.append(tuple(output.clone() for output in outputs)) + + torch.cuda.synchronize(topk_ids.device) + expected = results[0] + + for result in results[1:]: + assert all( + torch.equal(actual, reference) + for actual, reference in zip(result, expected) + ) + + +def test_moe_align_block_size_descriptor_reuses_matching_metadata(device): + if device != "cuda": + pytest.skip("`moe_align_block_size` requires the NVIDIA backend") + + num_experts = 4 + block_size = 2 + topk_ids = torch.tensor(((0, 1), (2, 3)), dtype=torch.int32, device=device) + expert_map = torch.arange(num_experts, dtype=torch.int32, device=device) + outputs = _make_outputs(topk_ids, num_experts, block_size) + operator = infini.ops.MoeAlignBlockSize( + topk_ids, + expert_map, + num_experts, + block_size, + *outputs, + ) + reused_topk_ids = topk_ids.clone() + reused_expert_map = expert_map.clone() + reused_outputs = tuple(torch.empty_like(output) for output in outputs) + + operator( + reused_topk_ids, + reused_expert_map, + num_experts, + block_size, + *reused_outputs, + ) + + _assert_matches_reference( + reused_topk_ids, + reused_expert_map, + num_experts, + block_size, + *reused_outputs, + ) + + +@pytest.mark.parametrize( + "metadata_change", + ( + "topk_shape", + "topk_strides", + "topk_dtype", + "sorted_token_ids_size", + "expert_map_shape", + "expert_map_presence", + ), +) +def test_moe_align_block_size_descriptor_rejects_changed_metadata( + metadata_change, device, implementation_index +): + if device != "cuda": + pytest.skip("`moe_align_block_size` requires the NVIDIA backend") + + assertion_probe = subprocess.run( + [ + sys.executable, + "-c", + _DESCRIPTOR_REUSE_SCRIPT, + "assertions_enabled_probe", + ], + capture_output=True, + text=True, + ) + + if assertion_probe.returncode == 0: + pytest.skip("descriptor validation requires an assertions-enabled build") + + result = subprocess.run( + [sys.executable, "-c", _DESCRIPTOR_REUSE_SCRIPT, metadata_change], + capture_output=True, + text=True, + ) + + assert result.returncode != 0 + assert "tensor metadata differs from its descriptor" in result.stderr + + +def test_moe_align_block_size_non_default_stream(device, implementation_index): + if device != "cuda": + pytest.skip("non-default CUDA streams require the NVIDIA backend") + + topk_ids = torch.tensor( + ((0, 1), (2, 3), (0, 3), (2, 1)), dtype=torch.int32, device=device + ) + num_experts = 4 + block_size = 4 + outputs = _make_outputs(topk_ids, num_experts, block_size) + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + + infini.ops.moe_align_block_size( + topk_ids, + num_experts, + block_size, + *outputs, + stream=stream.cuda_stream, + implementation_index=implementation_index, + ) + + stream.synchronize() + _assert_matches_reference(topk_ids, None, num_experts, block_size, *outputs) + + +def test_moe_align_block_size_device_guard(): + if not infini.ops.MoeAlignBlockSize.active_implementation_indices("nvidia"): + pytest.skip("device guard test requires the NVIDIA implementation") + + if torch.cuda.device_count() < 2: + pytest.skip("device guard test requires two NVIDIA GPUs") + + original_device = torch.cuda.current_device() + + try: + torch.cuda.set_device(0) + target_device = torch.device("cuda:1") + topk_ids = torch.tensor( + ((0, 1), (2, 3), (0, 3), (2, 1)), + dtype=torch.int32, + device=target_device, + ) + num_experts = 4 + block_size = 4 + outputs = _make_outputs(topk_ids, num_experts, block_size) + stream = torch.cuda.Stream(device=target_device) + stream.wait_stream(torch.cuda.current_stream(target_device)) + + infini.ops.moe_align_block_size( + topk_ids, + num_experts, + block_size, + *outputs, + stream=stream.cuda_stream, + ) + + assert torch.cuda.current_device() == 0 + stream.synchronize() + _assert_matches_reference(topk_ids, None, num_experts, block_size, *outputs) + finally: + torch.cuda.set_device(original_device) + + +def _make_outputs(topk_ids, num_experts, block_size): + numel = topk_ids.numel() + max_num_tokens_padded = numel + num_experts * (block_size - 1) + + if numel < num_experts: + max_num_tokens_padded = min(numel * block_size, max_num_tokens_padded) + + max_num_blocks = (max_num_tokens_padded + block_size - 1) // block_size + sorted_token_ids = torch.full( + (max_num_tokens_padded,), -2, dtype=torch.int32, device=topk_ids.device + ) + expert_ids = torch.full( + (max_num_blocks,), -2, dtype=torch.int32, device=topk_ids.device + ) + num_tokens_post_pad = torch.full( + (1,), -2, dtype=torch.int32, device=topk_ids.device + ) + + return sorted_token_ids, expert_ids, num_tokens_post_pad + + +def _moe_align_block_size( + topk_ids, + expert_map, + num_experts, + block_size, + sorted_token_ids, + expert_ids, + num_tokens_post_pad, + *, + implementation_index, +): + args = (topk_ids,) + if expert_map is not None: + args += (expert_map,) + args += ( + num_experts, + block_size, + sorted_token_ids, + expert_ids, + num_tokens_post_pad, + ) + infini.ops.moe_align_block_size( + *args, + stream=get_stream(topk_ids.device), + implementation_index=implementation_index, + ) + + +def _assert_matches_reference( + topk_ids, + expert_map, + num_experts, + block_size, + sorted_token_ids, + expert_ids, + num_tokens_post_pad, +): + flattened_experts = topk_ids.flatten().to(torch.int64) + + if expert_map is not None: + mapped_experts = torch.full_like(flattened_experts, -1) + valid = (flattened_experts >= 0) & (flattened_experts < num_experts) + mapped_experts[valid] = expert_map[flattened_experts[valid]].to(torch.int64) + else: + mapped_experts = flattened_experts + + expected_sorted_token_ids = [] + expected_blocks = [] + + for expert_id in range(num_experts): + token_ids = torch.nonzero(mapped_experts == expert_id).flatten().tolist() + + if not token_ids: + continue + + num_blocks = (len(token_ids) + block_size - 1) // block_size + expected_sorted_token_ids.extend(token_ids) + expected_sorted_token_ids.extend( + [topk_ids.numel()] * (num_blocks * block_size - len(token_ids)) + ) + expected_blocks.extend([expert_id] * num_blocks) + + expected_num_tokens = len(expected_blocks) * block_size + assert num_tokens_post_pad.item() == expected_num_tokens + assert expert_ids[: len(expected_blocks)].tolist() == expected_blocks + assert torch.all(expert_ids[len(expected_blocks) :] == -1) + assert sorted_token_ids[:expected_num_tokens].tolist() == expected_sorted_token_ids + assert torch.all(sorted_token_ids[expected_num_tokens:] == topk_ids.numel()) + + +_DESCRIPTOR_REUSE_SCRIPT = textwrap.dedent( + r""" + import sys + + import infini.ops + import torch + + + metadata_change = sys.argv[1] + num_experts = 4 + block_size = 2 + topk_ids = torch.tensor(((0, 1), (2, 3)), dtype=torch.int32, device="cuda") + expert_map = torch.arange(num_experts, dtype=torch.int32, device="cuda") + sorted_token_ids = torch.empty(8, dtype=torch.int32, device="cuda") + expert_ids = torch.empty(4, dtype=torch.int32, device="cuda") + num_tokens_post_pad = torch.empty(1, dtype=torch.int32, device="cuda") + operator = infini.ops.MoeAlignBlockSize( + topk_ids, + expert_map, + num_experts, + block_size, + sorted_token_ids, + expert_ids, + num_tokens_post_pad, + ) + + if metadata_change == "assertions_enabled_probe": + infini.ops.MoeAlignBlockSize( + topk_ids.reshape(4), + expert_map, + num_experts, + block_size, + sorted_token_ids, + expert_ids, + num_tokens_post_pad, + ) + elif metadata_change == "topk_shape": + topk_ids = topk_ids.reshape(1, 4) + elif metadata_change == "topk_strides": + topk_ids = topk_ids.T + elif metadata_change == "topk_dtype": + topk_ids = topk_ids.to(torch.int64) + elif metadata_change == "sorted_token_ids_size": + sorted_token_ids = torch.empty(9, dtype=torch.int32, device="cuda") + elif metadata_change == "expert_map_shape": + expert_map = expert_map[:3] + + if metadata_change == "expert_map_presence": + operator( + topk_ids, + num_experts, + block_size, + sorted_token_ids, + expert_ids, + num_tokens_post_pad, + ) + else: + operator( + topk_ids, + expert_map, + num_experts, + block_size, + sorted_token_ids, + expert_ids, + num_tokens_post_pad, + ) + torch.cuda.synchronize() + """ +)