Commit 972d1d953e for ffmpeg

commit 972d1d953ee611bef31c2194a33c8d4b0ff4388a
Author: Lynne <dev@lynne.ee>
Date:   Sat Sep 26 16:35:16 2026 +0900

    vulkan_ffv1: decode the unary prefix with a search over its ranges

    While the unary prefix of a symbol decodes ones, the distance
    range - low - 1 does not change, and the range of each level is the
    range of the level before scaled by its state. The ranges of all ten
    levels can therefore be computed before any of them is decided. They
    only shrink, so the levels whose range stays above
    max(range - low - 1, 0xff), those that decode a one without a
    renormalisation, form a prefix.

    The first five ranges are computed back to back, and a search that
    starts at the fifth level finds the end of the prefix with two or
    three compares instead of a branch per level. The upper five ranges
    are only computed when the prefix reaches past the fifth level. The
    zero flag is folded into the first level: its subtraction from the
    distance wraps when the flag is set, so the same compares reject that
    case. A level that needs a renormalisation refills and restarts the
    search with the levels already decided neutralised.

    Each exit decodes the mantissa and the sign in straight-line code
    specialised for its exponent, so the masks of the states a symbol used
    are constant expressions of it. The sign decision leaves its
    renormalisation to the caller.

    Decoding a 6464x4852 16-bit RGB frame with 1024 slices on an RX 6900
    XT, with the bitstream in VRAM, goes from 126.6/106.2/104.7 ms to
    99.1/72.4/69.8 ms with context model 1/0/2.

diff --git a/libavcodec/vulkan/ffv1_dec.comp.glsl b/libavcodec/vulkan/ffv1_dec.comp.glsl
index e5f2d79f1d..c5b276b6b8 100644
--- a/libavcodec/vulkan/ffv1_dec.comp.glsl
+++ b/libavcodec/vulkan/ffv1_dec.comp.glsl
@@ -110,6 +110,7 @@ void decode_line(ivec2 sp, int w,

         uint used, used_bits;
         int v = get_isymbol(st, pr[1], sgn, used, used_bits);
+        rac_renorm();
         uint vz = zero_extend(v, bits);
         rac_check_window();

diff --git a/libavcodec/vulkan/rangecoder_subgroup.glsl b/libavcodec/vulkan/rangecoder_subgroup.glsl
index 81937357f8..da787dc99b 100644
--- a/libavcodec/vulkan/rangecoder_subgroup.glsl
+++ b/libavcodec/vulkan/rangecoder_subgroup.glsl
@@ -131,6 +131,38 @@ bool get_rac_equi(void)
     return bit;
 }

+int get_isymbol_tail(int e, uint range, uint range1, uint sx, int pred, int sgn,
+                     out uint read, out uint bits)
+{
+    uint m[9];
+    [[unroll]] for (int k = 8; k >= 0; k--)
+        if (k + 1 < e)
+            m[k] = subgroupBroadcast(sx, 22 + k);
+    uint ss = subgroupBroadcast(sx, 10 + e);
+
+    rc.range = range - range1;
+    rc_dist -= range1;
+    rac_renorm();
+
+    uint a = 0;
+    [[unroll]] for (int k = 8; k >= 0; k--) {
+        if (k + 1 < e) {
+            a = (a << 1) + uint(get_rac_internal(rac_range1(rc.range, m[k])));
+            rac_renorm();
+        }
+    }
+
+    a += 1u << (e - 1);
+    int sa = int(a)*sgn;
+    int vp = pred + sa;
+    int vn = vp - 2*sa;
+    bool neg = get_rac_internal(rac_range1(rc.range, ss));
+    int v = neg ? vn : vp;
+    read = (2u << e) - 1u + (((1u << (e - 1)) - 1u) << 22) + (1u << (10 + e));
+    bits = (a << 22) + (1u << e) - 2u - (1u << (21 + e)) + (neg ? 1u << (10 + e) : 0u);
+    return v;
+}
+
 const int AVERROR_INVALIDDATA = -0x41444E49;

 int get_isymbol_esc(inout uint st, uint sx, int pred, int sgn, out uint read, out uint bits)
