Ling-3.0-tiny-RKNN / src /router.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
3.71 kB
#include "ling3/router.h"
#include <algorithm>
#include <array>
#include <cmath>
#include <limits>
#if defined(__aarch64__)
#include <arm_neon.h>
#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<float, kExpertCount> scores{};
std::array<float, kExpertCount> 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<float, kExpertGroups> group_scores{};
for (std::size_t group = 0; group < kExpertGroups; ++group) {
float first = -std::numeric_limits<float>::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<bool, kExpertGroups> selected_groups{};
for (std::size_t selection = 0; selection < kSelectedGroups; ++selection) {
std::size_t best_group = 0;
float best = -std::numeric_limits<float>::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<bool, kExpertCount> 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<float>::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<int>(expert);
}
}
route.experts[selection] = best_expert;
selected_experts[static_cast<std::size_t>(best_expert)] = true;
route.weights[selection] = scores[static_cast<std::size_t>(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