| 338 | } |
| 339 | |
| 340 | void 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 |
nothing calls this directly
no test coverage detected