Ling-3.0-tiny-RKNN / tests /router_test.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
1.51 kB
#include "ling3/router.h"
#include <algorithm>
#include <array>
#include <cmath>
#include <iostream>
int main() {
std::array<float, ling3::kExpertCount> logits {};
std::array<float, ling3::kExpertCount> bias {};
for (std::size_t group = 0; group < ling3::kExpertGroups; ++group) {
for (std::size_t local = 0; local < ling3::kExpertCount / ling3::kExpertGroups; ++local) {
const std::size_t expert = group * (ling3::kExpertCount / ling3::kExpertGroups) + local;
logits[expert] = -2.0F + static_cast<float>(group) * 0.25F +
static_cast<float>(local) * 0.005F;
}
}
const ling3::Route route = ling3::SelectRoute(logits.data(), bias.data());
std::array<bool, ling3::kExpertCount> seen {};
float sum = 0.0F;
for (std::size_t index = 0; index < route.experts.size(); ++index) {
const int expert = route.experts[index];
if (expert < 0 || expert >= static_cast<int>(ling3::kExpertCount) || seen[expert]) {
std::cerr << "route contains an invalid or duplicate expert\n";
return 1;
}
seen[expert] = true;
sum += route.weights[index];
if (expert / 16 < 4) {
std::cerr << "route selected an expert outside the four strongest groups\n";
return 1;
}
}
if (std::abs(sum - 2.5F) > 1.0e-5F) {
std::cerr << "route weights are not normalized to routed_scale\n";
return 1;
}
return 0;
}