#include "ling3/cpu_kernels.h" #include #include #include #if defined(__aarch64__) #include #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(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(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(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