vulkan_ffv1: speed up range coded decoding The range coded decoder spent most of its time on control flow and glue around the binary decisions of each symbol. Rework the symbol decoder and the per-sample loop: - States are pre-shifted so that each decision is a 32-bit high multiply. The next bitstream byte is prefetched from the window, and the window is checked once per sample rather than on every refill. - The unary prefix is decoded as an unrolled nest of decisions whose common path has no taken branches, with the zero flag fused into its first level. Each exit decodes the mantissa and the sign in straight-line code specialised for its exponent. - Escapes adapt their states in scalar registers, and the prefix is capped at 32 ones as in the software decoder, which the previous code did not do. - The context and the state load of the next sample are issued as soon as the value is known, before the state of the current sample is adapted. Inputs 0 and 3 of the quant tables are evaluated with a subgroup ballot when the tables allow it. - The top row is prefetched one 32-sample chunk ahead, and decoded samples are written once per chunk. - Slices with up to 3.5KiB of context states keep them in shared memory while decoding. - The quant tables are stored once, as int32, for both the encoder and the decoder, and bound with explicit ranges. Decoding a 6464x4852 16-bit RGB frame with 1024 slices on an RX 6900 XT goes from 80.7/70.3/68.7 ms to 53.4/43.0/37.9 ms (context model 1/0/2). diff --git a/libavcodec/ffv1_vulkan.c b/libavcodec/ffv1_vulkan.c --- a/libavcodec/ffv1_vulkan.c +++ b/libavcodec/ffv1_vulkan.c @@ -82,16 +82,54 @@ } } +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; uint8_t *buf_mapped; + int32_t (*quant_tables)[MAX_CONTEXT_INPUTS][MAX_QUANT_TABLE_SIZE]; size_t buf_len = 256*sizeof(uint32_t) + /* CRC */ 512*sizeof(uint8_t) + /* Rangecoder */ MAX_QUANT_TABLES* MAX_CONTEXT_INPUTS* - MAX_QUANT_TABLE_SIZE*sizeof(int16_t); + MAX_QUANT_TABLE_SIZE*sizeof(int32_t) + + sizeof(FFv1QuantBallot); RET(ff_vk_create_buf(s, vkb, buf_len, @@ -106,8 +144,13 @@ set_rc_state_tab(f, buf_mapped + 256*sizeof(uint32_t)); - memcpy(buf_mapped + 256*sizeof(uint32_t) + 512*sizeof(uint8_t), - f->quant_tables, sizeof(f->quant_tables)); + quant_tables = (void *)(buf_mapped + 256*sizeof(uint32_t) + 512*sizeof(uint8_t)); + for (int i = 0; i < MAX_QUANT_TABLES; i++) + for (int j = 0; j < MAX_CONTEXT_INPUTS; j++) + 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)); diff --git a/libavcodec/ffv1_vulkan.h b/libavcodec/ffv1_vulkan.h --- a/libavcodec/ffv1_vulkan.h +++ b/libavcodec/ffv1_vulkan.h @@ -30,6 +30,21 @@ 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 --- a/libavcodec/ffv1enc_vulkan.c +++ b/libavcodec/ffv1enc_vulkan.c @@ -1546,7 +1546,8 @@ &fv->enc, 0, 1, 0, &fv->consts_buf, 256*sizeof(uint32_t) + 512*sizeof(uint8_t), - VK_WHOLE_SIZE, + MAX_QUANT_TABLES*MAX_CONTEXT_INPUTS* + MAX_QUANT_TABLE_SIZE*sizeof(int32_t), 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 --- a/libavcodec/vulkan/ffv1_common.glsl +++ b/libavcodec/vulkan/ffv1_common.glsl @@ -184,9 +184,13 @@ } layout (set = 0, binding = 1, scalar) readonly uniform quant_buf { - int16_t quant_table[MAX_QUANT_TABLES] + int32_t quant_table[MAX_QUANT_TABLES] [MAX_CONTEXT_INPUTS] [MAX_QUANT_TABLE_SIZE]; +#ifdef DECODE + ivec2 quant_thresh[MAX_QUANT_TABLES][32]; + ivec2 quant_scale_off[MAX_QUANT_TABLES]; +#endif }; /* -1, { -1, 0 } */ @@ -216,8 +220,8 @@ #define RGB_LBUF (rgb_linecache - 1) #define LADDR(p) (ivec2((p).x, ((p).y & RGB_LBUF))) -ivec3 get_pred_top(IMG_QUALI uimage2D pred, ivec2 sp, ivec2 off, - uint comp, int sw, uint8_t quant_table_idx, bool extend_lookup) +ivec4 get_top(IMG_QUALI 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; @@ -226,28 +230,25 @@ 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])); - 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]; - + TYPE top2 = TYPE(0); if (has_extend_lookup && extend_lookup) { /* 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]); } - return ivec3(top[0], top[1], base); + return ivec4(top, top2); } #else #define LADDR(p) (p) -ivec3 get_pred_top(IMG_QUALI uimage2D pred, ivec2 sp, ivec2 off, - uint comp, int sw, uint8_t quant_table_idx, bool extend_lookup) +ivec4 get_top(IMG_QUALI 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,20 +263,32 @@ top[2] = TYPE(imageLoad(pred, sp + ivec2(min(1, sw - off.x - 1), -1))[comp]); } + 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 */ + +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) { - 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]; - } + 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); } -#endif /* RGB */ +ivec3 get_pred_top(IMG_QUALI 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) { diff --git a/libavcodec/vulkan/ffv1_dec.comp.glsl b/libavcodec/vulkan/ffv1_dec.comp.glsl --- a/libavcodec/vulkan/ffv1_dec.comp.glsl +++ b/libavcodec/vulkan/ffv1_dec.comp.glsl @@ -50,72 +50,9 @@ uint8_t slice_rc_state[]; }; -#define READ(idx) get_rac(subgroupBroadcast(st, uint(idx))) - -int get_isymbol(inout uint st, out uint read, out uint bits) -{ - uint s[11]; - [[unroll]] for (int i = 0; i < 11; i++) - s[i] = subgroupBroadcast(st, i); - uint p = subgroupClusteredOr(st << ((gl_LocalInvocationID.x & 3) << 3), 4); - - read = 1u; - bits = 1u; - if (get_rac(s[0])) - return 0; - - int e = int(get_rac_unary(st, s)); - - int n = e; - bool esc = c_bits > 10 && e == 11; - - int a = 1; - int sym_e = e + 10; - uint ss = subgroupBroadcast(st, sym_e); - - if (e <= 9) { - uint64_t q = ((pack64(uvec2(subgroupBroadcast(p, 20), subgroupBroadcast(p, 24))) >> 16) | - (uint64_t(subgroupBroadcast(p, 28)) << 48)) << ((9 - e) << 3); - if (((e - 1) & 8) != 0) - a = get_rac_bits(q, a, 8); - if (((e - 1) & 4) != 0) - a = get_rac_bits(q, a, 4); - if (((e - 1) & 2) != 0) - a = get_rac_bits(q, a, 2); - if (((e - 1) & 1) != 0) - a = get_rac_bits(q, a, 1); - } else { - if (esc) { - do { - if (gl_LocalInvocationID.x == 10) - st = zero_one_state[st + 256]; - e++; - } while (READ(10)); - - bool b = READ(31); - a = b ? 0x3 : 0x2; - for (e -= 2; e >= 11; e--) { - if (gl_LocalInvocationID.x == 31) - st = zero_one_state[st + (b ? 256 : 0)]; - b = READ(31); - a = (a << 1) | int(b); - } - } - - for (e += 20; e >= 22; e--) - a = (a << 1) | int(READ(e)); - } - - bool neg = get_rac(ss); - - uint m = uint(min(n - 1, 10)); - read = 1u | (((1u << min(n, 10)) - 1u) << 1) | - (((1u << m) - 1u) << 22) | (1u << sym_e); - bits = (((1u << (esc ? 9 : n - 1)) - 1u) << 1) | - ((uint(a) & ((1u << m) - 1u)) << 22) | (uint(neg) << sym_e); - - return neg ? -a : a; -} +layout (constant_id = 19) const uint lds_contexts = 0; +layout (constant_id = 20) const bool quant_ballot = false; +shared uint8_t lds_rc_state[codec_planes*CONTEXT_SIZE*(lds_contexts > 0 ? lds_contexts : 1)]; void decode_line_pcm(ivec2 sp, int w, int y, int p) { @@ -139,7 +76,7 @@ void decode_line(ivec2 sp, int w, int y, int p, int bits, uint state_off, - uint8_t quant_table_idx, int run_index) + uint8_t quant_table_idx, int run_index, bool ext, bool lds) { #ifndef RGB if (p > 0 && p < 3) { @@ -150,44 +87,101 @@ linecache_load(dec[p], sp, y, 0); - ivec3 top = get_pred_top(dec[p], sp, ivec2(min(int(gl_LocalInvocationID.x), w - 1), y), - 0, w, quant_table_idx, extend_lookup[quant_table_idx]); - ivec2 pr = get_pred_left(subgroupBroadcast(top, 0u), - quant_table_idx, extend_lookup[quant_table_idx]); - uint ctx = abs(pr[0]); - uint st = slice_rc_state[state_off + CONTEXT_SIZE*ctx + gl_LocalInvocationID.x]; - - for (int x = 0; x < w; x++) { - uint used, used_bits; - int diff = get_isymbol(st, used, used_bits); - if (pr[0] < 0) - diff = -diff; + ivec3 top = subgroupBroadcast(get_pred_top(dec[p], sp, ivec2(0, y), 0, w, + quant_table_idx, ext), 0u); + ivec2 pr = get_pred_left(top, quant_table_idx, ext); + int c = pr[0]; + int pred = pr[1]; + int sgn = c < 0 ? -1 : 1; + int tl = top.y; + int l = linecache[1]; + uint ctx = abs(c); + uint sbase = state_off + gl_LocalInvocationID.x; + uint soff = sbase + CONTEXT_SIZE*ctx; + uint ld = lds ? lds_rc_state[soff] : slice_rc_state[soff]; + uint8_t adapted = uint8_t(0); + bool same = false; + uint row = 0; + 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); + + ivec4 tr = get_top(dec[p], sp, ivec2(min(1 + int(gl_LocalInvocationID.x), w - 1), y), + 0, w, ext); + for (int x = 0; x < w; x += 32) { + ivec3 tn = get_pred_top_quant(tr, quant_table_idx, ext); + tn.z += qso.y; + tr = get_top(dec[p], sp, ivec2(min(x + 33 + int(gl_LocalInvocationID.x), w - 1), y), + 0, w, ext); + int gmin = min(tn.y - tn.x, 0); + int gmax = max(tn.y - tn.x, 0); + int n = min(w - x, 32); + + int j = 0; + do { + uint st = same ? uint(adapted) : ld; + int base = subgroupBroadcast(tn.z, j); + + uint used, used_bits; + int v = get_isymbol(st, pred, sgn, used, used_bits); + if (lds) + rac_renorm(); + uint vz = zero_extend(v, bits); +#ifdef FLOAT + v = int(vz); +#endif + int t = subgroupBroadcast(tn.y, j); + int lo = lds ? subgroupBroadcast(gmin, j) : 0; + int hi = lds ? subgroupBroadcast(gmax, j) : 0; + + uint nst_off = st + (subgroupInverseBallot(uvec4(used_bits, 0, 0, 0)) ? 256 : 0); + uint nst; + if (lds) + nst = zero_one_state[nst_off]; + + if (quant_ballot) { + uvec4 q0 = subgroupBallot(int(int8_t(v - tl)) >= qthr.x); + uvec4 q3 = subgroupBallot(ext && int(int8_t(l - v)) >= qthr.y); + c = base + int(subgroupBallotBitCount(q0)) + qso.x*int(subgroupBallotBitCount(q3)); + } else { + c = base + quant_table[quant_table_idx][0][(v - tl) & MAX_QUANT_TABLE_MASK]; + if (ext) + c += quant_table[quant_table_idx][3][(l - v) & MAX_QUANT_TABLE_MASK]; + } + uint ctx_prev = ctx; + uint soff_prev = soff; + ctx = abs(c); + soff = sbase + CONTEXT_SIZE*ctx; + if (!lds) { + same = ctx == ctx_prev; + if (!same) + ld = slice_rc_state[soff]; + rac_renorm(); + nst = zero_one_state[nst_off]; + } - uint v = zero_extend(pr[1] + diff, bits); - linecache_next(TYPE(v)); + adapted = uint8_t(subgroupInverseBallot(uvec4(used, 0, 0, 0)) ? nst : st); + if (lds) { + lds_rc_state[soff_prev] = adapted; + ld = lds_rc_state[soff]; + } else { + slice_rc_state[soff_prev] = adapted; + } + sgn = c < 0 ? -1 : 1; - bool upd = bitfieldExtract(used, int(gl_LocalInvocationID.x), 1) != 0; - if (upd) - st = zero_one_state[st + (bitfieldExtract(used_bits, int(gl_LocalInvocationID.x), 1) << 8)]; - - uint ctx_prev = ctx; - uint adapted = st; - if (x + 1 < w) { - if (((x + 1) & 31) == 0) - top = get_pred_top(dec[p], sp, ivec2(min(x + 1 + int(gl_LocalInvocationID.x), w - 1), y), - 0, w, quant_table_idx, extend_lookup[quant_table_idx]); - - pr = get_pred_left(subgroupBroadcast(top, uint((x + 1) & 31)), - quant_table_idx, extend_lookup[quant_table_idx]); - ctx = abs(pr[0]); - if (ctx != ctx_prev) - st = slice_rc_state[state_off + CONTEXT_SIZE*ctx + gl_LocalInvocationID.x]; - } + int vm = int(TYPE(vz)); + if (lds) + pred = clamp(t, vm + lo, vm + hi); + else + pred = subgroupBroadcast(clamp(tn.y, vm + gmin, vm + gmax), j); + row = gl_LocalInvocationID.x == j ? vz : row; + rac_check_window(); + + l = v; + tl = t; + } while (++j < n); - if (gl_LocalInvocationID.x == 0) - imageStore(dec[p], sp + LADDR(ivec2(x, y)), uvec4(v)); - if (upd) - slice_rc_state[state_off + CONTEXT_SIZE*ctx_prev + gl_LocalInvocationID.x] = uint8_t(adapted); + if (gl_LocalInvocationID.x < n) + imageStore(dec[p], sp + LADDR(ivec2(x + int(gl_LocalInvocationID.x), y)), uvec4(row)); } memoryBarrierImage(); @@ -214,7 +208,7 @@ void decode_line(ivec2 sp, int w, int y, int p, int bits, uint state_off, - uint8_t quant_table_idx, inout int run_index) + uint8_t quant_table_idx, inout int run_index, bool ext, bool lds) { #ifndef RGB if (p > 0 && p < 3) { @@ -232,7 +226,7 @@ ivec2 pos = sp + ivec2(x, y); int diff; ivec2 pr = get_pred(dec[p], sp, ivec2(x, y), 0, w, - quant_table_idx, extend_lookup[quant_table_idx]); + quant_table_idx, ext); uint vlc_off = state_off + abs(pr[0]); @@ -385,6 +379,17 @@ } #endif +void decode_plane_line(ivec2 sp, int w, int y, int p, int bits, uint state_off, + uint8_t quant_table_idx, inout int run_index, bool use_lds) +{ + if (use_lds) + decode_line(sp, w, y, p, bits, state_off, quant_table_idx, run_index, false, true); + else if (has_extend_lookup && extend_lookup[quant_table_idx]) + decode_line(sp, w, y, p, bits, state_off, quant_table_idx, run_index, true, false); + else + decode_line(sp, w, y, p, bits, state_off, quant_table_idx, run_index, false, false); +} + void decode_slice(in SliceContext sc, uint slice_idx) { int w = sc.slice_dim.x; @@ -442,25 +447,38 @@ #ifdef BAYER u8vec4 quant_table_idx = sc.quant_table_idx.xzyy; - u32vec4 slice_state_off = (slice_idx*codec_planes + - uvec4(0, 2, 1, 1))*plane_state_size; + uvec4 state_plane = uvec4(0, 2, 1, 1); #else u8vec4 quant_table_idx = sc.quant_table_idx.xyyz; - u32vec4 slice_state_off = (slice_idx*codec_planes + - uvec4(0, 1, 1, 2))*plane_state_size; + uvec4 state_plane = uvec4(0, 1, 1, 2); #endif + u32vec4 slice_state_off = (slice_idx*codec_planes + state_plane)*plane_state_size; #ifdef GOLOMB slice_state_off >>= 3; // division by VLC_STATE_SIZE golomb_init(); + bool use_lds = false; +#else + bool use_lds = lds_contexts > 0; + for (int i = 0; i < codec_planes; i++) + use_lds = use_lds && context_count[sc.quant_table_idx[i]] <= lds_contexts && + !extend_lookup[sc.quant_table_idx[i]]; + + if (use_lds) { + for (int i = 0; i < codec_planes; i++) + for (uint j = gl_LocalInvocationID.x; j < lds_contexts*CONTEXT_SIZE; j += CONTEXT_SIZE) + lds_rc_state[i*lds_contexts*CONTEXT_SIZE + j] = + slice_rc_state[(slice_idx*codec_planes + i)*plane_state_size + j]; + slice_state_off = state_plane*lds_contexts*CONTEXT_SIZE; + } #endif #ifdef BAYER int run_index = 0; for (int y = 0; y < bayer_h; y++) { for (int p = 0; p < 4; p++) - decode_line(sp, w, y, p, bits[p], - slice_state_off[p], quant_table_idx[p], run_index); + decode_plane_line(sp, w, y, p, bits[p], slice_state_off[p], + quant_table_idx[p], run_index, use_lds); writeout_bayer(slice_idx, sc, sp, w, y); } @@ -468,8 +486,8 @@ int run_index = 0; for (int y = 0; y < sc.slice_dim.y; y++) { for (int p = 0; p < color_planes; p++) - decode_line(sp, w, y, p, bits[p], - slice_state_off[p], quant_table_idx[p], run_index); + decode_plane_line(sp, w, y, p, bits[p], slice_state_off[p], + quant_table_idx[p], run_index, use_lds); writeout_rgb(slice_idx, sc, sp, w, y, true); } @@ -481,10 +499,18 @@ int run_index = 0; for (int y = 0; y < h; y++) - decode_line(sp, w, y, p, bits[p], - slice_state_off[p], quant_table_idx[p], run_index); + decode_plane_line(sp, w, y, p, bits[p], slice_state_off[p], + quant_table_idx[p], run_index, use_lds); } #endif + +#ifndef GOLOMB + if (use_lds) + for (int i = 0; i < codec_planes; i++) + for (uint j = gl_LocalInvocationID.x; j < lds_contexts*CONTEXT_SIZE; j += CONTEXT_SIZE) + slice_rc_state[(slice_idx*codec_planes + i)*plane_state_size + j] = + lds_rc_state[i*lds_contexts*CONTEXT_SIZE + j]; +#endif } void main(void) @@ -500,6 +526,9 @@ decode_slice(slice_ctx[slice_idx], slice_idx); if (gl_LocalInvocationID.x == 0) { +#ifndef GOLOMB + rc.bs_off += rc_pos; +#endif uint overread = 0; if (rc.bs_off >= (rc.bs_end + MAX_OVERREAD)) overread = rc.bs_off - rc.bs_end; diff --git a/libavcodec/vulkan/rangecoder_subgroup.glsl b/libavcodec/vulkan/rangecoder_subgroup.glsl --- a/libavcodec/vulkan/rangecoder_subgroup.glsl +++ b/libavcodec/vulkan/rangecoder_subgroup.glsl @@ -25,7 +25,6 @@ #extension GL_KHR_shader_subgroup_basic : require #extension GL_KHR_shader_subgroup_ballot : require -#extension GL_KHR_shader_subgroup_clustered : require #extension GL_KHR_shader_subgroup_rotate : require #define CONTEXT_SIZE 32 @@ -58,11 +57,15 @@ #ifdef DECODE uint rc_win; uint rc_dist; +uint rc_next; +uint rc_pos; void rac_load_window(void) { - uint o = (rc.bs_off & ~31u) + gl_SubgroupInvocationID; - rc_win = ~uint(o < rc.bs_end ? u8buf(uint64_t(slice_data) + o).v : uint8_t(0)) << 24; + rc.bs_off += rc_pos; + rc_pos &= ~31u; + uint o = rc.bs_off + gl_SubgroupInvocationID; + rc_win = uint(~(o < rc.bs_end ? u8buf(uint64_t(slice_data) + o).v : uint8_t(0))); } void rac_init_dec(in RangeCoder c) @@ -73,16 +76,35 @@ rc = c; rc_dist = rc.range - rc.low - 1; + rc_pos = 0; rac_load_window(); + rc_next = subgroupBroadcast(rc_win, 0); } -void refill(void) +void rac_check_window(void) { - uint b = subgroupBroadcast(rc_win, rc.bs_off & 31u); - if ((++rc.bs_off & 31u) == 0) + if (rc_pos > 10) rac_load_window(); +} + +void refill(void) +{ rc.range <<= 8; - rc_dist = unpack32(pack64(u32vec2(b, rc_dist)) << 8).y; + rc_dist = (rc_dist << 8) | rc_next; + rc_next = subgroupBroadcast(rc_win, ++rc_pos); +} + +void rac_renorm(void) +{ + if (expectEXT(rc.range < 0x100, false)) + refill(); +} + +uint rac_range1(uint range, uint state24) +{ + uint hi, lo; + umulExtended(range, state24, hi, lo); + return hi; } bool get_rac_internal(uint range1) @@ -92,71 +114,249 @@ uint distd = rc_dist - range1; rc.range = bit ? range1 : ranged; rc_dist = bit ? rc_dist : distd; - - if (expectEXT(rc.range < 0x100, false)) - refill(); - return bit; } -bool get_rac(uint state) +bool get_rac(uint state24) { - return get_rac_internal(rc.range * state >> 8); + bool bit = get_rac_internal(rac_range1(rc.range, state24)); + rac_renorm(); + return bit; } bool get_rac_equi(void) { - return get_rac_internal(rc.range >> 1); + rac_check_window(); + bool bit = get_rac_internal(rc.range >> 1); + rac_renorm(); + 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; } -int get_rac_bits(inout uint64_t q, int a, int n) +const int AVERROR_INVALIDDATA = -0x41444E49; + +int get_isymbol_esc(inout uint st, uint sx, int pred, int sgn, out uint read, out uint bits) { - [[unroll]] for (int k = 0; k < n; k++) { - a = (a << 1) + int(get_rac(uint(q >> 56))); - q <<= 8; + bool esc = c_bits > 10; + int n = 11; + if (esc) { + uint s10 = subgroupBroadcast(st, 10); + bool one; + [[dont_unroll]] do { + s10 = rangecoder_state[s10 + 256]; + n++; + rac_check_window(); + one = get_rac(s10 << 24); + } while (one && n < 33); + st = gl_SubgroupInvocationID == 10 ? s10 : st; + + if (one) { + read = 0x7FFu; + bits = 0x7FEu; + return pred + sgn*AVERROR_INVALIDDATA; + } } - return a; + + uint s31 = subgroupBroadcast(st, 31); + rac_check_window(); + bool b = get_rac(s31 << 24); + uint a = b ? 0x3 : 0x2; + for (n -= 2; n >= 11; n--) { + s31 = rangecoder_state[s31 + (b ? 256 : 0)]; + rac_check_window(); + b = get_rac(s31 << 24); + a = (a << 1) | uint(b); + } + st = gl_SubgroupInvocationID == 31 ? s31 : st; + + rac_check_window(); + [[unroll]] for (int k = 8; k >= 0; k--) + a = (a << 1) | uint(get_rac(subgroupBroadcast(sx, 22 + k))); + + 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); + return neg ? pred - sa : pred + sa; } -uint get_rac_unary(uint states, uint s[11]) +int get_isymbol(inout uint st, int pred, int sgn, out uint read, out uint bits) { + uint st24 = st << 24; + uint s[11]; + [[unroll]] for (int i = 0; i < 5; i++) + s[i] = subgroupBroadcast(st24, i); + + read = 1u; + bits = 1u; + 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); - uint range1; - uint i = 1; - [[unroll]] for (; i <= 10; i++) { - range1 = range * s[i] >> 8; - if (range1 <= lim) - break; - range = range1; - } - - if (expectEXT(range1 > rc_dist, false)) { - while (i <= 10) { - range = range1; - if (range < 0x100) { - rc.range = range; - refill(); - range = rc.range; - } - if (++i > 10) + int v; + uint sx = st24; + while (true) { + uint skip; + s[5] = subgroupBroadcast(sx, 5); + r[2] = rac_range1(r[1], s[2]); + if (r[1] > lim) { + s[6] = subgroupBroadcast(sx, 6); + r[3] = rac_range1(r[2], s[3]); + if (r[2] > lim) { + s[7] = subgroupBroadcast(sx, 7); + r[4] = rac_range1(r[3], s[4]); + if (r[3] > lim) { + s[8] = subgroupBroadcast(sx, 8); + r[5] = rac_range1(r[4], s[5]); + if (r[4] > lim) { + s[9] = subgroupBroadcast(sx, 9); + r[6] = rac_range1(r[5], s[6]); + if (r[5] > lim) { + s[10] = subgroupBroadcast(sx, 10); + r[7] = rac_range1(r[6], s[7]); + if (r[6] > lim) { + r[8] = rac_range1(r[7], s[8]); + if (r[7] > lim) { + r[9] = rac_range1(r[8], s[9]); + if (r[8] > lim) { + r[10] = rac_range1(r[9], s[10]); + 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[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[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[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[2] <= rc_dist) { + v = get_isymbol_tail(2, r[1], r[2], sx, pred, sgn, read, bits); break; - range1 = range * subgroupBroadcast(states, i) >> 8; - if (range1 <= rc_dist) + } 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]; } - if (i > 10) { - rc.range = range; - return i; - } - } - rc.range = range - range1; - rc_dist -= range1; - if (expectEXT(rc.range < 0x100, false)) refill(); - return i; + 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 < 5; i++) + s[i] = subgroupBroadcast(sx, i); + } + + return v; } #endif diff --git a/libavcodec/vulkan_ffv1.c b/libavcodec/vulkan_ffv1.c --- a/libavcodec/vulkan_ffv1.c +++ b/libavcodec/vulkan_ffv1.c @@ -919,7 +919,7 @@ dctx = (AVHWFramesContext *)fv->intermediate_frames_ref->data; } - SPEC_LIST_CREATE(sl, 15, 15*sizeof(uint32_t)) + SPEC_LIST_CREATE(sl, 17, 17*sizeof(uint32_t)) ff_ffv1_vk_set_common_sl(avctx, f, sl, sw_format); if (RGB_LINECACHE != 2) @@ -928,6 +928,24 @@ if (f->ec && !!(avctx->err_recognition & AV_EF_CRCCHECK)) SPEC_LIST_ADD(sl, 1, 32, 1); + if (f->ac != AC_GOLOMB_RICE) { + FFv1QuantBallot qb; + int lds_contexts = 0; + + /* Slices whose context states take up to 3.5KiB keep them in shared + * memory while decoding. Together with the 512 byte state transition + * table, this still lets 16 slices share 64KiB of shared memory. */ + for (int i = 0; i < f->quant_table_count; i++) + if (f->plane_count*CONTEXT_SIZE*f->context_count[i] <= 3584 && + !(f->quant_tables[i][3][127] || f->quant_tables[i][4][127])) + lds_contexts = FFMAX(lds_contexts, f->context_count[i]); + if (lds_contexts) + SPEC_LIST_ADD(sl, 19, 32, lds_contexts); + + if (ff_ffv1_vk_quant_ballot(f, &qb)) + SPEC_LIST_ADD(sl, 20, 32, 1); + } + /* Setup shader */ RET(init_setup_shader(f, &ctx->s, &ctx->exec_pool, &fv->setup, sl)); @@ -963,7 +981,9 @@ &fv->decode, 0, 1, 0, &fv->consts_buf, 256*sizeof(uint32_t) + 512*sizeof(uint8_t), - VK_WHOLE_SIZE, + MAX_QUANT_TABLES*MAX_CONTEXT_INPUTS* + MAX_QUANT_TABLE_SIZE*sizeof(int32_t) + + sizeof(FFv1QuantBallot), VK_FORMAT_UNDEFINED)); fail: