Commit d04c5385ce for ffmpeg

commit d04c5385ce721358da8a72139914ae8366686a0a
Author: Lynne <dev@lynne.ee>
Date:   Sat Sep 26 15:08:11 2026 +0900

    vulkan_ffv1: evaluate quant table inputs 0 and 3 with a subgroup ballot

    The two context inputs that depend on the sample to the left are the
    only ones on the dependency chain between decoding a sample and loading
    the states of the next one. When a quant table is a staircase over the
    8-bit difference, with at most 32 steps of the same size, its value is
    the number of thresholds the difference passes: one threshold per
    invocation, a comparison, a ballot and a popcount, with no memory access.
    The host checks the tables and stores the thresholds next to them in the
    constants buffer, and the decoder requests a subgroup of exactly
    CONTEXT_SIZE invocations when it uses them.

    get_pred is split into the part that depends on the row above and the
    part that depends on the sample to the left.

    Decoding a 6464x4852 16-bit RGB frame with 1024 slices on an RX 6900
    XT, with the bitstream in VRAM, goes from 171.6/143.7/139.6 ms to
    168.9/144.1/140.2 ms with context model 1/0/2: only context model 1
    has an extended lookup, and the current decoder is bound by shared
    memory and barriers. The rewrite of the range coded decoder relies on
    it.

diff --git a/libavcodec/ffv1_vulkan.c b/libavcodec/ffv1_vulkan.c
index 9c34a479d3..9d61f35946 100644
--- a/libavcodec/ffv1_vulkan.c
+++ b/libavcodec/ffv1_vulkan.c
@@ -82,6 +82,42 @@ static void set_rc_state_tab(FFV1Context *f, uint8_t *buf)
     }
 }