@@ -171,7 +203,7 @@ int get_isymbol_esc(inout uint st, uint sx, int pred, int sgn, out uint read, ou
     [[unroll]] for (int k = 8; k >= 0; k--)
         a = (a << 1) | uint(get_rac(subgroupBroadcast(sx, 22 + k)));

-    bool neg = get_rac(subgroupBroadcast(sx, 21));
+    bool neg = get_rac_internal(rac_range1(rc.range, subgroupBroadcast(sx, 21)));
     int sa = int(a)*sgn;
     read = 0xFFE007FFu;
     bits = ((esc ? 0x1FFu : 0x3FFu) << 1) | ((a & 0x3FFu) << 22) | (uint(neg) << 21);
@@ -181,27 +213,136 @@ int get_isymbol_esc(inout uint st, uint sx, int pred, int sgn, out uint read, ou
 int get_isymbol(inout uint st, int pred, int sgn, out uint read, out uint bits)
 {
     uint st24 = st << 24;
-    if (get_rac(subgroupBroadcast(st24, 0))) {
-        read = 1u;
-        bits = 1u;
-        return pred;
-    }
+    uint s[11];
+    [[unroll]] for (int i = 0; i < 11; i++)
+        s[i] = subgroupBroadcast(st24, i);

-    int e = 1;
-    while (e < 11 && get_rac(subgroupBroadcast(st24, e)))
-        e++;
-    if (e == 11)
-        return get_isymbol_esc(st, st24, pred, sgn, read, bits);
+    read = 1u;
+    bits = 1u;

-    uint a = 1u;
-    for (int k = e - 2; k >= 0; k--)
-        a = (a << 1) | uint(get_rac(subgroupBroadcast(st24, 22 + k)));
-    bool neg = get_rac(subgroupBroadcast(st24, 10 + e));
+    uint range = rc.range;
+    uint dist = rc_dist;
+    uint range0 = rac_range1(range, s[0]);
+    uint r[11];
+    r[0] = range - range0;
+    r[1] = rac_range1(r[0], s[1]);
+    rc.range = r[0];
+    rc_dist = dist - range0;
+    uint lim = max(rc_dist, 0xffu);

-    read = (2u << e) - 1u + (((1u << (e - 1)) - 1u) << 22) + (1u << (10 + e));
-    bits = (a << 22) + (1u << e) - 2u - (1u << (21 + e)) + (neg ? 1u << (10 + e) : 0u);
-    int sa = int(a)*sgn;
-    return neg ? pred - sa : pred + sa;
+    int v;
+    uint sx = st24;
+    while (true) {
+        uint skip;
+        [[unroll]] for (int i = 2; i < 6; i++)
+            r[i] = rac_range1(r[i - 1], s[i]);
+
+        if (r[5] > lim) {
+            [[unroll]] for (int i = 6; i < 11; i++)
+                r[i] = rac_range1(r[i - 1], s[i]);
+            if (r[6] > lim) {
+                if (r[7] > lim) {
+                    if (r[8] > lim) {
+                        if (r[9] > lim) {
+                            if (r[10] > lim) {
+                                rc.range = r[10];
+                                v = get_isymbol_esc(st, sx, pred, sgn, read, bits);
+                                break;
+                            } else if (r[10] <= rc_dist) {
+                                v = get_isymbol_tail(10, r[9], r[10], sx, pred, sgn, read, bits);
+                                break;
+                            } else {
+                                skip = 10;
+                                rc.range = r[10];
+                            }
+                        } else if (r[9] <= rc_dist) {
+                            v = get_isymbol_tail(9, r[8], r[9], sx, pred, sgn, read, bits);
+                            break;
+                        } else {
+                            skip = 9;
+                            rc.range = r[9];
+                        }
+                    } else if (r[8] <= rc_dist) {
+                        v = get_isymbol_tail(8, r[7], r[8], sx, pred, sgn, read, bits);
+                        break;
+                    } else {
+                        skip = 8;
+                        rc.range = r[8];
+                    }
+                } else if (r[7] <= rc_dist) {
+                    v = get_isymbol_tail(7, r[6], r[7], sx, pred, sgn, read, bits);
+                    break;
+                } else {
+                    skip = 7;
+                    rc.range = r[7];
+                }
+            } else if (r[6] <= rc_dist) {
+                v = get_isymbol_tail(6, r[5], r[6], sx, pred, sgn, read, bits);
+                break;
+            } else {
+                skip = 6;
+                rc.range = r[6];
+            }
+        } else if (r[4] > lim) {
+            if (r[5] <= rc_dist) {
+                v = get_isymbol_tail(5, r[4], r[5], sx, pred, sgn, read, bits);
+                break;
+            } else {
+                skip = 5;
+                rc.range = r[5];
+            }
+        } else if (r[3] > lim) {
+            if (r[4] <= rc_dist) {
+                v = get_isymbol_tail(4, r[3], r[4], sx, pred, sgn, read, bits);
+                break;
+            } else {
+                skip = 4;
+                rc.range = r[4];
+            }
+        } else if (r[2] > lim) {
+            if (r[3] <= rc_dist) {
+                v = get_isymbol_tail(3, r[2], r[3], sx, pred, sgn, read, bits);
+                break;
+            } else {
+                skip = 3;
+                rc.range = r[3];
+            }
+        } else if (r[1] > lim) {
+            if (r[2] <= rc_dist) {
+                v = get_isymbol_tail(2, r[1], r[2], sx, pred, sgn, read, bits);
+                break;
+            } else {
+                skip = 2;
+                rc.range = r[2];
+            }
+        } else if (range0 > min(dist, range - 0x100)) {
+            if (dist < range0) {
+                rc.range = range0;
+                rc_dist = dist;
+                v = pred;
+                break;
+            }
+            skip = 0;
+            rc.range = r[0];
+        } else if (r[1] <= rc_dist) {
+            v = get_isymbol_tail(1, r[0], r[1], sx, pred, sgn, read, bits);
+            break;
+        } else {
+            skip = 1;
+            rc.range = r[1];
+        }
+
+        refill();
+        lim = max(rc_dist, 0xffu);
+        r[0] = rc.range;
+        r[1] = skip > 0 ? rc.range + skip - 1 : rac_range1(rc.range, s[1]);
+        range0 = 0;
+        sx = subgroupInverseBallot(uvec4((2u << skip) - 2u, 0, 0, 0)) ? ~0u : sx;
+        [[unroll]] for (int i = 2; i < 11; i++)
+            s[i] = subgroupBroadcast(sx, i);
+    }
+
+    return v;
 }
 #endif