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
|