Compare commits

...
Author SHA1 Message Date
Sigbjørn Skjæret 4ffcbaf2ae add to graph_optimize 2026-09-23 10:01:35 +02:00
Sigbjørn Skjæret 1ce2eb90e0 hip fix 2026-09-23 09:49:41 +02:00
Sigbjørn Skjæret 504e7f062c add tests 2026-09-23 09:49:38 +02:00
Sigbjørn Skjæret c8ff7916a0 fuse grouped experts 2026-09-23 09:49:34 +02:00
4 changed files with 667 additions and 67 deletions
+377 -20
View File
@@ -2813,6 +2813,261 @@ static int ggml_cuda_try_gdn_cache_fusion(
return skip;
}
static bool ggml_cuda_topk_moe_grouped_experts(
const struct ggml_cgraph * cgraph,
int & node_idx,
const ggml_tensor * selection_src,
const ggml_tensor * probs_reshaped,
ggml_cuda_topk_moe_args & args) {
const int n_nodes = cgraph->n_nodes;
ggml_tensor ** nodes = cgraph->nodes;
// group reshape: [n_exp_per_group, n_expert_groups, n_tokens]
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_RESHAPE ||
nodes[node_idx]->src[0] != selection_src) {
return false;
}
const ggml_tensor * gr = nodes[node_idx];
const int64_t E = probs_reshaped->ne[1];
const int64_t P = gr->ne[0];
const int64_t G = gr->ne[1];
const int64_t T = gr->ne[2];
if (probs_reshaped->ne[0] != 1 ||
probs_reshaped->ne[2] != T ||
probs_reshaped->ne[3] != 1) {
return false;
}
if (P <= 0 || G <= 1 || T <= 0 || E <= 0 || P * G != E) {
return false;
}
// The current kernel assigns one thread per group.
if (G > WARP_SIZE) {
return false;
}
++node_idx;
// reshape to [1, n_exp_per_group, n_expert_groups, n_tokens]
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_RESHAPE ||
nodes[node_idx]->src[0] != gr) {
return false;
}
const ggml_tensor * r4 = nodes[node_idx];
if (r4->ne[0] != 1 || r4->ne[1] != P || r4->ne[2] != G || r4->ne[3] != T) {
return false;
}
++node_idx;
// argsort over experts within groups
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_ARGSORT ||
nodes[node_idx]->src[0] != gr) {
return false;
}
const ggml_tensor * ga_arg = nodes[node_idx];
++node_idx;
// view top-2 per group
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_VIEW ||
nodes[node_idx]->src[0] != ga_arg) {
return false;
}
const ggml_tensor * g2 = nodes[node_idx];
if (g2->ne[0] != 2 ||
g2->ne[1] != G || g2->ne[2] != T) {
return false;
}
++node_idx;
// gather top-k expert values per group
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_GET_ROWS) {
return false;
}
const ggml_tensor * get1 = nodes[node_idx];
if (get1->src[0] != r4 || get1->src[1] != g2) {
return false;
}
if (get1->ne[0] != 1 ||
get1->ne[1] != 2 ||
get1->ne[2] != G ||
get1->ne[3] != T) {
return false;
}
++node_idx;
// reshape to [2, n_expert_groups, n_tokens]
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_RESHAPE ||
nodes[node_idx]->src[0] != get1) {
return false;
}
const ggml_tensor * r5 = nodes[node_idx];
if (r5->ne[0] != 2 || r5->ne[1] != G || r5->ne[2] != T) {
return false;
}
++node_idx;
// sum top-k expert values per group
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_SUM_ROWS ||
nodes[node_idx]->src[0] != r5) {
return false;
}
const ggml_tensor * sum = nodes[node_idx];
if (sum->ne[0] != 1 || sum->ne[1] != G || sum->ne[2] != T) {
return false;
}
++node_idx;
// reshape group scores to [n_expert_groups, n_tokens]
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_RESHAPE ||
nodes[node_idx]->src[0] != sum) {
return false;
}
const ggml_tensor * r6 = nodes[node_idx];
if (r6->ne[0] != G || r6->ne[1] != T || r6->ne[2] != 1) {
return false;
}
++node_idx;
// argsort groups
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_ARGSORT ||
nodes[node_idx]->src[0] != r6) {
return false;
}
const ggml_tensor * earg = nodes[node_idx];
++node_idx;
// view top n_group_used groups
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_VIEW ||
nodes[node_idx]->src[0] != earg) {
return false;
}
const ggml_tensor * eg = nodes[node_idx];
if (eg->ne[0] < 1 || eg->ne[0] > G || eg->ne[1] != T) {
return false;
}
const int n_group_used = (int) eg->ne[0];
++node_idx;
// gather selected expert-group rows
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_GET_ROWS) {
return false;
}
const ggml_tensor * get2 = nodes[node_idx];
if (get2->src[0] != gr || get2->src[1] != eg) {
return false;
}
if (get2->ne[0] != P ||
get2->ne[1] != n_group_used ||
get2->ne[2] != T) {
return false;
}
++node_idx;
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_FILL ||
nodes[node_idx]->src[0] != gr) {
return false;
}
const ggml_tensor * fill = nodes[node_idx];
if (fill->ne[0] != P || fill->ne[1] != G || fill->ne[2] != T) {
return false;
}
if (ggml_get_op_params_f32(fill, 0) != -INFINITY) {
return false;
}
++node_idx;
// keep selected groups, -INFINITY elsewhere
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_SET_ROWS) {
return false;
}
const ggml_tensor * sel = nodes[node_idx];
if (sel->src[0] != get2 ||
sel->src[1] != eg ||
sel->src[2] != fill) {
return false;
}
if (sel->ne[0] != P || sel->ne[1] != G || sel->ne[2] != T) {
return false;
}
++node_idx;
// reshape masked scores back to [n_expert, n_tokens]
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_RESHAPE ||
nodes[node_idx]->src[0] != sel) {
return false;
}
const ggml_tensor * r7 = nodes[node_idx];
if (r7->ne[0] != E || r7->ne[1] != T || r7->ne[2] != 1) {
return false;
}
++node_idx;
args.grouped_experts = true;
args.n_expert_groups = (int) G;
args.n_exp_per_group = (int) P;
args.n_group_used = n_group_used;
return true;
}
static bool ggml_cuda_topk_moe_fusion(const struct ggml_cgraph * cgraph, int node_idx, ggml_cuda_topk_moe_args & args) {
args.sigmoid = false;
args.sqrt_softplus = false;
@@ -2820,6 +3075,7 @@ static bool ggml_cuda_topk_moe_fusion(const struct ggml_cgraph * cgraph, int nod
args.delayed_softmax = false;
args.prob_bias = false;
args.norm = false;
args.grouped_experts = false;
const int n_nodes = cgraph->n_nodes;
ggml_tensor ** nodes = cgraph->nodes;
@@ -2866,15 +3122,41 @@ static bool ggml_cuda_topk_moe_fusion(const struct ggml_cgraph * cgraph, int nod
args.prob_bias = true;
node_idx++;
}
// RESHAPE/ADD -> ARGSORT
if (node_idx >= n_nodes || nodes[node_idx]->op != GGML_OP_ARGSORT) {
return false;
const ggml_tensor * selection_src = args.prob_bias ? nodes[node_idx - 1] : nodes[node_idx - 2];
bool grouped_experts = false;
if (node_idx < n_nodes &&
nodes[node_idx]->op == GGML_OP_RESHAPE &&
nodes[node_idx]->src[0] == selection_src) {
const int saved_node_idx = node_idx;
grouped_experts = ggml_cuda_topk_moe_grouped_experts(cgraph, node_idx, selection_src, probs_reshaped, args);
if (!grouped_experts) {
node_idx = saved_node_idx;
}
}
if (args.prob_bias && nodes[node_idx]->src[0] != nodes[node_idx - 1]) {
return false;
} else if (!args.prob_bias && nodes[node_idx]->src[0] != nodes[node_idx - 2]) {
return false;
if (!grouped_experts) {
// RESHAPE/ADD -> ARGSORT
if (node_idx >= n_nodes || nodes[node_idx]->op != GGML_OP_ARGSORT) {
return false;
}
if (args.prob_bias && nodes[node_idx]->src[0] != nodes[node_idx - 1]) {
return false;
} else if (!args.prob_bias && nodes[node_idx]->src[0] != nodes[node_idx - 2]) {
return false;
}
} else {
// grouped path leaves node_idx at the final expert ARGSORT
if (node_idx >= n_nodes ||
nodes[node_idx]->op != GGML_OP_ARGSORT ||
nodes[node_idx]->src[0] != nodes[node_idx - 1]) {
return false;
}
}
node_idx++;
@@ -3491,14 +3773,52 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
}
const int i_probs = i + (int) ops.size() - 1; // last node of the gating activation
if (args.prob_bias) {
bias = cgraph->nodes[i_probs + 2]->src[1];
ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ADD, GGML_OP_ARGSORT, GGML_OP_VIEW,
GGML_OP_GET_ROWS });
out_nodes[0] = i_probs + 4;
if (args.grouped_experts) {
// weights reshape: [1, n_expert, n_tokens]
ops.push_back(GGML_OP_RESHAPE);
if (args.prob_bias) {
ops.push_back(GGML_OP_ADD);
bias = cgraph->nodes[i + (int) ops.size() - 1]->src[1];
}
// grouped top-k sequence
ops.insert(ops.end(), {
GGML_OP_RESHAPE, // selection_groups
GGML_OP_RESHAPE, // 4d rows for get_rows
GGML_OP_ARGSORT, // top-2 per group
GGML_OP_VIEW, // top-2 per group view
GGML_OP_GET_ROWS, // gather top-2 per group values
GGML_OP_RESHAPE, // [2, n_groups, n_tokens]
GGML_OP_SUM_ROWS, // group score
GGML_OP_RESHAPE, // [n_groups, n_tokens]
GGML_OP_ARGSORT, // top groups
GGML_OP_VIEW, // top groups view
GGML_OP_GET_ROWS, // gather selected group rows
GGML_OP_FILL, // -INFINITY template
GGML_OP_SET_ROWS, // mask unselected groups
GGML_OP_RESHAPE // [n_expert, n_tokens]
});
// final expert top-k
ops.push_back(GGML_OP_ARGSORT);
const int view_idx = i + (int) ops.size();
ops.push_back(GGML_OP_VIEW);
ops.push_back(GGML_OP_GET_ROWS);
out_nodes[0] = view_idx;
} else {
ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS });
out_nodes[0] = i_probs + 3;
if (args.prob_bias) {
bias = cgraph->nodes[i_probs + 2]->src[1];
ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ADD, GGML_OP_ARGSORT, GGML_OP_VIEW,
GGML_OP_GET_ROWS });
out_nodes[0] = i_probs + 4;
} else {
ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS });
out_nodes[0] = i_probs + 3;
}
}
ids = cgraph->nodes[out_nodes[0]];
@@ -4559,13 +4879,50 @@ static void ggml_backend_cuda_graph_optimize(ggml_backend_t backend, ggml_cgraph
}
const int i_probs = i + (int) ops.size() - 1; // last node of the gating activation
if (args.prob_bias) {
ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ADD, GGML_OP_ARGSORT, GGML_OP_VIEW,
GGML_OP_GET_ROWS });
out_nodes[0] = i_probs + 4;
if (args.grouped_experts) {
// weights reshape: [1, n_expert, n_tokens]
ops.push_back(GGML_OP_RESHAPE);
if (args.prob_bias) {
ops.push_back(GGML_OP_ADD);
}
// grouped top-k sequence
ops.insert(ops.end(), {
GGML_OP_RESHAPE, // selection_groups
GGML_OP_RESHAPE, // 4d rows for get_rows
GGML_OP_ARGSORT, // top-2 per group
GGML_OP_VIEW, // top-2 per group view
GGML_OP_GET_ROWS, // gather top-2 per group values
GGML_OP_RESHAPE, // [2, n_groups, n_tokens]
GGML_OP_SUM_ROWS, // group score
GGML_OP_RESHAPE, // [n_groups, n_tokens]
GGML_OP_ARGSORT, // top groups
GGML_OP_VIEW, // top groups view
GGML_OP_GET_ROWS, // gather selected group rows
GGML_OP_FILL, // -INFINITY template
GGML_OP_SET_ROWS, // mask unselected groups
GGML_OP_RESHAPE // [n_expert, n_tokens]
});
// final expert top-k
ops.push_back(GGML_OP_ARGSORT);
const int view_idx = i + (int) ops.size();
ops.push_back(GGML_OP_VIEW);
ops.push_back(GGML_OP_GET_ROWS);
out_nodes[0] = view_idx;
} else {
ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS });
out_nodes[0] = i_probs + 3;
if (args.prob_bias) {
ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ADD, GGML_OP_ARGSORT, GGML_OP_VIEW,
GGML_OP_GET_ROWS });
out_nodes[0] = i_probs + 4;
} else {
ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS });
out_nodes[0] = i_probs + 3;
}
}
ids = cgraph->nodes[out_nodes[0]];
+253 -44
View File
@@ -11,6 +11,11 @@ struct topk_moe_config {
bool use_sqrt_softplus;
bool with_norm;
bool delayed_softmax;
bool grouped_experts;
int n_expert_groups;
int n_exp_per_group;
int n_group_used;
};
// Warp-local softmax used for both the pre-top-k logits and the post-top-k delayed path.
@@ -78,12 +83,74 @@ __device__ void sqrt_softplus_warp_inplace(float (&vals)[experts_per_thread], co
}
}
__device__ __forceinline__ void topk_moe_top2_merge(float &a1, float &a2, float b1, float b2) {
// Merge two sorted top-2 pairs:
// a1 >= a2
// b1 >= b2
if (b1 > a1) {
a2 = max(a1, b2);
a1 = b1;
} else if (b1 > a2) {
a2 = b1;
}
if (b2 > a2) {
a2 = b2;
}
}
template <int experts_per_thread, bool has_bias>
__device__ __forceinline__ float topk_moe_group_score(const float (&wt)[experts_per_thread],
const float * bias,
const int n_experts,
const int n_exp_per_group,
const int group,
const int lane) {
const int start = group * n_exp_per_group;
const int end = start + n_exp_per_group;
float top1 = -INFINITY;
float top2 = -INFINITY;
#pragma unroll
for (int i = 0; i < experts_per_thread; ++i) {
const int e = i * WARP_SIZE + lane;
if (e >= start && e < end && e < n_experts) {
float v = wt[i];
if constexpr (has_bias) {
v += bias[e];
}
if (v > top1) {
top2 = top1;
top1 = v;
} else if (v > top2) {
top2 = v;
}
}
}
// Reduce the top-2 pair across the whole warp.
#pragma unroll
for (int mask = WARP_SIZE / 2; mask > 0; mask /= 2) {
const float o1 = __shfl_xor_sync(0xFFFFFFFF, top1, mask, WARP_SIZE);
const float o2 = __shfl_xor_sync(0xFFFFFFFF, top2, mask, WARP_SIZE);
topk_moe_top2_merge(top1, top2, o1, o2);
}
return (top2 == -INFINITY) ? top1 : (top1 + top2);
}
/*
This kernel does the following:
1. optionally softmax over the logits per token [n_experts, n_tokens]
2. argmax reduce over the top-k (n_experts_used) logits
3. write weights + ids to global memory
4. optionally normalize the weights or apply softmax over the selected logits
5. optionally select top-k grouped experts (n_group_used)
It is intended as fusion of softmax->top-k->get_rows pipeline for MoE models
*/
@@ -179,72 +246,208 @@ __global__ void topk_moe_cuda(const float * logits,
output_weights[i] = 0.f;
}
ggml_cuda_pdl_lc();
for (int k = 0; k < n_expert_used; k++) {
float max_val = wt[0];
int max_expert = threadIdx.x;
if (config.grouped_experts) {
const int G = config.n_expert_groups;
const int P = config.n_exp_per_group;
if constexpr (has_bias) {
float max_val_s = selection_wt[0];
// The score for each group must be reduced across the whole warp because
// a group is distributed over many lanes.
float gval = -INFINITY;
for (int g = 0; g < G; ++g) {
const float score = topk_moe_group_score<experts_per_thread, has_bias>(
wt, bias, n_experts, P, g, threadIdx.x
);
if (threadIdx.x == g) {
gval = score;
}
}
unsigned selected_mask = 0;
// Select top n_group_used groups.
for (int k = 0; k < config.n_group_used; ++k) {
float best = -INFINITY;
int best_group = WARP_SIZE;
int my_group = WARP_SIZE;
float my_score = -INFINITY;
if (threadIdx.x < G) {
// Skip groups that were already selected in a previous round.
if ((selected_mask & (1u << threadIdx.x)) == 0) {
my_group = threadIdx.x;
my_score = gval;
}
}
if (my_score > best || (my_score == best && my_group < best_group)) {
best = my_score;
best_group = my_group;
}
#pragma unroll
for (int i = 1; i < experts_per_thread; i++) {
for (int mask = WARP_SIZE / 2; mask > 0; mask /= 2) {
const float other = __shfl_xor_sync(0xFFFFFFFF, best, mask, WARP_SIZE);
const int other_group = __shfl_xor_sync(0xFFFFFFFF, best_group, mask, WARP_SIZE);
if (other > best || (other == best && other_group < best_group)) {
best = other;
best_group = other_group;
}
}
if (best_group < G) {
selected_mask |= (1u << best_group);
if (threadIdx.x == best_group) {
gval = -INFINITY;
}
}
}
// Build the masked expert selection scores.
// If bias is present, selection uses wt + bias, but output weights use wt.
float sel[experts_per_thread];
#pragma unroll
for (int i = 0; i < experts_per_thread; ++i) {
const int e = i * WARP_SIZE + threadIdx.x;
if ((n_experts % WARP_SIZE == 0 || e < n_experts) && P > 0) {
const int group = e / P;
if (group < G && ((selected_mask >> group) & 1u)) {
float v = wt[i];
if constexpr (has_bias) {
v += bias[e];
}
sel[i] = v;
} else {
sel[i] = -INFINITY;
}
} else {
sel[i] = -INFINITY;
}
}
ggml_cuda_pdl_lc();
// Final top-k expert selection over masked scores.
for (int k = 0; k < n_expert_used; ++k) {
float max_w = -INFINITY;
float max_s = -INFINITY;
int max_expert = threadIdx.x;
#pragma unroll
for (int i = 0; i < experts_per_thread; ++i) {
const int expert = threadIdx.x + i * WARP_SIZE;
if ((n_experts % WARP_SIZE == 0 || expert < n_experts) && selection_wt[i] > max_val_s) {
max_val = wt[i];
max_val_s = selection_wt[i];
if ((n_experts % WARP_SIZE == 0 || expert < n_experts) && sel[i] > max_s) {
max_s = sel[i];
max_w = wt[i];
max_expert = expert;
}
}
#pragma unroll
for (int mask = WARP_SIZE / 2; mask > 0; mask /= 2) {
const float val = __shfl_xor_sync(0xFFFFFFFF, max_val, mask, WARP_SIZE);
const float val_s = __shfl_xor_sync(0xFFFFFFFF, max_val_s, mask, WARP_SIZE);
const int expert = __shfl_xor_sync(0xFFFFFFFF, max_expert, mask, WARP_SIZE);
if (val_s > max_val_s || (val_s == max_val_s && expert < max_expert)) {
max_val = val;
max_val_s = val_s;
max_expert = expert;
const float other_s = __shfl_xor_sync(0xFFFFFFFF, max_s, mask, WARP_SIZE);
const float other_w = __shfl_xor_sync(0xFFFFFFFF, max_w, mask, WARP_SIZE);
const int other_exp = __shfl_xor_sync(0xFFFFFFFF, max_expert, mask, WARP_SIZE);
if (other_s > max_s || (other_s == max_s && other_exp < max_expert)) {
max_s = other_s;
max_w = other_w;
max_expert = other_exp;
}
}
if ((k & (WARP_SIZE - 1)) == threadIdx.x) {
output_weights[k / WARP_SIZE] = max_w;
}
if ((max_expert & (WARP_SIZE - 1)) == threadIdx.x) {
selection_wt[max_expert / WARP_SIZE] = -INFINITY;
ids[k] = max_expert;
if (config.with_norm) {
wt_sum += max_w;
}
sel[max_expert / WARP_SIZE] = -INFINITY;
}
} else {
}
} else {
ggml_cuda_pdl_lc();
for (int k = 0; k < n_expert_used; k++) {
float max_val = wt[0];
int max_expert = threadIdx.x;
if constexpr (has_bias) {
float max_val_s = selection_wt[0];
#pragma unroll
for (int i = 1; i < experts_per_thread; i++) {
const int expert = threadIdx.x + i * WARP_SIZE;
if ((n_experts % WARP_SIZE == 0 || expert < n_experts) && wt[i] > max_val) {
max_val = wt[i];
max_expert = expert;
for (int i = 1; i < experts_per_thread; i++) {
const int expert = threadIdx.x + i * WARP_SIZE;
if ((n_experts % WARP_SIZE == 0 || expert < n_experts) && selection_wt[i] > max_val_s) {
max_val = wt[i];
max_val_s = selection_wt[i];
max_expert = expert;
}
}
#pragma unroll
for (int mask = WARP_SIZE / 2; mask > 0; mask /= 2) {
const float val = __shfl_xor_sync(0xFFFFFFFF, max_val, mask, WARP_SIZE);
const float val_s = __shfl_xor_sync(0xFFFFFFFF, max_val_s, mask, WARP_SIZE);
const int expert = __shfl_xor_sync(0xFFFFFFFF, max_expert, mask, WARP_SIZE);
if (val_s > max_val_s || (val_s == max_val_s && expert < max_expert)) {
max_val = val;
max_val_s = val_s;
max_expert = expert;
}
}
if ((max_expert & (WARP_SIZE - 1)) == threadIdx.x) {
selection_wt[max_expert / WARP_SIZE] = -INFINITY;
}
} else {
#pragma unroll
for (int i = 1; i < experts_per_thread; i++) {
const int expert = threadIdx.x + i * WARP_SIZE;
if ((n_experts % WARP_SIZE == 0 || expert < n_experts) && wt[i] > max_val) {
max_val = wt[i];
max_expert = expert;
}
}
#pragma unroll
for (int mask = WARP_SIZE / 2; mask > 0; mask /= 2) {
const float val = __shfl_xor_sync(0xFFFFFFFF, max_val, mask, WARP_SIZE);
const int expert = __shfl_xor_sync(0xFFFFFFFF, max_expert, mask, WARP_SIZE);
if (val > max_val || (val == max_val && expert < max_expert)) {
max_val = val;
max_expert = expert;
}
}
if ((max_expert & (WARP_SIZE - 1)) == threadIdx.x) {
wt[max_expert / WARP_SIZE] = -INFINITY;
}
}
#pragma unroll
for (int mask = WARP_SIZE / 2; mask > 0; mask /= 2) {
const float val = __shfl_xor_sync(0xFFFFFFFF, max_val, mask, WARP_SIZE);
const int expert = __shfl_xor_sync(0xFFFFFFFF, max_expert, mask, WARP_SIZE);
if (val > max_val || (val == max_val && expert < max_expert)) {
max_val = val;
max_expert = expert;
}
if ((k & (WARP_SIZE - 1)) == threadIdx.x) {
output_weights[k / WARP_SIZE] = max_val;
}
if ((max_expert & (WARP_SIZE - 1)) == threadIdx.x) {
wt[max_expert / WARP_SIZE] = -INFINITY;
}
}
if ((k & (WARP_SIZE - 1)) == threadIdx.x) {
output_weights[k / WARP_SIZE] = max_val;
}
if ((max_expert & (WARP_SIZE - 1)) == threadIdx.x) {
ids[k] = max_expert;
if (config.with_norm) {
wt_sum += max_val;
ids[k] = max_expert;
if (config.with_norm) {
wt_sum += max_val;
}
}
}
}
@@ -385,6 +588,12 @@ void ggml_cuda_op_topk_moe(ggml_backend_cuda_context & ctx,
config.use_sqrt_softplus = args.sqrt_softplus;
config.with_norm = with_norm;
config.delayed_softmax = args.delayed_softmax;
config.grouped_experts = args.grouped_experts;
config.n_expert_groups = args.n_expert_groups;
config.n_exp_per_group = args.n_exp_per_group;
config.n_group_used = args.n_group_used;
GGML_ASSERT(!(config.grouped_experts && config.delayed_softmax));
if (bias) {
launch_topk_moe_cuda<true>(ctx, logits_d, weights_d, ids_d, bias_d, n_rows, n_experts, n_expert_used, clamp_val,
+5
View File
@@ -14,6 +14,11 @@ struct ggml_cuda_topk_moe_args {
bool prob_bias{};
bool norm{};
bool scale{};
bool grouped_experts{};
int n_expert_groups{};
int n_exp_per_group{};
int n_group_used{};
};
void ggml_cuda_op_topk_moe(ggml_backend_cuda_context & ctx,
+32 -3
View File
@@ -6778,6 +6778,7 @@ struct test_topk_moe : public test_case {
const bool bias_probs;
const MoeGatingFunc gating_func;
const float scale_w;
const int n_expert_groups;
ggml_tensor * weights {};
ggml_tensor * selected_experts {};
@@ -6786,17 +6787,21 @@ struct test_topk_moe : public test_case {
bool with_norm = false,
bool bias_probs = false,
MoeGatingFunc gating_func = GATING_FUNC_SOFTMAX,
float scale_w = 0.0f) :
float scale_w = 0.0f,
int n_expert_groups = 1) :
ne(ne),
n_expert_used(n_expert_used),
with_norm(with_norm),
bias_probs(bias_probs),
gating_func(gating_func),
scale_w(scale_w) {
scale_w(scale_w),
n_expert_groups(n_expert_groups) {
GGML_ASSERT(n_expert_used <= ne[0]);
GGML_ASSERT(ne[0] % n_expert_groups == 0);
GGML_ASSERT(n_expert_used <= ne[0] / n_expert_groups);
}
std::string vars() override { return VARS_TO_STR6(ne, n_expert_used, with_norm, bias_probs, gating_func, scale_w); }
std::string vars() override { return VARS_TO_STR7(ne, n_expert_used, with_norm, bias_probs, gating_func, scale_w, n_expert_groups); }
std::string op_desc(ggml_tensor * t) override {
GGML_UNUSED(t);
@@ -6808,6 +6813,8 @@ struct test_topk_moe : public test_case {
ggml_tensor * build_graph(ggml_context * ctx) override {
const int n_expert = ne[0];
const int n_tokens = ne[1];
const int n_exp_per_group = n_expert / n_expert_groups;
const int n_group_used = n_expert_groups / 2;
ggml_tensor * logits = ggml_new_tensor(ctx, GGML_TYPE_F32, 4, ne.data());
ggml_tensor * probs =
@@ -6824,6 +6831,24 @@ struct test_topk_moe : public test_case {
ggml_set_name(selection_probs, "selection_probs");
}
if (n_expert_groups > 1) {
ggml_tensor * selection_groups = ggml_reshape_3d(ctx, selection_probs, n_exp_per_group, n_expert_groups, n_tokens); // [n_exp_per_group, n_expert_groups, n_tokens]
ggml_tensor * group_scores = ggml_argsort_top_k(ctx, selection_groups, 2); // [2, n_expert_groups, n_tokens]
group_scores = ggml_get_rows(ctx, ggml_reshape_4d(ctx, selection_groups, 1, selection_groups->ne[0], selection_groups->ne[1], selection_groups->ne[2]), group_scores); // [1, 2, n_expert_groups, n_tokens]
group_scores = ggml_sum_rows(ctx, ggml_reshape_3d(ctx, group_scores, group_scores->ne[1], group_scores->ne[2], group_scores->ne[3])); // [1, n_expert_groups, n_tokens]
group_scores = ggml_reshape_2d(ctx, group_scores, group_scores->ne[1], group_scores->ne[2]); // [n_expert_groups, n_tokens]
ggml_tensor * expert_groups = ggml_argsort_top_k(ctx, group_scores, n_group_used); // [n_group_used, n_tokens]
ggml_set_name(expert_groups, "ffn_moe_group_topk");
selection_probs = ggml_get_rows(ctx, selection_groups, expert_groups); // [n_exp_per_group, n_group_used, n_tokens]
selection_probs = ggml_set_rows(ctx, ggml_fill(ctx, selection_groups, -INFINITY), selection_probs, expert_groups); // [n_exp_per_group, n_expert_groups, n_tokens]
selection_probs = ggml_reshape_2d(ctx, selection_probs, n_expert, n_tokens); // [n_expert, n_tokens]
ggml_set_name(selection_probs, "ffn_moe_probs_masked");
}
selected_experts = ggml_argsort_top_k(ctx, selection_probs, n_expert_used); // [n_expert_used, n_tokens]
ggml_set_name(selected_experts, "selected_experts");
@@ -10987,6 +11012,10 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
test_cases.emplace_back(new test_topk_moe({32, 8, 1, 1}, 4, with_norm, bias_probs, gate, scale_w));
test_cases.emplace_back(new test_topk_moe({32, 8, 1, 1}, 8, with_norm, bias_probs, gate, scale_w));
test_cases.emplace_back(new test_topk_moe({32, 9, 1, 1}, 8, with_norm, bias_probs, gate, scale_w));
// grouped experts
test_cases.emplace_back(new test_topk_moe({128, 1, 1, 1}, 8, with_norm, bias_probs, gate, scale_w, 16));
test_cases.emplace_back(new test_topk_moe({256, 22, 1, 1}, 6, with_norm, bias_probs, gate, scale_w, 32));
}
}
}