MCPcopy Create free account
hub / github.com/cactus-compute/cactus / compute_moe_layer_node

Function compute_moe_layer_node

cactus/graph/graph_ops_nn.cpp:340–512  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

338}
339
340void compute_moe_layer_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes, const std::unordered_map<size_t, size_t>& node_index_map) {
341 const size_t num_experts = node.params.num_experts;
342 const size_t top_k = node.params.num_experts_per_tok;
343 const bool normalize_routing = node.params.normalize_routing;
344 const float eps = node.params.epsilon;
345 const float routed_scaling_factor = node.params.scalar;
346 const bool gated = node.params.moe_gated;
347 const Activation activation = node.params.activation;
348 const size_t base_inputs = gated ? (3 + 3 * num_experts) : (3 + 2 * num_experts);
349 bool has_per_expert_scale = node.input_ids.size() == base_inputs + 1;
350 if (node.input_ids.size() != base_inputs && node.input_ids.size() != base_inputs + 1) {
351 throw std::runtime_error("moe_layer expects " + std::to_string(base_inputs) + " or " + std::to_string(base_inputs + 1) + " inputs, got " + std::to_string(node.input_ids.size()));
352 }
353
354 const auto& hidden_buffer = get_input(node, 0, nodes, node_index_map);
355 const auto& routing_buffer = get_input(node, 1, nodes, node_index_map);
356 const auto& topk_idx_buffer = get_input(node, 2, nodes, node_index_map);
357
358 if (hidden_buffer.precision != Precision::FP16 || node.output_buffer.precision != Precision::FP16) {
359 throw std::runtime_error("moe_layer expects FP16 hidden/output");
360 }
361 if (topk_idx_buffer.precision != Precision::FP32) {
362 throw std::runtime_error("moe_layer expects FP32 topk indices");
363 }
364
365 const __fp16* expert_scales_fp16 = nullptr;
366 if (has_per_expert_scale) {
367 const auto& scale_buffer = get_input(node, base_inputs, nodes, node_index_map);
368 if (scale_buffer.precision != Precision::FP16) {
369 throw std::runtime_error("moe_layer expects FP16 per_expert_scale");
370 }
371 expert_scales_fp16 = scale_buffer.data_as<__fp16>();
372 }
373
374 const size_t token_count = hidden_buffer.shape[0];
375 const size_t hidden_dim = hidden_buffer.shape[1];
376 const size_t total_num_experts = routing_buffer.shape[1];
377
378 const auto& w1_0_buffer = get_input(node, 3, nodes, node_index_map);
379 const size_t expert_intermediate_dim = w1_0_buffer.shape[0];
380
381 const auto* hidden = hidden_buffer.data_as<__fp16>();
382 auto* output = node.output_buffer.data_as<__fp16>();
383 const auto* topk_idx = topk_idx_buffer.data_as<float>();
384 const auto* routing_fp16 = routing_buffer.precision == Precision::FP16 ? routing_buffer.data_as<__fp16>() : nullptr;
385 const auto* routing_fp32 = routing_buffer.precision == Precision::FP32 ? routing_buffer.data_as<float>() : nullptr;
386
387 auto routing_prob = [&](size_t tok, size_t exp) -> float {
388 const size_t offset = tok * total_num_experts + exp;
389 if (routing_fp16) return static_cast<float>(routing_fp16[offset]);
390 return routing_fp32[offset];
391 };
392
393 ensure_moe_buffers(token_count, hidden_dim, expert_intermediate_dim, num_experts, top_k);
394
395 size_t* expert_offsets = moe_expert_offsets_buf.data();
396 size_t* expert_tokens_flat = moe_expert_tokens_buf.data();
397

Callers

nothing calls this directly

Calls 11

ensure_moe_buffersFunction · 0.85
moe_matmulFunction · 0.85
cactus_gelu_f16Function · 0.85
cactus_gelu_f16_erfFunction · 0.85
cactus_relu_f16Function · 0.85
cactus_silu_f16Function · 0.85
cactus_multiply_f16Function · 0.85
cactus_add_scaled_f16Function · 0.85
sizeMethod · 0.80
dataMethod · 0.80
is_grouped_int8Method · 0.80

Tested by

no test coverage detected