E5 NEON-vectorize C-collect (fp32/int32/int16)

Replace scalar per-element accumulate/scale in the result-collect loops
with NEON 4-wide FMA helpers (rknpu_axpy_f32/s32/s16). Helps the decode
path (M=1, single-threaded collect) and the prefill CPU tail.

Experiment (ornith-1.0-9b-Q8_0, -t 4, ulimit -n 65536):
  tg128  2.93 -> 2.98 t/s  (vectorized decode collect)
  pp300  36.53 -> 36.50    (flat; NPU-bound)
  pp512  42.09 -> 42.10    (flat)
  pp700  42.37 -> 41.52    (within noise)
Output coherent (FMA rounding: not byte-identical bar, per E5 spec).

Co-Authored-By: Claude Haiku 4.5 <noreply@anthropic.com>
This commit is contained in:
mfritsche
2026-07-09 01:52:20 +02:00
parent 3705e84765
commit 05cf2f9dd5
+47 -15
View File
@@ -465,6 +465,50 @@ static std::shared_ptr<rknn_tensor_mem> get_tensor_buffer(
return mem_shared;
}
#include <arm_neon.h>
// NEON collect helpers: dst[i] = src[i]*scale (accumulate=false)
// or dst[i] += src[i]*scale (accumulate=true)
static inline void rknpu_axpy_f32(float* dst, const float* src, int n, float scale, bool accumulate) {
const float32x4_t vs = vdupq_n_f32(scale);
int i = 0;
if (accumulate) {
for (; i + 4 <= n; i += 4)
vst1q_f32(dst + i, vfmaq_f32(vld1q_f32(dst + i), vld1q_f32(src + i), vs));
for (; i < n; ++i) dst[i] += src[i] * scale;
} else {
for (; i + 4 <= n; i += 4)
vst1q_f32(dst + i, vmulq_f32(vld1q_f32(src + i), vs));
for (; i < n; ++i) dst[i] = src[i] * scale;
}
}
static inline void rknpu_axpy_s32(float* dst, const int32_t* src, int n, float scale, bool accumulate) {
const float32x4_t vs = vdupq_n_f32(scale);
int i = 0;
if (accumulate) {
for (; i + 4 <= n; i += 4)
vst1q_f32(dst + i, vfmaq_f32(vld1q_f32(dst + i), vcvtq_f32_s32(vld1q_s32(src + i)), vs));
for (; i < n; ++i) dst[i] += (float)src[i] * scale;
} else {
for (; i + 4 <= n; i += 4)
vst1q_f32(dst + i, vmulq_f32(vcvtq_f32_s32(vld1q_s32(src + i)), vs));
for (; i < n; ++i) dst[i] = (float)src[i] * scale;
}
}
static inline void rknpu_axpy_s16(float* dst, const int16_t* src, int n, float scale, bool accumulate) {
const float32x4_t vs = vdupq_n_f32(scale);
int i = 0;
if (accumulate) {
for (; i + 4 <= n; i += 4)
vst1q_f32(dst + i, vfmaq_f32(vld1q_f32(dst + i), vcvtq_f32_s32(vmovl_s16(vld1_s16(src + i))), vs));
for (; i < n; ++i) dst[i] += (float)src[i] * scale;
} else {
for (; i + 4 <= n; i += 4)
vst1q_f32(dst + i, vmulq_f32(vcvtq_f32_s32(vmovl_s16(vld1_s16(src + i))), vs));
for (; i < n; ++i) dst[i] = (float)src[i] * scale;
}
}
static enum ggml_status ggml_backend_rknpu_graph_compute(ggml_backend_t backend, struct ggml_cgraph* cgraph) {
auto* backend_ctx = (ggml_backend_rknpu_context*)backend->context;
@@ -762,11 +806,7 @@ static enum ggml_status ggml_backend_rknpu_graph_compute(ggml_backend_t backend,
float* dst_ptr = dst_data + (size_t)m * N + N_offset;
float* src_ptr = src_segment_base + (size_t)m * N_segment;
if (single_k_segment) {
for(int n=0; n<N_segment; ++n) dst_ptr[n] = src_ptr[n] * dequant_scale;
} else {
for(int n=0; n<N_segment; ++n) dst_ptr[n] += src_ptr[n] * dequant_scale;
}
rknpu_axpy_f32(dst_ptr, src_ptr, N_segment, dequant_scale, !single_k_segment);
}
break;
}
@@ -781,11 +821,7 @@ static enum ggml_status ggml_backend_rknpu_graph_compute(ggml_backend_t backend,
float* dst_ptr = dst_data + (size_t)m * N + N_offset;
int32_t* src_ptr = (int32_t*)mem_C_segments[idx]->virt_addr + (size_t)m * N_segment;
if (single_k_segment) {
for(int n=0; n<N_segment; ++n) dst_ptr[n] = (float)src_ptr[n] * dequant_scale;
} else {
for(int n=0; n<N_segment; ++n) dst_ptr[n] += (float)src_ptr[n] * dequant_scale;
}
rknpu_axpy_s32(dst_ptr, src_ptr, N_segment, dequant_scale, !single_k_segment);
}
break;
}
@@ -800,11 +836,7 @@ static enum ggml_status ggml_backend_rknpu_graph_compute(ggml_backend_t backend,
float* dst_ptr = dst_data + (size_t)m * N + N_offset;
int16_t* src_ptr = (int16_t*)mem_C_segments[idx]->virt_addr + (size_t)m * N_segment;
if (single_k_segment) {
for(int n=0; n<N_segment; ++n) dst_ptr[n] = (float)src_ptr[n] * dequant_scale;
} else {
for(int n=0; n<N_segment; ++n) dst_ptr[n] += (float)src_ptr[n] * dequant_scale;
}
rknpu_axpy_s16(dst_ptr, src_ptr, N_segment, dequant_scale, !single_k_segment);
}
break;
}