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