diff --git a/internal/ailego/math/Makefile b/internal/ailego/math/Makefile new file mode 100644 index 0000000..887dde3 --- /dev/null +++ b/internal/ailego/math/Makefile @@ -0,0 +1,29 @@ +GO ?= go +GOAT ?= $(GO) tool goat +GOAT_TARGET_AMD64 ?= amd64 +GOAT_TARGET_ARM64 ?= arm64 +GOAT_OUTPUT_DIR ?= .goat/mathutil +MV ?= mv + +.PHONY: generate avx2 avx512 neon clean + +generate: avx2 avx512 neon clean + +avx2: + mkdir -p $(GOAT_OUTPUT_DIR) + $(GOAT) src/fht_avx2.c --target $(GOAT_TARGET_AMD64) -O3 -mavx2 -o $(GOAT_OUTPUT_DIR) + $(MV) $(GOAT_OUTPUT_DIR)/fht_avx2.go $(GOAT_OUTPUT_DIR)/fht_avx2.s . + +avx512: + mkdir -p $(GOAT_OUTPUT_DIR) + $(GOAT) src/fht_avx512.c --target $(GOAT_TARGET_AMD64) -O3 -mavx512f -mavx512dq -o $(GOAT_OUTPUT_DIR) + $(MV) $(GOAT_OUTPUT_DIR)/fht_avx512.go $(GOAT_OUTPUT_DIR)/fht_avx512.s . + +neon: + mkdir -p $(GOAT_OUTPUT_DIR) + $(GOAT) src/fht_neon.c --target $(GOAT_TARGET_ARM64) -O3 -o $(GOAT_OUTPUT_DIR) + $(MV) $(GOAT_OUTPUT_DIR)/fht_neon.go $(GOAT_OUTPUT_DIR)/fht_neon.s . + +clean: + $(RM) src/*.o src/*.s + $(RM) -r .goat diff --git a/internal/ailego/math/fht.go b/internal/ailego/math/fht.go index 5307dea..5313b18 100644 --- a/internal/ailego/math/fht.go +++ b/internal/ailego/math/fht.go @@ -24,16 +24,29 @@ var ( ErrShortSignBits = errors.New("ailego: sign-bit buffer is too short") ) +type fhtUnaryKernel func(data []float32) +type fhtFlipSignsKernel func(signs []byte, data []float32) + +type fhtKernels struct { + flipSigns fhtFlipSignsKernel + kacWalk fhtUnaryKernel + inverseKacWalk fhtUnaryKernel + inPlace fhtUnaryKernel +} + +var activeFHTKernels = fhtKernels{ + flipSigns: fhtFlipSignsScalar, + kacWalk: fhtKacWalkScalar, + inverseKacWalk: fhtInverseKacWalkScalar, + inPlace: fhtInPlaceScalar, +} + // FHTFlipSigns negates elements selected by the little-endian bits in signs. func FHTFlipSigns(signs []byte, data []float32) error { if len(signs) < (len(data)+7)/8 { return ErrShortSignBits } - for index := range data { - if signs[index/8]&(1< +#include + +void xvec_avx2_fht_flip_signs(uint8_t *signs, float *data, int64_t size) { + int64_t simd_end = size & ~31LL; + for (int64_t index = 0; index < simd_end; index += 32) { + uint32_t bits; + __builtin_memcpy(&bits, signs + index / 8, sizeof(bits)); + for (int64_t block = 0; block < 4; block++) { + uint64_t byte = (bits >> (block * 8)) & 0xff; + volatile uint64_t mask0 = ((byte & 0x01) << 31) | ((byte & 0x02) << 62); + volatile uint64_t mask1 = ((byte & 0x04) << 29) | ((byte & 0x08) << 60); + volatile uint64_t mask2 = ((byte & 0x10) << 27) | ((byte & 0x20) << 58); + volatile uint64_t mask3 = ((byte & 0x40) << 25) | ((byte & 0x80) << 56); + __m256i mask = _mm256_set_epi64x(mask3, mask2, mask1, mask0); + __m256 values = _mm256_loadu_ps(data + index + block * 8); + values = _mm256_xor_ps(values, _mm256_castsi256_ps(mask)); + _mm256_storeu_ps(data + index + block * 8, values); + } + } + for (int64_t index = simd_end; index < size; index++) { + if (signs[index / 8] & (1u << (index % 8))) { + data[index] = -data[index]; + } + } +} + +void xvec_avx2_fht_kac_walk(float *data, int64_t size) { + int64_t half = size / 2; + int64_t base = size % 2; + int64_t offset = base + half; + int64_t simd_end = half & ~7LL; + for (int64_t index = 0; index < simd_end; index += 8) { + __m256 left = _mm256_loadu_ps(data + index); + __m256 right = _mm256_loadu_ps(data + index + offset); + _mm256_storeu_ps(data + index, _mm256_add_ps(left, right)); + _mm256_storeu_ps(data + index + offset, _mm256_sub_ps(left, right)); + } + for (int64_t index = simd_end; index < half; index++) { + float left = data[index]; + float right = data[index + offset]; + data[index] = left + right; + data[index + offset] = left - right; + } + +} + +void xvec_avx2_fht_inverse_kac_walk(float *data, int64_t size) { + int64_t half = size / 2; + int64_t base = size % 2; + int64_t offset = base + half; + int64_t simd_end = half & ~7LL; + volatile float scale_scalar = 0.5f; + const __m256 scale = _mm256_set1_ps(scale_scalar); + for (int64_t index = 0; index < simd_end; index += 8) { + __m256 left = _mm256_loadu_ps(data + index); + __m256 right = _mm256_loadu_ps(data + index + offset); + _mm256_storeu_ps(data + index, _mm256_mul_ps(_mm256_add_ps(left, right), scale)); + _mm256_storeu_ps(data + index + offset, _mm256_mul_ps(_mm256_sub_ps(left, right), scale)); + } + for (int64_t index = simd_end; index < half; index++) { + float left = data[index]; + float right = data[index + offset]; + data[index] = (left + right) * scale_scalar; + data[index + offset] = (left - right) * scale_scalar; + } +} + +void xvec_avx2_fht_in_place(float *data, int64_t size) { + for (int64_t width = 1; width < size; width <<= 1) { + int64_t step = width << 1; + int64_t simd_end = width & ~7LL; + for (int64_t block = 0; block < size; block += step) { + for (int64_t index = 0; index < simd_end; index += 8) { + __m256 left = _mm256_loadu_ps(data + block + index); + __m256 right = _mm256_loadu_ps(data + block + index + width); + _mm256_storeu_ps(data + block + index, _mm256_add_ps(left, right)); + _mm256_storeu_ps(data + block + index + width, _mm256_sub_ps(left, right)); + } + for (int64_t index = simd_end; index < width; index++) { + float left = data[block + index]; + float right = data[block + index + width]; + data[block + index] = left + right; + data[block + index + width] = left - right; + } + } + } +} diff --git a/internal/ailego/math/src/fht_avx512.c b/internal/ailego/math/src/fht_avx512.c new file mode 100644 index 0000000..694e14c --- /dev/null +++ b/internal/ailego/math/src/fht_avx512.c @@ -0,0 +1,99 @@ +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include + +void xvec_avx512_fht_flip_signs(uint8_t *signs, float *data, int64_t size) { + int64_t simd_end = size & ~63LL; + volatile int32_t sign = (int32_t)0x80000000u; + const __m512 sign_bit = _mm512_castsi512_ps(_mm512_set1_epi32(sign)); + for (int64_t index = 0; index < simd_end; index += 64) { + uint64_t bits; + __builtin_memcpy(&bits, signs + index / 8, sizeof(bits)); + for (int64_t block = 0; block < 4; block++) { + __mmask16 mask = (__mmask16)(bits >> (block * 16)); + __m512 values = _mm512_loadu_ps(data + index + block * 16); + values = _mm512_mask_xor_ps(values, mask, values, sign_bit); + _mm512_storeu_ps(data + index + block * 16, values); + } + } + for (int64_t index = simd_end; index < size; index++) { + if (signs[index / 8] & (1u << (index % 8))) { + data[index] = -data[index]; + } + } +} + +void xvec_avx512_fht_kac_walk(float *data, int64_t size) { + int64_t half = size / 2; + int64_t base = size % 2; + int64_t offset = base + half; + int64_t simd_end = half & ~15LL; + for (int64_t index = 0; index < simd_end; index += 16) { + __m512 left = _mm512_loadu_ps(data + index); + __m512 right = _mm512_loadu_ps(data + index + offset); + _mm512_storeu_ps(data + index, _mm512_add_ps(left, right)); + _mm512_storeu_ps(data + index + offset, _mm512_sub_ps(left, right)); + } + for (int64_t index = simd_end; index < half; index++) { + float left = data[index]; + float right = data[index + offset]; + data[index] = left + right; + data[index + offset] = left - right; + } + +} + +void xvec_avx512_fht_inverse_kac_walk(float *data, int64_t size) { + int64_t half = size / 2; + int64_t base = size % 2; + int64_t offset = base + half; + int64_t simd_end = half & ~15LL; + volatile float scale_scalar = 0.5f; + const __m512 scale = _mm512_set1_ps(scale_scalar); + for (int64_t index = 0; index < simd_end; index += 16) { + __m512 left = _mm512_loadu_ps(data + index); + __m512 right = _mm512_loadu_ps(data + index + offset); + _mm512_storeu_ps(data + index, _mm512_mul_ps(_mm512_add_ps(left, right), scale)); + _mm512_storeu_ps(data + index + offset, _mm512_mul_ps(_mm512_sub_ps(left, right), scale)); + } + for (int64_t index = simd_end; index < half; index++) { + float left = data[index]; + float right = data[index + offset]; + data[index] = (left + right) * scale_scalar; + data[index + offset] = (left - right) * scale_scalar; + } +} + +void xvec_avx512_fht_in_place(float *data, int64_t size) { + for (int64_t width = 1; width < size; width <<= 1) { + int64_t step = width << 1; + int64_t simd_end = width & ~15LL; + for (int64_t block = 0; block < size; block += step) { + for (int64_t index = 0; index < simd_end; index += 16) { + __m512 left = _mm512_loadu_ps(data + block + index); + __m512 right = _mm512_loadu_ps(data + block + index + width); + _mm512_storeu_ps(data + block + index, _mm512_add_ps(left, right)); + _mm512_storeu_ps(data + block + index + width, _mm512_sub_ps(left, right)); + } + for (int64_t index = simd_end; index < width; index++) { + float left = data[block + index]; + float right = data[block + index + width]; + data[block + index] = left + right; + data[block + index + width] = left - right; + } + } + } +} diff --git a/internal/ailego/math/src/fht_neon.c b/internal/ailego/math/src/fht_neon.c new file mode 100644 index 0000000..7112dba --- /dev/null +++ b/internal/ailego/math/src/fht_neon.c @@ -0,0 +1,101 @@ +// Copyright 2026-present the xvec project +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include + +void xvec_neon_fht_flip_signs(uint8_t *signs, float *data, int64_t size) { + int64_t simd_end = size & ~3LL; + for (int64_t index = 0; index < simd_end; index += 4) { + uint16_t bits = signs[index / 8]; + if (index / 8 + 1 < (size + 7) / 8) { + bits |= (uint16_t)signs[index / 8 + 1] << 8; + } + bits >>= index % 8; + volatile uint64_t mask0 = ((uint64_t)(bits & 0x01) << 31) | + ((uint64_t)(bits & 0x02) << 62); + volatile uint64_t mask1 = ((uint64_t)(bits & 0x04) << 29) | + ((uint64_t)(bits & 0x08) << 60); + uint64x2_t mask64 = vcombine_u64(vcreate_u64(mask0), vcreate_u64(mask1)); + uint32x4_t values = vreinterpretq_u32_f32(vld1q_f32(data + index)); + vst1q_f32(data + index, vreinterpretq_f32_u32(veorq_u32(values, vreinterpretq_u32_u64(mask64)))); + } + for (int64_t index = simd_end; index < size; index++) { + if (signs[index / 8] & (1u << (index % 8))) { + data[index] = -data[index]; + } + } +} + +void xvec_neon_fht_kac_walk(float *data, int64_t size) { + int64_t half = size / 2; + int64_t base = size % 2; + int64_t offset = base + half; + int64_t simd_end = half & ~3LL; + for (int64_t index = 0; index < simd_end; index += 4) { + float32x4_t left = vld1q_f32(data + index); + float32x4_t right = vld1q_f32(data + index + offset); + vst1q_f32(data + index, vaddq_f32(left, right)); + vst1q_f32(data + index + offset, vsubq_f32(left, right)); + } + for (int64_t index = simd_end; index < half; index++) { + float left = data[index]; + float right = data[index + offset]; + data[index] = left + right; + data[index + offset] = left - right; + } + +} + +void xvec_neon_fht_inverse_kac_walk(float *data, int64_t size) { + int64_t half = size / 2; + int64_t base = size % 2; + int64_t offset = base + half; + + int64_t simd_end = half & ~3LL; + const float32x4_t scale = vdupq_n_f32(0.5f); + for (int64_t index = 0; index < simd_end; index += 4) { + float32x4_t left = vld1q_f32(data + index); + float32x4_t right = vld1q_f32(data + index + offset); + vst1q_f32(data + index, vmulq_f32(vaddq_f32(left, right), scale)); + vst1q_f32(data + index + offset, vmulq_f32(vsubq_f32(left, right), scale)); + } + for (int64_t index = simd_end; index < half; index++) { + float left = data[index]; + float right = data[index + offset]; + data[index] = (left + right) * 0.5f; + data[index + offset] = (left - right) * 0.5f; + } +} + +void xvec_neon_fht_in_place(float *data, int64_t size) { + for (int64_t width = 1; width < size; width <<= 1) { + int64_t step = width << 1; + int64_t simd_end = width & ~3LL; + for (int64_t block = 0; block < size; block += step) { + for (int64_t index = 0; index < simd_end; index += 4) { + float32x4_t left = vld1q_f32(data + block + index); + float32x4_t right = vld1q_f32(data + block + index + width); + vst1q_f32(data + block + index, vaddq_f32(left, right)); + vst1q_f32(data + block + index + width, vsubq_f32(left, right)); + } + for (int64_t index = simd_end; index < width; index++) { + float left = data[block + index]; + float right = data[block + index + width]; + data[block + index] = left + right; + data[block + index + width] = left - right; + } + } + } +}