File size: 3,393 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
97
98
99
100
#include "ling3/cpu_kernels.h"

#include <cmath>
#include <cstring>
#include <cstdlib>

#if defined(__aarch64__)
#include <arm_neon.h>
#endif

#if defined(LING3_WITH_VECTOR_EXP)
extern "C" float32x4_t _ZGVnN4v_expf(float32x4_t);
#endif

namespace ling3 {

float BFloat16ToFloat(std::uint16_t value) noexcept {
    std::uint32_t bits = static_cast<std::uint32_t>(value) << 16U;
    float result;
    std::memcpy(&result, &bits, sizeof(result));
    return result;
}

std::uint16_t FloatToBFloat16(float value) noexcept {
    std::uint32_t bits;
    std::memcpy(&bits, &value, sizeof(bits));
    const std::uint32_t rounded = bits + 0x7FFFU + ((bits >> 16U) & 1U);
    return static_cast<std::uint16_t>(rounded >> 16U);
}

void RmsNorm(
    const float * input,
    const float * weight,
    float * output,
    std::size_t count,
    float epsilon) noexcept {
    float sum = 0.0F;
#if defined(__aarch64__)
    float32x4_t accum0 = vdupq_n_f32(0.0F);
    float32x4_t accum1 = vdupq_n_f32(0.0F);
    std::size_t index = 0;
    for (; index + 8 <= count; index += 8) {
        const float32x4_t a = vld1q_f32(input + index);
        const float32x4_t b = vld1q_f32(input + index + 4);
        accum0 = vfmaq_f32(accum0, a, a);
        accum1 = vfmaq_f32(accum1, b, b);
    }
    sum = vaddvq_f32(vaddq_f32(accum0, accum1));
    for (; index < count; ++index) sum += input[index] * input[index];
#else
    for (std::size_t index = 0; index < count; ++index) sum += input[index] * input[index];
#endif
    const float inverse_rms = 1.0F / std::sqrt(sum / static_cast<float>(count) + epsilon);
    for (std::size_t index = 0; index < count; ++index) {
        output[index] = input[index] * inverse_rms * weight[index];
    }
}

void SiluMultiply(const float * gate, const float * up, float * output, std::size_t count) noexcept {
    std::size_t index = 0;
#if defined(LING3_WITH_VECTOR_EXP)
    static const bool vector_math = std::getenv("LING3_VECTOR_MATH") != nullptr;
    if (vector_math) {
        for (; index + 4 <= count; index += 4) {
            const auto value = vld1q_f32(gate + index);
            const auto divisor = vaddq_f32(vdupq_n_f32(1.0F), _ZGVnN4v_expf(vnegq_f32(value)));
            auto result = vdivq_f32(value, divisor);
            if (up) result = vmulq_f32(result, vld1q_f32(up + index));
            vst1q_f32(output + index, result);
        }
    }
#endif
    for (; index < count; ++index) {
        const float value = gate[index];
        const float result = value / (1.0F + std::exp(-value));
        output[index] = up ? result * up[index] : result;
    }
}

void Silu(const float * input, float * output, std::size_t count) noexcept {
    SiluMultiply(input, nullptr, output, count);
}

void WeightedAccumulate(const float * input, float weight, float * output, std::size_t count) noexcept {
#if defined(__aarch64__)
    const float32x4_t scale = vdupq_n_f32(weight);
    std::size_t index = 0;
    for (; index + 4 <= count; index += 4) {
        const float32x4_t source = vld1q_f32(input + index);
        const float32x4_t destination = vld1q_f32(output + index);
        vst1q_f32(output + index, vfmaq_f32(destination, source, scale));
    }
    for (; index < count; ++index) output[index] += input[index] * weight;
#else
    for (std::size_t index = 0; index < count; ++index) output[index] += input[index] * weight;
#endif
}

} // namespace ling3