#include "rivide/pqc/ntt_simd.h"
void rivide_simd_poly_add_reduce(int16_t *r, const int16_t *a, const int16_t *b, int16_t q) {
size_t i;
(void)q;
#if defined(RIVIDE_NTT_AVX2_ENABLED)
for (i = 0; i < 256; i += 16) {
__m256i va = _mm256_loadu_si256((const __m256i *)(const void *)(a + i));
__m256i vb = _mm256_loadu_si256((const __m256i *)(const void *)(b + i));
__m256i vr = _mm256_add_epi16(va, vb);
_mm256_storeu_si256((__m256i *)(void *)(r + i), vr);
}
#elif defined(RIVIDE_NTT_NEON_ENABLED)
for (i = 0; i < 256; i += 8) {
int16x8_t va = vld1q_s16(a + i);
int16x8_t vb = vld1q_s16(b + i);
int16x8_t vr = vaddq_s16(va, vb);
vst1q_s16(r + i, vr);
}
#else
for (i = 0; i < 256; i++) {
r[i] = (int16_t)(a[i] + b[i]);
}
#endif
}
void rivide_simd_poly_sub_reduce(int16_t *r, const int16_t *a, const int16_t *b, int16_t q) {
size_t i;
(void)q;
#if defined(RIVIDE_NTT_AVX2_ENABLED)
for (i = 0; i < 256; i += 16) {
__m256i va = _mm256_loadu_si256((const __m256i *)(const void *)(a + i));
__m256i vb = _mm256_loadu_si256((const __m256i *)(const void *)(b + i));
__m256i vr = _mm256_sub_epi16(va, vb);
_mm256_storeu_si256((__m256i *)(void *)(r + i), vr);
}
#elif defined(RIVIDE_NTT_NEON_ENABLED)
for (i = 0; i < 256; i += 8) {
int16x8_t va = vld1q_s16(a + i);
int16x8_t vb = vld1q_s16(b + i);
int16x8_t vr = vsubq_s16(va, vb);
vst1q_s16(r + i, vr);
}
#else
for (i = 0; i < 256; i++) {
r[i] = (int16_t)(a[i] - b[i]);
}
#endif
}
void rivide_simd_poly_pointwise_montgomery(int16_t *r, const int16_t *a, const int16_t *b,
int16_t q, int32_t qinv) {
size_t i;
(void)qinv;
for (i = 0; i < 256; i++) {
int32_t prod = (int32_t)a[i] * (int32_t)b[i];
int16_t t = (int16_t)((int64_t)prod * 62209);
r[i] = (int16_t)((prod - (int32_t)t * q) >> 16);
}
}