Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
295 changes: 145 additions & 150 deletions source/backend/arm82/Arm82Functions.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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_);
Expand All @@ -94,161 +114,76 @@ static void MNNFusedGatedDeltaFp16(float* S_, const float* k_, const float* q_,
auto vIn = reinterpret_cast<const __fp16*>(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<float>(k[i]);
const float qi = static_cast<float>(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<float>(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) {
Expand All @@ -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<const __fp16*>(S_);
auto k = reinterpret_cast<const __fp16*>(k_);
auto q = reinterpret_cast<const __fp16*>(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<float>(k[i]));
oq = vfmaq_n_f32(oq, s, static_cast<float>(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<float>(S[i * dv + j]);
ok += s * static_cast<float>(k[i]);
oq += s * static_cast<float>(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<const __fp16*>(k_);
auto delta = reinterpret_cast<const __fp16*>(delta_);
const float32x4_t vDecay = vdupq_n_f32(decay);
for (size_t i = 0; i < dk; ++i) {
float k_val = static_cast<float>(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<float>(row[j]) + k_val * static_cast<float>(delta[j]));
}
}
}
#endif // __aarch64__ && MNN_USE_NEON

namespace MNN {
Expand Down Expand Up @@ -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.
Expand Down
Loading