File size: 3,711 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
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
#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