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;
}