20 std::optional<Tensor> expert_map,
Tensor output)
21 :
num_tokens_{input.ndim() == 3 ? input.size(0) : 0},
22 topk_{input.ndim() == 3 ? input.size(1) : 0},
31 topk_ids && topk_ids->ndim() == 2 ? topk_ids->stride(0) : 0},
33 topk_ids && topk_ids->ndim() == 2 ? topk_ids->stride(1) : 0},
37 expert_map && expert_map->ndim() == 1 ? expert_map->stride(0) : 0},
39 assert(input.ndim() == 3 && output.ndim() == 2 &&
40 "`MoeSum` requires `[num_tokens, topk, hidden_size]` input and "
41 "`[num_tokens, hidden_size]` output");
43 "`MoeSum` output shape is incompatible with the input");
44 assert(
topk_ > 0 &&
"`MoeSum` requires at least one top-k slot");
45 assert((
dtype_ == DataType::kFloat32 ||
dtype_ == DataType::kFloat16 ||
46 dtype_ == DataType::kBFloat16) &&
47 "`MoeSum` supports float32, float16, and bfloat16 inputs");
48 assert(output.dtype() ==
dtype_ &&
49 "`MoeSum` input and output dtypes must match");
50 assert(output.IsContiguous() &&
"`MoeSum` requires contiguous output");
52 constexpr auto kMaxSignedIndex =
53 static_cast<Tensor::Size
>(std::numeric_limits<int64_t>::max());
56 "`MoeSum` dimensions must fit signed index arithmetic");
57 assert(
num_tokens_ <= std::numeric_limits<int>::max() &&
58 "`MoeSum` token count exceeds the CUDA grid limit");
61 "`MoeSum` output size must fit signed index arithmetic");
63 const auto offsets_fit = [](
const auto& shape,
const auto& strides) {
65 constexpr auto kMaxOffset =
66 static_cast<uint64_t>(std::numeric_limits<int64_t>::max());
68 for (Tensor::Size dim = 0; dim < shape.size(); ++dim) {
69 if (
static_cast<uint64_t>(shape[dim]) > kMaxOffset) {
72 if (strides[dim] < 0) {
75 if (shape[dim] == 0) {
79 const auto extent =
static_cast<uint64_t>(shape[dim] - 1);
80 const auto stride =
static_cast<uint64_t>(strides[dim]);
81 if (extent != 0 && stride > kMaxOffset / extent) {
85 const auto term = extent * stride;
86 if (term > kMaxOffset - max_offset) {
94 assert(offsets_fit(input.shape(), input.strides()) &&
95 "`MoeSum` input requires non-negative strides with signed offsets");
97 const auto same_device_as_input = [&](
const Tensor tensor) {
98 return tensor.device().type() == input.device().type() &&
99 tensor.device().index() == input.device().index();
101 assert(same_device_as_input(output) &&
102 "`MoeSum` input and output must be on the same device");
103 assert((!expert_map || topk_ids) &&
104 "`MoeSum` expert_map requires topk_ids");
107 assert(topk_ids->ndim() == 2 && topk_ids->size(0) ==
num_tokens_ &&
108 topk_ids->size(1) ==
topk_ &&
109 "`MoeSum` topk_ids must have shape `[num_tokens, topk]`");
112 "`MoeSum` topk_ids must have int32 or int64 dtype");
113 assert(same_device_as_input(*topk_ids) &&
114 "`MoeSum` topk_ids must be on the input device");
115 assert(offsets_fit(topk_ids->shape(), topk_ids->strides()) &&
116 "`MoeSum` topk_ids requires non-negative strides with signed "
121 assert(expert_map->ndim() == 1 &&
122 expert_map->dtype() == DataType::kInt32 &&
123 "`MoeSum` expert_map must be a 1D int32 tensor");
124 assert(same_device_as_input(*expert_map) &&
125 "`MoeSum` expert_map must be on the input device");
126 assert(offsets_fit(expert_map->shape(), expert_map->strides()) &&
127 "`MoeSum` expert_map requires non-negative strides with signed "
142 std::optional<Tensor> expert_map,
143 const Tensor output)
const {
144 const auto same_device_as_descriptor = [&](
const Tensor tensor) {
149 input.ndim() == 3 && input.size(0) ==
num_tokens_ &&
152 same_device_as_descriptor(input) && output.ndim() == 2 &&
155 same_device_as_descriptor(output) &&
159 if (matches && topk_ids) {
160 matches = topk_ids->ndim() == 2 && topk_ids->size(0) ==
num_tokens_ &&
161 topk_ids->size(1) ==
topk_ &&
165 same_device_as_descriptor(*topk_ids);
168 if (matches && expert_map) {
169 matches = expert_map->ndim() == 1 &&
172 expert_map->dtype() == DataType::kInt32 &&
173 same_device_as_descriptor(*expert_map);
176 assert(matches &&
"`MoeSum` call metadata must match descriptor");