#include "ling3/router.h" #include #include #include #include int main() { std::array logits {}; std::array 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(group) * 0.25F + static_cast(local) * 0.005F; } } const ling3::Route route = ling3::SelectRoute(logits.data(), bias.data()); std::array 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(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; }