File size: 1,507 Bytes
3fd1a35 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 | #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;
}
|