Commit a122b3c3b6 for ffmpeg

commit a122b3c3b66f7efd802385b34fba3cb1329d69be
Author: Lynne <dev@lynne.ee>
Date:   Thu Jun 4 17:50:57 2026 +0900

    lavu/tx: add AArch64 NEON inverse MDCT

    Same as the x86 version.

diff --git a/libavutil/aarch64/tx_float_init.c b/libavutil/aarch64/tx_float_init.c
index a049562609..70be9e8483 100644
--- a/libavutil/aarch64/tx_float_init.c
+++ b/libavutil/aarch64/tx_float_init.c
@@ -37,6 +37,7 @@ TX_DECL_FN(fft_sr,    neon)
 TX_DECL_FN(fft_sr_ns, neon)
 TX_DECL_FN(fft_pfa_15xM, neon)
 TX_DECL_FN(fft_pfa_15xM_ns, neon)
+TX_DECL_FN(mdct_inv, neon)

 static av_cold int neon_init(AVTXContext *s, const FFTXCodelet *cd,
                              uint64_t flags, FFTXCodeletOptions *opts,
@@ -121,6 +122,49 @@ static av_cold int fft_pfa_init(AVTXContext *s, const FFTXCodelet *cd,
     return 0;
 }

+/* Inverse MDCT: a pre-rotation, an in-place len/2 complex FFT, and a
+ * post-rotation. Mirrors the generic ff_tx_mdct_init / x86 m_inv_init, but the
+ * subtransform is called with a normal blr (no FF_TX_ASM_CALL on aarch64). */
+static av_cold int mdct_inv_init(AVTXContext *s, const FFTXCodelet *cd,
+                                 uint64_t flags, FFTXCodeletOptions *opts,
+                                 int len, int inv, const void *scale)
+{
+    int ret;
+    FFTXCodeletOptions sub_opts = { .map_dir = FF_TX_MAP_GATHER };
+
+    /* The pre-rotation processes two output complex at a time, so len/2 must
+     * be even.  Real codecs always satisfy this; bail out otherwise so the
+     * generic C MDCT is used. */
+    if (len & 3)
+        return AVERROR(ENOSYS);
+
+    s->scale_d = *((const float *)scale);
+    s->scale_f = s->scale_d;
+
+    flags &= ~FF_TX_OUT_OF_PLACE; /* The subtransform is in-place */
+    flags |=  AV_TX_INPLACE;
+    flags |=  FF_TX_PRESHUFFLE;   /* This function handles the permute step */
+
+    if ((ret = ff_tx_init_subtx(s, TX_TYPE(FFT), flags, &sub_opts, len >> 1,
+                                inv, scale)))
+        return ret;
+
+    s->map = av_malloc((len >> 1)*sizeof(*s->map));
+    if (!s->map)
+        return AVERROR(ENOMEM);
+
+    memcpy(s->map, s->sub->map, (len >> 1)*sizeof(*s->map));
+
+    if ((ret = ff_tx_mdct_gen_exp_float(s, s->map)))
+        return ret;
+
+    /* Pre-double the map indices (saves a shift in the hot path). */
+    for (int i = 0; i < (len >> 1); i++)
+        s->map[i] <<= 1;
+
+    return 0;
+}
+
 const FFTXCodelet * const ff_tx_codelet_list_float_aarch64[] = {
     TX_DEF(fft2,      FFT,  2,  2, 2, 0, 128, NULL,      neon, NEON, AV_TX_INPLACE, 0),
     TX_DEF(fft2,      FFT,  2,  2, 2, 0, 192, neon_init, neon, NEON, AV_TX_INPLACE | FF_TX_PRESHUFFLE, 0),
@@ -142,5 +186,7 @@ const FFTXCodelet * const ff_tx_codelet_list_float_aarch64[] = {
     TX_DEF(fft_pfa_15xM,    FFT, 60, TX_LEN_UNLIMITED, 15, 2, 128, fft_pfa_init, neon, NEON, AV_TX_INPLACE, 0),
     TX_DEF(fft_pfa_15xM_ns, FFT, 60, TX_LEN_UNLIMITED, 15, 2, 192, fft_pfa_init, neon, NEON, AV_TX_INPLACE | FF_TX_PRESHUFFLE, 0),

+    TX_DEF(mdct_inv, MDCT, 16, TX_LEN_UNLIMITED, 2, TX_FACTOR_ANY, 256, mdct_inv_init, neon, NEON, FF_TX_INVERSE_ONLY, 0),
+
     NULL,
 };
diff --git a/libavutil/aarch64/tx_float_neon.S b/libavutil/aarch64/tx_float_neon.S
index 2db0e95740..96aadbab6c 100644
--- a/libavutil/aarch64/tx_float_neon.S
+++ b/libavutil/aarch64/tx_float_neon.S
@@ -770,6 +770,128 @@ endfunc
 PFA_15_FN float,    0
 PFA_15_FN ns_float, 1

+// Inverse MDCT (see ff_tx_mdct_inv): pre-rotation, in-place N/2pt FFT,
+// post-rotation. out is the N/2-complex buffer z, in the N-real input.
+// Where ld2 can split the re/im planes the multiplies are planar and fused
+// (fmul + fmls/fmla); the gather and the odd tail use interleaved pairs
+// (trn1/trn2 broadcast, rev64 swap, the v31 fmla).
+// AVTXContext offsets: len=0, map=8 (doubled in init), exp=16, sub=32, fn=40.
+function ff_tx_mdct_inv_float_neon, export=1
+        AARCH64_SIGN_LINK_REGISTER
+        stp             x29, x30, [sp, #-48]!
+        mov             x29, sp
+        stp             x19, x20, [sp, #16]
+        stp             x21, x22, [sp, #32]
+
+        mov             x19, x1                     // z
+        mov             x20, x0                     // ctx
+        ldr             w21, [x0, #0]               // N
+        ldr             x22, [x0, #16]              // exp
+        ldr             x4,  [x0, #8]               // map (doubled)
+
+        LOAD_SUBADD
+
+        sub             x5,  x21, #1
+        madd            x14, x5,  x3,  x2           // in2 = in + (N-1)*stride
+        lsr             w7,  w21, #1                // len2
+
+        // pre-rotation via gather: z[i] = (in2[-k], in1[k])*exp[i], 2 cx/iter
+        mov             x13, x19
+        mov             x15, x22
+1:
+        ldp             w5,  w6,  [x4], #8          // k0, k1
+        madd            x8,  x5,  x3,  x2           // &in[k0]      (im0)
+        msub            x9,  x5,  x3,  x14          // &in2[-k0]    (re0)
+        madd            x10, x6,  x3,  x2           // &in[k1]      (im1)
+        msub            x11, x6,  x3,  x14          // &in2[-k1]    (re1)
+        ldr             s0,  [x9]
+        ld1             { v0.s }[1], [x8]
+        ld1             { v0.s }[2], [x11]
+        ld1             { v0.s }[3], [x10]          // (re0, im0, re1, im1) = tmp
+        ld1             { v1.4s }, [x15], #16        // exp[i,i+1]
+        trn1            v2.4s, v0.4s, v0.4s         // re dup
+        trn2            v3.4s, v0.4s, v0.4s         // im dup
+        rev64           v4.4s, v1.4s                // exp swapped
+        fmul            v5.4s, v2.4s, v1.4s
+        fmul            v6.4s, v3.4s, v4.4s
+        fmla            v5.4s, v6.4s, v31.4s        // z = tmp*exp
+        st1             { v5.4s }, [x13], #16
+        subs            w7,  w7,  #2
+        b.gt            1b
+
+        ldr             x5,  [x20, #40]             // fn[0]
+        ldr             x0,  [x20, #32]             // sub[0]
+        mov             x1,  x19
+        mov             x2,  x19
+        mov             x3,  #8
+        blr             x5                          // N/2pt FFT, in-place
+
+        // post-rotation over symmetric pairs (i0 = len4+i, i1 = len4-1-i),
+        // 2 pairs/iter, planes ld2-split: r = swap(z)*swap(e), the im parts
+        // crossing partners (z[i0] = (r_i0.re, r_i1.im) and vice versa), so
+        // each side stores its re plane zipped with the other's reversed ims
+        add             x15, x22, x21, lsl #2       // exp_post = exp + len2
+        add             x8,  x19, x21, lsl #1       // p_i0 = z + len4
+        sub             x9,  x8,  #16               // p_i1 = z + len4 - 2
+        add             x10, x15, x21, lsl #1       // e_i0 = exp_post + len4
+        sub             x11, x10, #16               // e_i1 = exp_post + len4 - 2
+        lsr             w7,  w21, #2                // len4
+        lsr             w6,  w7,  #1                // pairs of pairs
+        cbz             w6,  8f                     // N == 4: lone middle pair
+2:
+        ld2             { v0.2s, v1.2s }, [x8]       // z asc (i0, i0+1): re, im planes
+        ld2             { v2.2s, v3.2s }, [x9]       // z desc (i1-1, i1)
+        ld2             { v4.2s, v5.2s }, [x10], #16 // e asc: er, ei
+        ld2             { v6.2s, v7.2s }, [x11]      // e desc
+        sub             x11, x11, #16
+        fmul            v16.2s, v1.2s, v5.2s        // asc:  z.im*e.im
+        fmul            v18.2s, v1.2s, v4.2s        // asc:  z.im*e.re
+        fmul            v20.2s, v3.2s, v7.2s        // desc: z.im*e.im
+        fmul            v22.2s, v3.2s, v6.2s        // desc: z.im*e.re
+        fmls            v16.2s, v0.2s, v4.2s        // re plane (r_i0,   r_i0+1)
+        fmla            v18.2s, v0.2s, v5.2s        // im plane (r_i0,   r_i0+1)
+        fmls            v20.2s, v2.2s, v6.2s        // re plane (r_i1-1, r_i1)
+        fmla            v22.2s, v2.2s, v7.2s        // im plane (r_i1-1, r_i1)
+        rev64           v22.2s, v22.2s              // (r_i1.im,   r_i1-1.im)
+        rev64           v18.2s, v18.2s              // (r_i0+1.im, r_i0.im)
+        zip1            v0.4s, v16.4s, v22.4s       // (z[i0], z[i0+1])
+        zip1            v2.4s, v20.4s, v18.4s       // (z[i1-1], z[i1])
+        st1             { v0.4s }, [x8], #16
+        st1             { v2.4s }, [x9]
+        sub             x9,  x9,  #16
+        subs            w6,  w6,  #1
+        b.gt            2b
+8:
+        // odd len4 (N % 8 == 4): one leftover pair, (i0, i1) = (len2-1, 0).
+        // x8/x10 already point at it; x9/x11 sit one complex below their slot.
+        tbz             w7,  #0, 9f
+        LOAD_SUBADD     // v31 (clobbered by the FFT)
+        add             x9,  x9,  #8
+        add             x11, x11, #8
+        ldr             d0,  [x8]                   // z[i0]
+        ld1             { v0.d }[1], [x9]           // v0 = (z[i0], z[i1])
+        ldr             d1,  [x10]                  // exp[i0]
+        ld1             { v1.d }[1], [x11]          // v1 = (exp[i0], exp[i1])
+        rev64           v2.4s, v0.4s                // swap(z)   = a
+        rev64           v3.4s, v1.4s                // swap(exp) = b
+        trn1            v4.4s, v2.4s, v2.4s         // a.re dup
+        trn2            v5.4s, v2.4s, v2.4s         // a.im dup
+        fmul            v6.4s, v4.4s, v3.4s         // a.re*b
+        fmul            v7.4s, v5.4s, v1.4s         // a.im*b_swap (b_swap = orig exp)
+        fmla            v6.4s, v7.4s, v31.4s        // v6 = (r0.re, r0.im, r1.re, r1.im)
+        mov             v16.16b, v6.16b
+        ins             v16.s[1], v6.s[3]           // (r0.re, r1.im, r1.re, r1.im)
+        ins             v16.s[3], v6.s[1]           // (r0.re, r1.im, r1.re, r0.im)
+        st1             { v16.d }[0], [x8]          // z[i0] = (r0.re, r1.im)
+        st1             { v16.d }[1], [x9]          // z[i1] = (r1.re, r0.im)
+9:
+        ldp             x19, x20, [sp, #16]
+        ldp             x21, x22, [sp, #32]
+        ldp             x29, x30, [sp], #48
+        AARCH64_VALIDATE_LINK_REGISTER
+        ret
+endfunc
+
 .macro SETUP_SR_RECOMB len, re, im, dec
         ldr             w5, =(\len - 4*7)
         movrel          \re, X(ff_tx_tab_\len\()_float)