diff --git a/source/backend/arm82/Arm82Functions.cpp b/source/backend/arm82/Arm82Functions.cpp index cc13f2bd5..853698260 100644 --- a/source/backend/arm82/Arm82Functions.cpp +++ b/source/backend/arm82/Arm82Functions.cpp @@ -86,6 +86,26 @@ void MNNDecayRankOneUpdateFp16(float* S, const float* k, const float* delta, flo // which matters at d_v=128 where holding both out_k and out_q for the // whole row would otherwise exceed the 32 NEON v register budget. // ────────────────────────────────────────────────────────────────────────── +// BUG FIX (bugfix/linear-attn-fp16): The previous implementation used fp16 +// accumulators throughout (float16x8_t + vfmaq_f16 / vmulq_f16 / vsubq_f16). +// For Qwen3.5's linear_attention the state matrix S is recurrent -- errors +// from each timestep's k·k^T rank-1 update compound over prefill (thousands +// of tokens) and decode. fp16 (10-bit mantissa) cannot represent the +// accumulated sum accurately; the state drifts toward a degenerate rank +// structure and the model starts emitting repetitive tokens on decode +// ("edge devices and edge devices", "The output is the output of the model"). +// +// This rewrite keeps fp16 storage (S / k / q / v / out) but promotes ALL +// arithmetic to fp32 via vcvt_f32_f16 / vcvt_f16_f32 in-loop. The chunk +// size drops from 32 to 16 to keep a full column strip of fp32 accumulators +// (4 x float32x4_t = 16 elements) in registers. Compared to the fp16 path +// this costs 2 vcvt insns per 4 fp16 loads and doubles the register +// pressure for accumulators, but eliminates the drift entirely -- and +// scalars decay/beta/kq stay fp32 rather than being narrowed to fp16. +// +// Follows the fp32 MNNFusedGatedDeltaDefault layout in CommonOptFunction.cpp +// but replaces vld1q_f32 with vcvt_f32_f16(vld1_f16(...)) on load and +// vst1_f32 with vst1_f16(vcvt_f16_f32(...)) on store. static void MNNFusedGatedDeltaFp16(float* S_, const float* k_, const float* q_, const float* v_, float* out_, float decay, float beta, float kq, size_t dk, size_t dv) { auto S = reinterpret_cast<__fp16*>(S_); @@ -94,161 +114,76 @@ static void MNNFusedGatedDeltaFp16(float* S_, const float* k_, const float* q_, auto vIn = reinterpret_cast(v_); auto out = reinterpret_cast<__fp16*>(out_); - const __fp16 decayH = static_cast<__fp16>(decay); - const __fp16 betaH = static_cast<__fp16>(beta); - const __fp16 kqH = static_cast<__fp16>(kq); - const float16x8_t vDecay = vdupq_n_f16(decayH); - const float16x8_t vBeta = vdupq_n_f16(betaH); - const float16x8_t vKq = vdupq_n_f16(kqH); + const float32x4_t vDecay = vdupq_n_f32(decay); + const float32x4_t vBeta = vdupq_n_f32(beta); + const float32x4_t vKq = vdupq_n_f32(kq); - const size_t kChunk = 32; + const size_t kChunk = 16; size_t j = 0; for (; j + kChunk <= dv; j += kChunk) { - // ── Pass 1: out_k = S^T @ k, out_q = S^T @ q for this column chunk ── - float16x8_t ok0 = vdupq_n_f16((__fp16)0), ok1 = vdupq_n_f16((__fp16)0), ok2 = vdupq_n_f16((__fp16)0), - ok3 = vdupq_n_f16((__fp16)0); - float16x8_t oq0 = vdupq_n_f16((__fp16)0), oq1 = vdupq_n_f16((__fp16)0), oq2 = vdupq_n_f16((__fp16)0), - oq3 = vdupq_n_f16((__fp16)0); - size_t i = 0; - // Unroll i by 8: load 8 k & q scalars at once, then use fma-by-lane - // to amortize the scalar broadcast across 8 row iterations. - for (; i + 8 <= dk; i += 8) { - float16x8_t kVec = vld1q_f16(k + i); - float16x8_t qVec = vld1q_f16(q + i); -#define LANE_STEP(lane) \ - { \ - const __fp16* row = S + (i + (lane)) * dv + j; \ - float16x8_t s0 = vld1q_f16(row); \ - float16x8_t s1 = vld1q_f16(row + 8); \ - float16x8_t s2 = vld1q_f16(row + 16); \ - float16x8_t s3 = vld1q_f16(row + 24); \ - ok0 = vfmaq_laneq_f16(ok0, s0, kVec, (lane)); \ - ok1 = vfmaq_laneq_f16(ok1, s1, kVec, (lane)); \ - ok2 = vfmaq_laneq_f16(ok2, s2, kVec, (lane)); \ - ok3 = vfmaq_laneq_f16(ok3, s3, kVec, (lane)); \ - oq0 = vfmaq_laneq_f16(oq0, s0, qVec, (lane)); \ - oq1 = vfmaq_laneq_f16(oq1, s1, qVec, (lane)); \ - oq2 = vfmaq_laneq_f16(oq2, s2, qVec, (lane)); \ - oq3 = vfmaq_laneq_f16(oq3, s3, qVec, (lane)); \ - } - LANE_STEP(0); - LANE_STEP(1); - LANE_STEP(2); - LANE_STEP(3); - LANE_STEP(4); - LANE_STEP(5); - LANE_STEP(6); - LANE_STEP(7); -#undef LANE_STEP - } - // Tail rows (dk % 8) — fall back to the broadcast form. - for (; i < dk; ++i) { + // ── Pass 1: out_k = S^T @ k, out_q = S^T @ q for this column chunk (fp32 accum) ── + float32x4_t ok0 = vdupq_n_f32(0.0f), ok1 = vdupq_n_f32(0.0f); + float32x4_t ok2 = vdupq_n_f32(0.0f), ok3 = vdupq_n_f32(0.0f); + float32x4_t oq0 = vdupq_n_f32(0.0f), oq1 = vdupq_n_f32(0.0f); + float32x4_t oq2 = vdupq_n_f32(0.0f), oq3 = vdupq_n_f32(0.0f); + for (size_t i = 0; i < dk; ++i) { + const float ki = static_cast(k[i]); + const float qi = static_cast(q[i]); const __fp16* row = S + i * dv + j; - float16x8_t s0 = vld1q_f16(row); - float16x8_t s1 = vld1q_f16(row + 8); - float16x8_t s2 = vld1q_f16(row + 16); - float16x8_t s3 = vld1q_f16(row + 24); - __fp16 ki = k[i]; - __fp16 qi = q[i]; - ok0 = vfmaq_n_f16(ok0, s0, ki); - ok1 = vfmaq_n_f16(ok1, s1, ki); - ok2 = vfmaq_n_f16(ok2, s2, ki); - ok3 = vfmaq_n_f16(ok3, s3, ki); - oq0 = vfmaq_n_f16(oq0, s0, qi); - oq1 = vfmaq_n_f16(oq1, s1, qi); - oq2 = vfmaq_n_f16(oq2, s2, qi); - oq3 = vfmaq_n_f16(oq3, s3, qi); - } - - // ── Inline analytic correction (regs only) ── - float16x8_t v0 = vld1q_f16(vIn + j); - float16x8_t v1 = vld1q_f16(vIn + j + 8); - float16x8_t v2 = vld1q_f16(vIn + j + 16); - float16x8_t v3 = vld1q_f16(vIn + j + 24); + float32x4_t s0 = vcvt_f32_f16(vld1_f16(row)); + float32x4_t s1 = vcvt_f32_f16(vld1_f16(row + 4)); + float32x4_t s2 = vcvt_f32_f16(vld1_f16(row + 8)); + float32x4_t s3 = vcvt_f32_f16(vld1_f16(row + 12)); + ok0 = vfmaq_n_f32(ok0, s0, ki); + ok1 = vfmaq_n_f32(ok1, s1, ki); + ok2 = vfmaq_n_f32(ok2, s2, ki); + ok3 = vfmaq_n_f32(ok3, s3, ki); + oq0 = vfmaq_n_f32(oq0, s0, qi); + oq1 = vfmaq_n_f32(oq1, s1, qi); + oq2 = vfmaq_n_f32(oq2, s2, qi); + oq3 = vfmaq_n_f32(oq3, s3, qi); + } + + // ── Inline analytic correction (fp32 regs) ── + float32x4_t v0 = vcvt_f32_f16(vld1_f16(vIn + j)); + float32x4_t v1 = vcvt_f32_f16(vld1_f16(vIn + j + 4)); + float32x4_t v2 = vcvt_f32_f16(vld1_f16(vIn + j + 8)); + float32x4_t v3 = vcvt_f32_f16(vld1_f16(vIn + j + 12)); // delta = beta * (v - decay * out_k) - float16x8_t d0 = vmulq_f16(vBeta, vsubq_f16(v0, vmulq_f16(vDecay, ok0))); - float16x8_t d1 = vmulq_f16(vBeta, vsubq_f16(v1, vmulq_f16(vDecay, ok1))); - float16x8_t d2 = vmulq_f16(vBeta, vsubq_f16(v2, vmulq_f16(vDecay, ok2))); - float16x8_t d3 = vmulq_f16(vBeta, vsubq_f16(v3, vmulq_f16(vDecay, ok3))); + float32x4_t d0 = vmulq_f32(vBeta, vsubq_f32(v0, vmulq_f32(vDecay, ok0))); + float32x4_t d1 = vmulq_f32(vBeta, vsubq_f32(v1, vmulq_f32(vDecay, ok1))); + float32x4_t d2 = vmulq_f32(vBeta, vsubq_f32(v2, vmulq_f32(vDecay, ok2))); + float32x4_t d3 = vmulq_f32(vBeta, vsubq_f32(v3, vmulq_f32(vDecay, ok3))); // out = decay * out_q + kq * delta - float16x8_t o0 = vfmaq_f16(vmulq_f16(vDecay, oq0), vKq, d0); - float16x8_t o1 = vfmaq_f16(vmulq_f16(vDecay, oq1), vKq, d1); - float16x8_t o2 = vfmaq_f16(vmulq_f16(vDecay, oq2), vKq, d2); - float16x8_t o3 = vfmaq_f16(vmulq_f16(vDecay, oq3), vKq, d3); - vst1q_f16(out + j, o0); - vst1q_f16(out + j + 8, o1); - vst1q_f16(out + j + 16, o2); - vst1q_f16(out + j + 24, o3); - - // ── Pass 2: S = decay * S + k ⊗ delta (delta still in regs) ── - size_t i2 = 0; - for (; i2 + 8 <= dk; i2 += 8) { - float16x8_t kVec = vld1q_f16(k + i2); -#define ROW_UPDATE(lane) \ - { \ - __fp16* row = S + (i2 + (lane)) * dv + j; \ - float16x8_t s0 = vld1q_f16(row); \ - float16x8_t s1 = vld1q_f16(row + 8); \ - float16x8_t s2 = vld1q_f16(row + 16); \ - float16x8_t s3 = vld1q_f16(row + 24); \ - float16x8_t r0 = vfmaq_laneq_f16(vmulq_f16(vDecay, s0), d0, kVec, (lane)); \ - float16x8_t r1 = vfmaq_laneq_f16(vmulq_f16(vDecay, s1), d1, kVec, (lane)); \ - float16x8_t r2 = vfmaq_laneq_f16(vmulq_f16(vDecay, s2), d2, kVec, (lane)); \ - float16x8_t r3 = vfmaq_laneq_f16(vmulq_f16(vDecay, s3), d3, kVec, (lane)); \ - vst1q_f16(row, r0); \ - vst1q_f16(row + 8, r1); \ - vst1q_f16(row + 16, r2); \ - vst1q_f16(row + 24, r3); \ - } - ROW_UPDATE(0); - ROW_UPDATE(1); - ROW_UPDATE(2); - ROW_UPDATE(3); - ROW_UPDATE(4); - ROW_UPDATE(5); - ROW_UPDATE(6); - ROW_UPDATE(7); -#undef ROW_UPDATE - } - for (; i2 < dk; ++i2) { + float32x4_t o0 = vfmaq_f32(vmulq_f32(vDecay, oq0), vKq, d0); + float32x4_t o1 = vfmaq_f32(vmulq_f32(vDecay, oq1), vKq, d1); + float32x4_t o2 = vfmaq_f32(vmulq_f32(vDecay, oq2), vKq, d2); + float32x4_t o3 = vfmaq_f32(vmulq_f32(vDecay, oq3), vKq, d3); + vst1_f16(out + j, vcvt_f16_f32(o0)); + vst1_f16(out + j + 4, vcvt_f16_f32(o1)); + vst1_f16(out + j + 8, vcvt_f16_f32(o2)); + vst1_f16(out + j + 12, vcvt_f16_f32(o3)); + + // ── Pass 2: S = decay * S + k ⊗ delta (delta d0..d3 still in fp32 regs) ── + for (size_t i2 = 0; i2 < dk; ++i2) { + const float ki2 = static_cast(k[i2]); __fp16* row = S + i2 * dv + j; - float16x8_t s0 = vld1q_f16(row); - float16x8_t s1 = vld1q_f16(row + 8); - float16x8_t s2 = vld1q_f16(row + 16); - float16x8_t s3 = vld1q_f16(row + 24); - __fp16 ki = k[i2]; - float16x8_t r0 = vfmaq_n_f16(vmulq_f16(vDecay, s0), d0, ki); - float16x8_t r1 = vfmaq_n_f16(vmulq_f16(vDecay, s1), d1, ki); - float16x8_t r2 = vfmaq_n_f16(vmulq_f16(vDecay, s2), d2, ki); - float16x8_t r3 = vfmaq_n_f16(vmulq_f16(vDecay, s3), d3, ki); - vst1q_f16(row, r0); - vst1q_f16(row + 8, r1); - vst1q_f16(row + 16, r2); - vst1q_f16(row + 24, r3); - } - } - - // ── Tail (chunks of 8) ── - for (; j + 8 <= dv; j += 8) { - float16x8_t ok = vdupq_n_f16((__fp16)0); - float16x8_t oq = vdupq_n_f16((__fp16)0); - for (size_t i = 0; i < dk; ++i) { - float16x8_t s = vld1q_f16(S + i * dv + j); - ok = vfmaq_n_f16(ok, s, k[i]); - oq = vfmaq_n_f16(oq, s, q[i]); - } - float16x8_t vv = vld1q_f16(vIn + j); - float16x8_t d = vmulq_f16(vBeta, vsubq_f16(vv, vmulq_f16(vDecay, ok))); - float16x8_t o = vfmaq_f16(vmulq_f16(vDecay, oq), vKq, d); - vst1q_f16(out + j, o); - for (size_t i = 0; i < dk; ++i) { - float16x8_t s = vld1q_f16(S + i * dv + j); - float16x8_t r = vfmaq_n_f16(vmulq_f16(vDecay, s), d, k[i]); - vst1q_f16(S + i * dv + j, r); - } - } - - // ── Scalar tail (defensive; d_v < multiple of 8 not used in current models) ── + float32x4_t s0 = vcvt_f32_f16(vld1_f16(row)); + float32x4_t s1 = vcvt_f32_f16(vld1_f16(row + 4)); + float32x4_t s2 = vcvt_f32_f16(vld1_f16(row + 8)); + float32x4_t s3 = vcvt_f32_f16(vld1_f16(row + 12)); + float32x4_t r0 = vfmaq_n_f32(vmulq_f32(vDecay, s0), d0, ki2); + float32x4_t r1 = vfmaq_n_f32(vmulq_f32(vDecay, s1), d1, ki2); + float32x4_t r2 = vfmaq_n_f32(vmulq_f32(vDecay, s2), d2, ki2); + float32x4_t r3 = vfmaq_n_f32(vmulq_f32(vDecay, s3), d3, ki2); + vst1_f16(row, vcvt_f16_f32(r0)); + vst1_f16(row + 4, vcvt_f16_f32(r1)); + vst1_f16(row + 8, vcvt_f16_f32(r2)); + vst1_f16(row + 12, vcvt_f16_f32(r3)); + } + } + + // ── Scalar tail (defensive; d_v < 16 remainder) ── for (; j < dv; ++j) { float ok = 0.0f, oq = 0.0f; for (size_t i = 0; i < dk; ++i) { @@ -264,6 +199,64 @@ static void MNNFusedGatedDeltaFp16(float* S_, const float* k_, const float* q_, } } } + +// BUG FIX (bugfix/linear-attn-fp16): same as MNNFusedGatedDeltaFp16 above but +// for the decode-path helpers. The original asm implementations in +// MNNRankOneUpdateFp16.S accumulate in fp16, which drifts across ~60 decode +// steps and produces repetitive tokens even after the prefill state is +// clean. These C++ replacements override the asm assignments below. +static void MNNDualMatVecFp16_Fp32Accum(const float* S_, const float* k_, const float* q_, float* out_k_, + float* out_q_, size_t dk, size_t dv) { + auto S = reinterpret_cast(S_); + auto k = reinterpret_cast(k_); + auto q = reinterpret_cast(q_); + auto out_k = reinterpret_cast<__fp16*>(out_k_); + auto out_q = reinterpret_cast<__fp16*>(out_q_); + size_t j = 0; + for (; j + 4 <= dv; j += 4) { + float32x4_t ok = vdupq_n_f32(0.0f); + float32x4_t oq = vdupq_n_f32(0.0f); + for (size_t i = 0; i < dk; ++i) { + float32x4_t s = vcvt_f32_f16(vld1_f16(S + i * dv + j)); + ok = vfmaq_n_f32(ok, s, static_cast(k[i])); + oq = vfmaq_n_f32(oq, s, static_cast(q[i])); + } + vst1_f16(out_k + j, vcvt_f16_f32(ok)); + vst1_f16(out_q + j, vcvt_f16_f32(oq)); + } + for (; j < dv; ++j) { + float ok = 0.0f, oq = 0.0f; + for (size_t i = 0; i < dk; ++i) { + float s = static_cast(S[i * dv + j]); + ok += s * static_cast(k[i]); + oq += s * static_cast(q[i]); + } + out_k[j] = static_cast<__fp16>(ok); + out_q[j] = static_cast<__fp16>(oq); + } +} + +static void MNNDecayRankOneUpdateFp16_Fp32Accum(float* S_, const float* k_, const float* delta_, float decay, + size_t dk, size_t dv) { + auto S = reinterpret_cast<__fp16*>(S_); + auto k = reinterpret_cast(k_); + auto delta = reinterpret_cast(delta_); + const float32x4_t vDecay = vdupq_n_f32(decay); + for (size_t i = 0; i < dk; ++i) { + float k_val = static_cast(k[i]); + __fp16* row = S + i * dv; + size_t j = 0; + for (; j + 4 <= dv; j += 4) { + float32x4_t s = vcvt_f32_f16(vld1_f16(row + j)); + float32x4_t d = vcvt_f32_f16(vld1_f16(delta + j)); + float32x4_t r = vfmaq_n_f32(vmulq_f32(vDecay, s), d, k_val); + vst1_f16(row + j, vcvt_f16_f32(r)); + } + for (; j < dv; ++j) { + row[j] = static_cast<__fp16>(decay * static_cast(row[j]) + k_val * static_cast(delta[j])); + } + } +} #endif // __aarch64__ && MNN_USE_NEON namespace MNN { @@ -2908,8 +2901,10 @@ bool Arm82Functions::init() { // LinearAttention fp16 kernels FUNC_PTR_ASSIGN(gInstance->MNNRankOneUpdate, MNNRankOneUpdateFp16); - FUNC_PTR_ASSIGN(gInstance->MNNDualMatVec, MNNDualMatVecFp16); - FUNC_PTR_ASSIGN(gInstance->MNNDecayRankOneUpdate, MNNDecayRankOneUpdateFp16); + // Override the fp16-accumulator asm helpers with fp32-accumulator C++ versions; + // see the bugfix comment above MNNFusedGatedDeltaFp16. + FUNC_PTR_ASSIGN(gInstance->MNNDualMatVec, MNNDualMatVecFp16_Fp32Accum); + FUNC_PTR_ASSIGN(gInstance->MNNDecayRankOneUpdate, MNNDecayRankOneUpdateFp16_Fp32Accum); #if defined(__aarch64__) && defined(MNN_USE_NEON) // Fused kernel uses NEON intrinsics directly (not extern asm), so the // assignment must follow the same guard as the function body above.