+int ff_ffv1_vk_quant_ballot(const FFV1Context *f, FFv1QuantBallot *qb)
+{
+    int ok = 1;
+
+    for (int i = 0; i < MAX_QUANT_TABLES; i++)
+        for (int k = 0; k < 32; k++)
+            qb->thresh[i][k][0] = qb->thresh[i][k][1] = 128;
+    memset(qb->scale_off, 0, sizeof(qb->scale_off));
+
+    for (int i = 0; i < f->quant_table_count; i++) {
+        for (int j = 0; j < 2; j++) {
+            const int16_t *qt = f->quant_tables[i][3*j];
+            int n = 0, scale = 0;
+
+            for (int d = -127; d < 128; d++) {
+                int step = qt[d & 255] - qt[(d - 1) & 255];
+                if (!step)
+                    continue;
+                if (!scale)
+                    scale = step;
+                if (step != scale || n == 32 || (!j && step != 1)) {
+                    ok = 0;
+                    break;
+                }
+                qb->thresh[i][n++][j] = d;
+            }
+
+            if (j)
+                qb->scale_off[i][0] = scale;
+            qb->scale_off[i][1] += qt[128];
+        }
+    }
+
+    return ok;
+}
+
 int ff_ffv1_vk_init_consts(FFVulkanContext *s, FFVkBuffer *vkb, FFV1Context *f)
 {
     int err;
@@ -92,7 +128,8 @@ int ff_ffv1_vk_init_consts(FFVulkanContext *s, FFVkBuffer *vkb, FFV1Context *f)
                      512*sizeof(uint8_t) + /* Rangecoder */
                      MAX_QUANT_TABLES*
                      MAX_CONTEXT_INPUTS*
-                     MAX_QUANT_TABLE_SIZE*sizeof(int32_t);
+                     MAX_QUANT_TABLE_SIZE*sizeof(int32_t) +
+                     sizeof(FFv1QuantBallot);

     RET(ff_vk_create_buf(s, vkb,
                          buf_len,
@@ -113,6 +150,8 @@ int ff_ffv1_vk_init_consts(FFVulkanContext *s, FFVkBuffer *vkb, FFV1Context *f)
             for (int k = 0; k < MAX_QUANT_TABLE_SIZE; k++)
                 quant_tables[i][j][k] = f->quant_tables[i][j][k];

+    ff_ffv1_vk_quant_ballot(f, (FFv1QuantBallot *)(quant_tables + MAX_QUANT_TABLES));
+
     RET(ff_vk_unmap_buffer(s, vkb, 1));

 fail:
diff --git a/libavcodec/ffv1_vulkan.h b/libavcodec/ffv1_vulkan.h
index fd69acf0a1..f4b6f167be 100644
--- a/libavcodec/ffv1_vulkan.h
+++ b/libavcodec/ffv1_vulkan.h
@@ -30,6 +30,21 @@ void ff_ffv1_vk_set_common_sl(AVCodecContext *avctx, FFV1Context *f,

 int ff_ffv1_vk_init_consts(FFVulkanContext *s, FFVkBuffer *vkb, FFV1Context *f);

+/* Context inputs 0 and 3 of the quant tables as a function of the signed
+ * 8-bit difference d, with one threshold per subgroup invocation:
+ * q(d) = q(-128) + scale*#{k : thresh[k] <= d}, where scale is 1 for input 0.
+ * scale_off holds the scale of input 3 and the sum of both q(-128). */
+typedef struct FFv1QuantBallot {
+    int32_t thresh[MAX_QUANT_TABLES][32][2];
+    int32_t scale_off[MAX_QUANT_TABLES][2];
+} FFv1QuantBallot;
+
+/**
+ * Fill in the ballot quantizers of all quant tables.
+ * Returns 1 if all of them can be evaluated with a ballot.
+ */
+int ff_ffv1_vk_quant_ballot(const FFV1Context *f, FFv1QuantBallot *qb);
+
 typedef struct FFv1ShaderParams {
     VkDeviceAddress slice_data;

diff --git a/libavcodec/ffv1enc_vulkan.c b/libavcodec/ffv1enc_vulkan.c
index 5de2c15087..1f0d839ded 100644
--- a/libavcodec/ffv1enc_vulkan.c
+++ b/libavcodec/ffv1enc_vulkan.c
@@ -1390,7 +1390,8 @@ static av_cold int vulkan_encode_ffv1_init(AVCodecContext *avctx)
                                         &fv->consts_buf,
                                         256*sizeof(uint32_t) + 512*sizeof(uint8_t),
                                         MAX_QUANT_TABLES*MAX_CONTEXT_INPUTS*
-                                        MAX_QUANT_TABLE_SIZE*sizeof(int32_t),
+                                        MAX_QUANT_TABLE_SIZE*sizeof(int32_t) +
+                                        sizeof(FFv1QuantBallot),
                                         VK_FORMAT_UNDEFINED));
     RET(ff_vk_shader_update_desc_buffer(&fv->s, &fv->exec_pool.contexts[0],
                                         &fv->enc, 0, 2, 0,
diff --git a/libavcodec/vulkan/ffv1_common.glsl b/libavcodec/vulkan/ffv1_common.glsl
index a3177b0d75..22fef6c969 100644
--- a/libavcodec/vulkan/ffv1_common.glsl
+++ b/libavcodec/vulkan/ffv1_common.glsl
@@ -183,6 +183,8 @@ layout (set = 0, binding = 1, scalar) readonly uniform quant_buf {
     int32_t quant_table[MAX_QUANT_TABLES]
                        [MAX_CONTEXT_INPUTS]
                        [MAX_QUANT_TABLE_SIZE];
+    ivec2 quant_thresh[MAX_QUANT_TABLES][32];
+    ivec2 quant_scale_off[MAX_QUANT_TABLES];
 };

 /* -1, { -1, 0 } */
@@ -206,8 +208,8 @@ shared VTYPE2 linecache;
 #define RGB_LBUF (rgb_linecache - 1)
 #define LADDR(p) (ivec2((p).x, ((p).y & RGB_LBUF)))

-ivec2 get_pred(readonly uimage2D pred, ivec2 sp, ivec2 off,
-               uint comp, int sw, uint8_t quant_table_idx, bool extend_lookup)
+ivec4 get_top(readonly uimage2D pred, ivec2 sp, ivec2 off,
+              uint comp, int sw, bool extend_lookup)
 {
     ivec2 yoff_border1 = expectEXT(off.x == 0, false) ? off + ivec2(1, -1) : off;

@@ -216,38 +218,25 @@ ivec2 get_pred(readonly uimage2D pred, ivec2 sp, ivec2 off,
                          TYPE(imageLoad(pred, sp + LADDR(off + ivec2(0, -1)))[comp]),
                          TYPE(imageLoad(pred, sp + LADDR(off + ivec2(min(1, sw - off.x - 1), -1)))[comp]));

-    /* Normally, we'd need to check if off != ivec2(0, 0) here, since otherwise, we must
-     * return zero. However, ivec2(-1,  0) + ivec2(1, -1) == ivec2(0, -1), e.g. previous
-     * row, 0 offset, same slice, which is zero since we zero out the buffer for RGB */
-    TYPE cur = linecache[1];
-
-    int base = quant_table[quant_table_idx][0][(cur    - top[0]) & MAX_QUANT_TABLE_MASK] +
-               quant_table[quant_table_idx][1][(top[0] - top[1]) & MAX_QUANT_TABLE_MASK] +
-               quant_table[quant_table_idx][2][(top[1] - top[2]) & MAX_QUANT_TABLE_MASK];
-
+    TYPE top2 = TYPE(0);
     if (has_extend_lookup && extend_lookup) {
-        TYPE cur2 = linecache[0];
-        base += quant_table[quant_table_idx][3][(cur2 - cur) & MAX_QUANT_TABLE_MASK];
-
         /* top-2 became current upon swap when rgb_linecache == 2 */
         ivec2 top2_off = off;
         if (rgb_linecache != 2)
             top2_off += ivec2(0, -2);

-        TYPE top2 = TYPE(imageLoad(pred, sp + LADDR(top2_off))[comp]);
-        base += quant_table[quant_table_idx][4][(top2 - top[1]) & MAX_QUANT_TABLE_MASK];
+        top2 = TYPE(imageLoad(pred, sp + LADDR(top2_off))[comp]);
     }

-    /* context, prediction */
-    return ivec2(base, predict(cur, VTYPE2(top)));
+    return ivec4(top, top2);
 }

 #else

 #define LADDR(p) (p)

-ivec2 get_pred(readonly uimage2D pred, ivec2 sp, ivec2 off,
-               uint comp, int sw, uint8_t quant_table_idx, bool extend_lookup)
+ivec4 get_top(readonly uimage2D pred, ivec2 sp, ivec2 off,
+              uint comp, int sw, bool extend_lookup)
 {
     ivec2 yoff_border1 = off.x == 0 ? ivec2(1, -1) : ivec2(0, 0);
     sp += off;
@@ -262,27 +251,57 @@ ivec2 get_pred(readonly uimage2D pred, ivec2 sp, ivec2 off,
         top[2] = TYPE(imageLoad(pred, sp + ivec2(min(1, sw - off.x - 1), -1))[comp]);
     }

-    TYPE cur = linecache[1];
+    TYPE top2 = TYPE(0);
+    if (has_extend_lookup && extend_lookup && off.y > 1)
+        top2 = TYPE(imageLoad(pred, sp + ivec2(0, -2))[comp]);
+
+    return ivec4(top, top2);
+}
+
+#endif /* RGB */

-    int base = quant_table[quant_table_idx][0][(cur - top[0]) & MAX_QUANT_TABLE_MASK] +
-               quant_table[quant_table_idx][1][(top[0] - top[1]) & MAX_QUANT_TABLE_MASK] +
+ivec3 get_pred_top_quant(ivec4 top, uint8_t quant_table_idx, bool extend_lookup)
+{
+    int base = quant_table[quant_table_idx][1][(top[0] - top[1]) & MAX_QUANT_TABLE_MASK] +
                quant_table[quant_table_idx][2][(top[1] - top[2]) & MAX_QUANT_TABLE_MASK];

+    if (has_extend_lookup && extend_lookup)
+        base += quant_table[quant_table_idx][4][(top[3] - top[1]) & MAX_QUANT_TABLE_MASK];
+
+    return ivec3(top[0], top[1], base);
+}
+
+ivec3 get_pred_top(readonly uimage2D pred, ivec2 sp, ivec2 off,
+                   uint comp, int sw, uint8_t quant_table_idx, bool extend_lookup)
+{
+    return get_pred_top_quant(get_top(pred, sp, off, comp, sw, extend_lookup),
+                              quant_table_idx, extend_lookup);
+}
+
+ivec2 get_pred_left(ivec3 top, uint8_t quant_table_idx, bool extend_lookup)
+{
+    /* Normally, we'd need to check if off != ivec2(0, 0) here, since otherwise, we must
+     * return zero. However, ivec2(-1,  0) + ivec2(1, -1) == ivec2(0, -1), e.g. previous
+     * row, 0 offset, same slice, which is zero since we zero out the buffer for RGB */
+    TYPE cur = linecache[1];
+
+    int base = top[2] + quant_table[quant_table_idx][0][(cur - top[0]) & MAX_QUANT_TABLE_MASK];
+
     if (has_extend_lookup && extend_lookup) {
         TYPE cur2 = linecache[0];
         base += quant_table[quant_table_idx][3][(cur2 - cur) & MAX_QUANT_TABLE_MASK];
-
-        TYPE top2 = TYPE(0);
-        if (off.y > 1)
-            top2 = TYPE(imageLoad(pred, sp + ivec2(0, -2))[comp]);
-        base += quant_table[quant_table_idx][4][(top2 - top[1]) & MAX_QUANT_TABLE_MASK];
     }

     /* context, prediction */
-    return ivec2(base, predict(cur, VTYPE2(top)));
+    return ivec2(base, predict(cur, top.xy));
 }

-#endif /* RGB */
+ivec2 get_pred(readonly uimage2D pred, ivec2 sp, ivec2 off,
+               uint comp, int sw, uint8_t quant_table_idx, bool extend_lookup)
+{
+    return get_pred_left(get_pred_top(pred, sp, off, comp, sw, quant_table_idx, extend_lookup),
+                         quant_table_idx, extend_lookup);
+}

 void linecache_load(readonly uimage2D src, ivec2 sp, int y, uint comp)
 {
diff --git a/libavcodec/vulkan/ffv1_dec.comp.glsl b/libavcodec/vulkan/ffv1_dec.comp.glsl
index 8e52ac2aeb..270d513b48 100644
--- a/libavcodec/vulkan/ffv1_dec.comp.glsl
+++ b/libavcodec/vulkan/ffv1_dec.comp.glsl
@@ -22,6 +22,7 @@

 #pragma shader_stage(compute)
 #extension GL_GOOGLE_include_directive : require
+#extension GL_KHR_shader_subgroup_ballot : require

 #define DECODE
 #include "common.glsl"
@@ -47,6 +48,8 @@ layout (set = 1, binding = 3, scalar) buffer slice_state_buf {
     uint8_t slice_rc_state[];
 };

+layout (constant_id = 20) const bool quant_ballot = false;
+
 #define READ(idx) get_rac_state(idx)
 shared int sym_e;
 shared bool rc_dec[CONTEXT_SIZE];
@@ -130,9 +133,23 @@ void decode_line(ivec2 sp, int w,

     linecache_load(dec[p], sp, y, 0);

+    bool ext = extend_lookup[quant_table_idx];
+    ivec2 qthr = quant_ballot ? quant_thresh[quant_table_idx][gl_LocalInvocationID.x] : ivec2(0);
+    ivec2 qso = quant_ballot ? quant_scale_off[quant_table_idx] : ivec2(0);
+
     for (int x = 0; x < w; x++) {
-        ivec2 pr = get_pred(dec[p], sp, ivec2(x, y), 0, w,
-                            quant_table_idx, extend_lookup[quant_table_idx]);
+        ivec2 pr;
+        if (quant_ballot) {
+            ivec3 top = get_pred_top(dec[p], sp, ivec2(x, y), 0, w, quant_table_idx, ext);
+            TYPE cur = linecache[1];
+            uvec4 q0 = subgroupBallot(int(int8_t(cur - top[0])) >= qthr.x);
+            uvec4 q3 = subgroupBallot(ext && int(int8_t(linecache[0] - cur)) >= qthr.y);
+            pr = ivec2(top[2] + qso.y + int(subgroupBallotBitCount(q0)) +
+                       qso.x*int(subgroupBallotBitCount(q3)),
+                       predict(cur, top.xy));
+        } else {
+            pr = get_pred(dec[p], sp, ivec2(x, y), 0, w, quant_table_idx, ext);
+        }

         uint rc_off = state_off + CONTEXT_SIZE*abs(pr[0]) + gl_LocalInvocationID.x;

diff --git a/libavcodec/vulkan_ffv1.c b/libavcodec/vulkan_ffv1.c
index 161f21e7cb..fd272e5293 100644
--- a/libavcodec/vulkan_ffv1.c
+++ b/libavcodec/vulkan_ffv1.c
@@ -687,13 +687,13 @@ static int init_decode_shader(FFV1Context *f, FFVulkanContext *s,
                               AVHWFramesContext *dec_frames_ctx,
                               AVHWFramesContext *out_frames_ctx,
                               VkSpecializationInfo *sl, int ac, int rgb,
-                              int bayer)
+                              int bayer, int ballot)
 {
     int err;

     uint32_t wg_x = ac != AC_GOLOMB_RICE ? CONTEXT_SIZE : 1;
     ff_vk_shader_load(shd, VK_SHADER_STAGE_COMPUTE_BIT, sl,
-                      (uint32_t []) { wg_x, 1, 1 }, 0);
+                      (uint32_t []) { wg_x, 1, 1 }, ballot ? CONTEXT_SIZE : 0);

     ff_vk_shader_add_push_const(shd, 0, sizeof(FFv1ShaderParams),
                                 VK_SHADER_STAGE_COMPUTE_BIT);
@@ -856,6 +856,8 @@ static void vk_decode_ffv1_uninit(FFVulkanDecodeShared *ctx)
 static int vk_decode_ffv1_init(AVCodecContext *avctx)
 {
     int err;
+    int ballot = 0;
+    FFv1QuantBallot qb;
     FFV1Context *f = avctx->priv_data;
     FFVulkanDecodeContext *dec = avctx->internal->hwaccel_priv_data;
     FFVulkanDecodeShared *ctx = NULL;
@@ -897,7 +899,7 @@ static int vk_decode_ffv1_init(AVCodecContext *avctx)
         dctx = (AVHWFramesContext *)fv->intermediate_frames_ref->data;
     }

-    SPEC_LIST_CREATE(sl, 15, 15*sizeof(uint32_t))
+    SPEC_LIST_CREATE(sl, 16, 16*sizeof(uint32_t))
     ff_ffv1_vk_set_common_sl(avctx, f, sl, sw_format);

     if (RGB_LINECACHE != 2)
@@ -906,6 +908,16 @@ static int vk_decode_ffv1_init(AVCodecContext *avctx)
     if (f->ec && !!(avctx->err_recognition & AV_EF_CRCCHECK))
         SPEC_LIST_ADD(sl, 1, 32, 1);

+    /* The ballot quantizer holds one threshold per invocation, so the
+     * workgroup must be a single subgroup of CONTEXT_SIZE invocations */
+    if (f->ac != AC_GOLOMB_RICE &&
+        ctx->s.subgroup_props.minSubgroupSize <= CONTEXT_SIZE &&
+        ctx->s.subgroup_props.maxSubgroupSize >= CONTEXT_SIZE &&
+        (ctx->s.subgroup_props.requiredSubgroupSizeStages & VK_SHADER_STAGE_COMPUTE_BIT))
+        ballot = ff_ffv1_vk_quant_ballot(f, &qb);
+    if (ballot)
+        SPEC_LIST_ADD(sl, 20, 32, 1);
+
     /* Setup shader */
     RET(init_setup_shader(f, &ctx->s, &ctx->exec_pool, &fv->setup, sl));

@@ -914,7 +926,7 @@ static int vk_decode_ffv1_init(AVCodecContext *avctx)

     /* Decode shaders */
     RET(init_decode_shader(f, &ctx->s, &ctx->exec_pool, &fv->decode,
-                           dctx, hwfc, sl, f->ac, is_rgb, f->bayer));
+                           dctx, hwfc, sl, f->ac, is_rgb, f->bayer, ballot));

     /* Init static data */
     RET(ff_ffv1_vk_init_consts(&ctx->s, &fv->consts_buf, f));
@@ -942,7 +954,8 @@ static int vk_decode_ffv1_init(AVCodecContext *avctx)
                                         &fv->consts_buf,
                                         256*sizeof(uint32_t) + 512*sizeof(uint8_t),
                                         MAX_QUANT_TABLES*MAX_CONTEXT_INPUTS*
-                                        MAX_QUANT_TABLE_SIZE*sizeof(int32_t),
+                                        MAX_QUANT_TABLE_SIZE*sizeof(int32_t) +
+                                        sizeof(FFv1QuantBallot),
                                         VK_FORMAT_UNDEFINED));

 fail: