#include "ling3/router.h" #include #include #include #include #if defined(__aarch64__) #include #endif namespace ling3 { void RouterLogitsF32(const float * hidden, const float * weights, float * logits) noexcept { for (std::size_t expert = 0; expert < kExpertCount; ++expert) { const float * row = weights + expert * kHiddenSize; float sum = 0.0F; #if defined(__aarch64__) float32x4_t sum0 = vdupq_n_f32(0.0F); float32x4_t sum1 = vdupq_n_f32(0.0F); for (std::size_t index = 0; index < kHiddenSize; index += 8) { sum0 = vfmaq_f32(sum0, vld1q_f32(hidden + index), vld1q_f32(row + index)); sum1 = vfmaq_f32(sum1, vld1q_f32(hidden + index + 4), vld1q_f32(row + index + 4)); } sum = vaddvq_f32(vaddq_f32(sum0, sum1)); #else for (std::size_t index = 0; index < kHiddenSize; ++index) sum += hidden[index] * row[index]; #endif logits[expert] = sum; } } Route SelectRoute(const float * logits, const float * expert_bias, float routed_scale) noexcept { std::array scores{}; std::array routing_scores{}; for (std::size_t expert = 0; expert < kExpertCount; ++expert) { const float score = 1.0F / (1.0F + std::exp(-logits[expert])); scores[expert] = score; routing_scores[expert] = score + expert_bias[expert]; } std::array group_scores{}; for (std::size_t group = 0; group < kExpertGroups; ++group) { float first = -std::numeric_limits::infinity(); float second = first; for (std::size_t local = 0; local < kExpertCount / kExpertGroups; ++local) { const float value = routing_scores[group * (kExpertCount / kExpertGroups) + local]; if (value > first) { second = first; first = value; } else if (value > second) { second = value; } } group_scores[group] = first + second; } std::array selected_groups{}; for (std::size_t selection = 0; selection < kSelectedGroups; ++selection) { std::size_t best_group = 0; float best = -std::numeric_limits::infinity(); for (std::size_t group = 0; group < kExpertGroups; ++group) { if (!selected_groups[group] && group_scores[group] > best) { best = group_scores[group]; best_group = group; } } selected_groups[best_group] = true; } Route route; std::array selected_experts{}; float weight_sum = 0.0F; for (std::size_t selection = 0; selection < kExpertsPerToken; ++selection) { int best_expert = 0; float best = -std::numeric_limits::infinity(); for (std::size_t expert = 0; expert < kExpertCount; ++expert) { const std::size_t group = expert / (kExpertCount / kExpertGroups); if (selected_groups[group] && !selected_experts[expert] && routing_scores[expert] > best) { best = routing_scores[expert]; best_expert = static_cast(expert); } } route.experts[selection] = best_expert; selected_experts[static_cast(best_expert)] = true; route.weights[selection] = scores[static_cast(best_expert)]; weight_sum += route.weights[selection]; } const float normalization = routed_scale / std::max(weight_sum, 1.0e-20F); for (float & weight : route.weights) weight *= normalization; return route; } } // namespace ling3