From f9ea7f057df927d99539e149da37e166a929e958 Mon Sep 17 00:00:00 2001 From: GF Date: Sat, 26 Sep 2026 03:03:26 -0400 Subject: [PATCH] perf(metal): use 32-bit checked arithmetic and one thread per HP block Every reconstruction kernel widened each add, subtract, and multiply to 64-bit, which Apple GPUs emulate, and returned through a branch after each operation. Arithmetic now uses exact 32-bit overflow predicates (sign-bit tests for add/sub, mulhi for quantizer products, range tests for the times-three steps) accumulated into a sticky flag that each kernel tests once before storing. The transforms become straight-line code; per-phase status codes are unchanged, and results derived from an overflowed intermediate are never stored. The HP kernel ran one threadgroup per macroblock with thread 0 applying HP prediction serially between two barriers, leaving half of each SIMD group idle for 16-block luma macroblocks. It now runs one thread per 4x4 block in a flat grid; each thread accumulates its own prediction chain in the normative order, so every partial-sum overflow check matches the serial traversal. Block rows are stored as aligned int4 writes, and plane ABI construction now rejects sample planes that are not four-sample aligned. The output kernels converted color, including chroma upsampling, once per output channel. Each pixel is now loaded and converted once, and premultiplied stores scale alpha once per pixel. Chroma upsampling computes the weighted average exactly in 32 bits by splitting each operand into 8q + r, and unsigned premultiplication uses 32-bit division (65535^2 + 32767 fits in u32). --- crates/jxr-metal/src/abi.rs | 6 + crates/jxr-metal/src/encode.rs | 35 +- crates/jxr-metal/src/encode/batch.rs | 24 +- crates/jxr-metal/src/kernels/common.metal | 171 ++++---- .../src/kernels/first_transform.metal | 31 +- .../src/kernels/highpass_transform.metal | 138 +++--- crates/jxr-metal/src/kernels/output.metal | 403 ++++++++++-------- crates/jxr-metal/src/kernels/overlap.metal | 247 +++++------ 8 files changed, 540 insertions(+), 515 deletions(-) diff --git a/crates/jxr-metal/src/abi.rs b/crates/jxr-metal/src/abi.rs index 9e925bd..632a566 100644 --- a/crates/jxr-metal/src/abi.rs +++ b/crates/jxr-metal/src/abi.rs @@ -122,6 +122,12 @@ impl JxrPlaneAbi { let u32_value = |value: usize, reason| { u32::try_from(value).map_err(|_| MetalError::InvalidPlan { reason }) }; + // The HP kernel stores each 4x4 block row as one aligned `int4`. + if !plane.sample_offset.is_multiple_of(4) || !plane.sample_width.is_multiple_of(4) { + return Err(MetalError::InvalidPlan { + reason: "sample plane is not aligned for four-sample block rows", + }); + } Ok(Self { macroblock_offset: u32_value( plane.macroblock_offset, diff --git a/crates/jxr-metal/src/encode.rs b/crates/jxr-metal/src/encode.rs index 1c1f630..7e39e52 100644 --- a/crates/jxr-metal/src/encode.rs +++ b/crates/jxr-metal/src/encode.rs @@ -5,7 +5,7 @@ use jxr_core::device_plan::{SAMPLE_OFFSET, SURFACE_OFFSET}; use j2k_metal_support::{ checked_buffer_fill_bytes, checked_command_buffer, checked_compute_command_encoder, checked_event, checked_private_buffer, checked_shared_buffer_with_slice, dispatch_1d_pipeline, - dispatch_2d_pipeline, mtl_size, one_d_threads_per_group, + dispatch_2d_pipeline, }; use jxr_core::{OverlapMode, SurfaceLayout}; use objc2::{rc::Retained, runtime::ProtocolObject}; @@ -525,6 +525,7 @@ fn encode_highpass( plane: JxrPlaneAbi, macroblock_count: usize, ) -> Result<(), MetalError> { + let threads = highpass_thread_count(macroblock_count, plane.block_columns, plane.block_rows)?; encoder.setComputePipelineState(&runtime.hp_transform); encoder.bind_buffer(0, &arena.packed, arena.packed_offset_bytes)?; encoder.bind_buffer(1, &arena.macroblocks, 0)?; @@ -532,25 +533,25 @@ fn encode_highpass( encoder.bind_buffer(3, buffers.samples.buffer(), 0)?; encoder.bind_buffer(4, buffers.status.buffer(), 0)?; encoder.bind_bytes(5, &plane)?; - let width = (runtime.hp_transform.threadExecutionWidth() as u64).max(16); - if width > runtime.hp_transform.maxTotalThreadsPerThreadgroup() as u64 { - return Err(MetalError::InvalidPlan { - reason: "Metal pipeline cannot host one transform macroblock", - }); - } - encoder.dispatchThreadgroups_threadsPerThreadgroup( - mtl_size( - u64::try_from(macroblock_count).map_err(|_| MetalError::InvalidPlan { - reason: "macroblock threadgroup count exceeds u64", - })?, - 1, - 1, - ), - one_d_threads_per_group(width), - ); + dispatch_1d_pipeline(encoder, &runtime.hp_transform, threads); Ok(()) } +/// Returns the one-thread-per-block grid width for the HP transform kernel. +fn highpass_thread_count( + macroblock_count: usize, + block_columns: impl Into, + block_rows: impl Into, +) -> Result { + u64::try_from(macroblock_count) + .ok() + .and_then(|count| count.checked_mul(block_columns.into() * block_rows.into())) + .filter(|&threads| u32::try_from(threads).is_ok()) + .ok_or(MetalError::InvalidPlan { + reason: "HP transform block count exceeds the Metal ABI", + }) +} + fn encode_overlap_schedule( device: &ProtocolObject, encoder: &ProtocolObject, diff --git a/crates/jxr-metal/src/encode/batch.rs b/crates/jxr-metal/src/encode/batch.rs index e1ee50a..040553f 100644 --- a/crates/jxr-metal/src/encode/batch.rs +++ b/crates/jxr-metal/src/encode/batch.rs @@ -327,20 +327,16 @@ fn encode_highpass_transforms( buffers: &BatchBuffers, ) -> Result<(), MetalError> { let plane_count = buffers.plane_inputs[0].len(); - let width = (runtime.batch_hp_transform.threadExecutionWidth() as u64).max(16); - if width > runtime.batch_hp_transform.maxTotalThreadsPerThreadgroup() as u64 { - return Err(invalid( - "batch HP pipeline cannot host one transform macroblock", - )); + let mut work = 0; + for plane in buffers.plane_inputs.iter().flat_map(|planes| planes.iter()) { + work = work.max(super::highpass_thread_count( + plane.macroblock_count, + plane.block_columns, + plane.block_rows, + )?); } - let work = buffers - .plane_inputs - .iter() - .flat_map(|planes| planes.iter()) - .map(|plane| plane.macroblock_count) - .max() - .unwrap_or(0); let batch = batch_dispatch(plans.len(), plane_count)?; + let width = runtime.batch_hp_transform.threadExecutionWidth() as u64; encoder.setComputePipelineState(&runtime.batch_hp_transform); encoder.bind_buffer(0, &buffers.packed, 0)?; encoder.bind_buffer(1, &buffers.macroblocks, 0)?; @@ -349,9 +345,9 @@ fn encode_highpass_transforms( encoder.bind_buffer(4, buffers.status.buffer(), 0)?; encoder.bind_buffer(5, &buffers.planes, 0)?; encoder.bind_bytes(6, &batch)?; - encoder.dispatchThreadgroups_threadsPerThreadgroup( + encoder.dispatchThreads_threadsPerThreadgroup( mtl_size( - u64::try_from(work).map_err(|_| invalid("batch HP work exceeds u64"))?, + work, u64::from(batch.image_count), u64::from(batch.plane_count), ), diff --git a/crates/jxr-metal/src/kernels/common.metal b/crates/jxr-metal/src/kernels/common.metal index 61f7732..8d934a1 100644 --- a/crates/jxr-metal/src/kernels/common.metal +++ b/crates/jxr-metal/src/kernels/common.metal @@ -12,117 +12,108 @@ inline void jxr_fail(device atomic_uint *status, uint code) { status, &expected, code, memory_order_relaxed, memory_order_relaxed)) {} } -inline bool jxr_narrow(long value, thread int &result, device atomic_uint *status, uint code) { - if (value < long(INT_MIN) || value > long(INT_MAX)) { - jxr_fail(status, code); - return false; - } - result = int(value); - return true; +// Checked 32-bit arithmetic. Each helper returns the wrapped two's-complement +// result and ORs an exact overflow predicate into a sticky flag. Kernels test +// the flag once before storing, which keeps the transforms branch-free and +// avoids emulated 64-bit integer math on Apple GPUs. A result derived from an +// overflowed intermediate is never stored because the flag stays set. +inline int jxr_add(int a, int b, thread bool &overflow) { + const int result = as_type(as_type(a) + as_type(b)); + overflow |= ((a ^ result) & (b ^ result)) < 0; + return result; } -inline bool jxr_add(int a, int b, thread int &result, device atomic_uint *status, uint code) { - return jxr_narrow(long(a) + long(b), result, status, code); +inline int jxr_sub(int a, int b, thread bool &overflow) { + const int result = as_type(as_type(a) - as_type(b)); + overflow |= ((a ^ b) & (a ^ result)) < 0; + return result; } -inline bool jxr_sub(int a, int b, thread int &result, device atomic_uint *status, uint code) { - return jxr_narrow(long(a) - long(b), result, status, code); +inline int jxr_neg(int value, thread bool &overflow) { + overflow |= value == INT_MIN; + return as_type(0u - as_type(value)); } -inline bool jxr_mul(int a, uint b, thread int &result, device atomic_uint *status, uint code) { - return jxr_narrow(long(a) * long(b), result, status, code); +inline int jxr_mul(int value, uint factor, thread bool &overflow) { + const bool negative = value < 0; + const uint magnitude = negative ? 0u - as_type(value) : as_type(value); + const uint product = magnitude * factor; + overflow |= mulhi(magnitude, factor) != 0u || product > (negative ? 0x80000000u : 0x7fffffffu); + return as_type(negative ? 0u - product : product); } -inline bool jxr_mul3_round(int value, int rounding, uint shift, thread int &result, - device atomic_uint *status, uint code) { - int product; - int rounded; - return jxr_narrow(long(value) * 3l, product, status, code) - && jxr_add(product, rounding, rounded, status, code) - && (result = rounded >> shift, true); +inline int jxr_mul3(int value, thread bool &overflow) { + overflow |= value > INT_MAX / 3 || value < INT_MIN / 3; + return as_type(as_type(value) * 3u); } -inline bool jxr_t2x2h(thread int v[4], int rounding, device atomic_uint *status, uint code) { - int difference; - int midpoint; - int old2 = v[2]; - return jxr_add(v[0], v[3], v[0], status, code) - && jxr_sub(v[1], v[2], v[1], status, code) - && jxr_sub(v[0], v[1], difference, status, code) - && jxr_add(difference, rounding, midpoint, status, code) - && (midpoint >>= 1, true) - && jxr_sub(midpoint, v[3], v[2], status, code) - && jxr_sub(midpoint, old2, v[3], status, code) - && jxr_sub(v[0], v[3], v[0], status, code) - && jxr_add(v[1], v[2], v[1], status, code); +inline int jxr_mul3_round(int value, int rounding, uint shift, thread bool &overflow) { + return jxr_add(jxr_mul3(value, overflow), rounding, overflow) >> shift; } -inline bool jxr_todd(thread int v[4], device atomic_uint *status, uint code) { - int temporary; - return jxr_add(v[1], v[3], v[1], status, code) - && jxr_sub(v[0], v[2], v[0], status, code) - && jxr_sub(v[3], v[1] >> 1, v[3], status, code) - && jxr_add(v[0], 1, temporary, status, code) - && jxr_add(v[2], temporary >> 1, v[2], status, code) - && jxr_mul3_round(v[1], 4, 3, temporary, status, code) - && jxr_sub(v[0], temporary, v[0], status, code) - && jxr_mul3_round(v[0], 4, 3, temporary, status, code) - && jxr_add(v[1], temporary, v[1], status, code) - && jxr_mul3_round(v[3], 4, 3, temporary, status, code) - && jxr_sub(v[2], temporary, v[2], status, code) - && jxr_mul3_round(v[2], 4, 3, temporary, status, code) - && jxr_add(v[3], temporary, v[3], status, code) - && jxr_add(v[1], 1, temporary, status, code) - && jxr_sub(v[2], temporary >> 1, v[2], status, code) - && jxr_add(v[0], 1, temporary, status, code) - && jxr_sub(temporary >> 1, v[3], v[3], status, code) - && jxr_add(v[1], v[2], v[1], status, code) - && jxr_sub(v[0], v[3], v[0], status, code); +inline void jxr_t2x2h(thread int v[4], int rounding, thread bool &overflow) { + const int old2 = v[2]; + v[0] = jxr_add(v[0], v[3], overflow); + v[1] = jxr_sub(v[1], v[2], overflow); + const int midpoint = jxr_add(jxr_sub(v[0], v[1], overflow), rounding, overflow) >> 1; + v[2] = jxr_sub(midpoint, v[3], overflow); + v[3] = jxr_sub(midpoint, old2, overflow); + v[0] = jxr_sub(v[0], v[3], overflow); + v[1] = jxr_add(v[1], v[2], overflow); } -inline bool jxr_todd_odd(thread int v[4], device atomic_uint *status, uint code) { - int temporary; - int first; - int second; - return jxr_add(v[3], v[0], v[3], status, code) - && jxr_sub(v[2], v[1], v[2], status, code) - && (first = v[3] >> 1, second = v[2] >> 1, true) - && jxr_sub(v[0], first, v[0], status, code) - && jxr_add(v[1], second, v[1], status, code) - && jxr_mul3_round(v[1], 3, 3, temporary, status, code) - && jxr_sub(v[0], temporary, v[0], status, code) - && jxr_mul3_round(v[0], 3, 2, temporary, status, code) - && jxr_add(v[1], temporary, v[1], status, code) - && jxr_mul3_round(v[1], 4, 3, temporary, status, code) - && jxr_sub(v[0], temporary, v[0], status, code) - && jxr_sub(v[1], second, v[1], status, code) - && jxr_add(v[0], first, v[0], status, code) - && jxr_add(v[2], v[1], v[2], status, code) - && jxr_sub(v[3], v[0], v[3], status, code) - && jxr_narrow(-long(v[1]), v[1], status, code) - && jxr_narrow(-long(v[2]), v[2], status, code); +inline void jxr_todd(thread int v[4], thread bool &overflow) { + v[1] = jxr_add(v[1], v[3], overflow); + v[0] = jxr_sub(v[0], v[2], overflow); + v[3] = jxr_sub(v[3], v[1] >> 1, overflow); + v[2] = jxr_add(v[2], jxr_add(v[0], 1, overflow) >> 1, overflow); + v[0] = jxr_sub(v[0], jxr_mul3_round(v[1], 4, 3, overflow), overflow); + v[1] = jxr_add(v[1], jxr_mul3_round(v[0], 4, 3, overflow), overflow); + v[2] = jxr_sub(v[2], jxr_mul3_round(v[3], 4, 3, overflow), overflow); + v[3] = jxr_add(v[3], jxr_mul3_round(v[2], 4, 3, overflow), overflow); + v[2] = jxr_sub(v[2], jxr_add(v[1], 1, overflow) >> 1, overflow); + v[3] = jxr_sub(jxr_add(v[0], 1, overflow) >> 1, v[3], overflow); + v[1] = jxr_add(v[1], v[2], overflow); + v[0] = jxr_sub(v[0], v[3], overflow); } -inline bool jxr_transform_group(thread int c[16], uint4 indices, uint kind, int rounding, - device atomic_uint *status, uint code) { +inline void jxr_todd_odd(thread int v[4], thread bool &overflow) { + v[3] = jxr_add(v[3], v[0], overflow); + v[2] = jxr_sub(v[2], v[1], overflow); + const int first = v[3] >> 1; + const int second = v[2] >> 1; + v[0] = jxr_sub(v[0], first, overflow); + v[1] = jxr_add(v[1], second, overflow); + v[0] = jxr_sub(v[0], jxr_mul3_round(v[1], 3, 3, overflow), overflow); + v[1] = jxr_add(v[1], jxr_mul3_round(v[0], 3, 2, overflow), overflow); + v[0] = jxr_sub(v[0], jxr_mul3_round(v[1], 4, 3, overflow), overflow); + v[1] = jxr_sub(v[1], second, overflow); + v[0] = jxr_add(v[0], first, overflow); + v[2] = jxr_add(v[2], v[1], overflow); + v[3] = jxr_sub(v[3], v[0], overflow); + v[1] = jxr_neg(v[1], overflow); + v[2] = jxr_neg(v[2], overflow); +} + +inline void jxr_transform_group(thread int c[16], uint4 indices, uint kind, int rounding, + thread bool &overflow) { int v[4] = { c[indices.x], c[indices.y], c[indices.z], c[indices.w] }; - bool ok = kind == 0u ? jxr_t2x2h(v, rounding, status, code) - : (kind == 1u ? jxr_todd(v, status, code) : jxr_todd_odd(v, status, code)); - if (!ok) return false; + if (kind == 0u) jxr_t2x2h(v, rounding, overflow); + else if (kind == 1u) jxr_todd(v, overflow); + else jxr_todd_odd(v, overflow); c[indices.x] = v[0]; c[indices.y] = v[1]; c[indices.z] = v[2]; c[indices.w] = v[3]; - return true; } -inline bool jxr_inverse_transform(thread int c[16], device atomic_uint *status, uint code) { +inline void jxr_inverse_transform(thread int c[16], thread bool &overflow) { int input[16]; for (uint i = 0; i < 16; ++i) input[i] = c[i]; for (uint i = 0; i < 16; ++i) c[JXR_INVERSE_PERMUTATION[i]] = input[i]; - return jxr_transform_group(c, uint4(0,1,4,5), 0, 1, status, code) - && jxr_transform_group(c, uint4(2,3,6,7), 1, 0, status, code) - && jxr_transform_group(c, uint4(8,12,9,13), 1, 0, status, code) - && jxr_transform_group(c, uint4(10,11,14,15), 2, 0, status, code) - && jxr_transform_group(c, uint4(0,3,12,15), 0, 0, status, code) - && jxr_transform_group(c, uint4(5,6,9,10), 0, 0, status, code) - && jxr_transform_group(c, uint4(1,2,13,14), 0, 0, status, code) - && jxr_transform_group(c, uint4(4,7,8,11), 0, 0, status, code); + jxr_transform_group(c, uint4(0,1,4,5), 0, 1, overflow); + jxr_transform_group(c, uint4(2,3,6,7), 1, 0, overflow); + jxr_transform_group(c, uint4(8,12,9,13), 1, 0, overflow); + jxr_transform_group(c, uint4(10,11,14,15), 2, 0, overflow); + jxr_transform_group(c, uint4(0,3,12,15), 0, 0, overflow); + jxr_transform_group(c, uint4(5,6,9,10), 0, 0, overflow); + jxr_transform_group(c, uint4(1,2,13,14), 0, 0, overflow); + jxr_transform_group(c, uint4(4,7,8,11), 0, 0, overflow); } diff --git a/crates/jxr-metal/src/kernels/first_transform.metal b/crates/jxr-metal/src/kernels/first_transform.metal index 9ca6000..7f5331f 100644 --- a/crates/jxr-metal/src/kernels/first_transform.metal +++ b/crates/jxr-metal/src/kernels/first_transform.metal @@ -9,42 +9,43 @@ inline void jxr_dequantize_first_transform_one( const uint metadata_index = plane.macroblock_offset + gid; const JxrMacroblockAbi metadata = macroblocks[metadata_index]; const uint block_count = plane.block_columns * plane.block_rows; + bool overflow = false; int low[16]; for (uint i = 0; i < 16; ++i) low[i] = 0; - if (!jxr_mul(packed[metadata.coefficient_offset], metadata.quantizer_dc, low[0], status, 1u)) return; + low[0] = jxr_mul(packed[metadata.coefficient_offset], metadata.quantizer_dc, overflow); if (metadata.bands != 0u) { + const uint stride = metadata.bands >= 2u ? 16u : 1u; for (uint block = 1; block < block_count; ++block) { - const uint source = metadata.bands >= 2u - ? metadata.coefficient_offset + block * 16u - : metadata.coefficient_offset + block; - if (!jxr_mul(packed[source], metadata.quantizer_low_pass, low[block], status, 1u)) return; + low[block] = jxr_mul(packed[metadata.coefficient_offset + block * stride], + metadata.quantizer_low_pass, overflow); } } if (plane.block_columns == 4u) { - if (!jxr_inverse_transform(low, status, 1u)) return; + jxr_inverse_transform(low, overflow); } else if (plane.block_rows == 2u) { int values[4] = { low[0], low[1], low[2], low[3] }; - if (!jxr_t2x2h(values, 0, status, 1u)) return; + jxr_t2x2h(values, 0, overflow); low[0] = values[0]; low[1] = values[2]; low[2] = values[1]; low[3] = values[3]; } else { - int pair0; - int pair1; - int temporary; - if (!jxr_add(low[4], 1, temporary, status, 1u) - || !jxr_sub(low[0], temporary >> 1, pair0, status, 1u) - || !jxr_add(low[4], pair0, pair1, status, 1u)) return; + const int pair0 = jxr_sub(low[0], jxr_add(low[4], 1, overflow) >> 1, overflow); + const int pair1 = jxr_add(low[4], pair0, overflow); low[0] = pair0; low[4] = pair1; int first[4] = { low[0], low[1], low[2], low[3] }; int second[4] = { low[4], low[6], low[5], low[7] }; - if (!jxr_t2x2h(first, 0, status, 1u) || !jxr_t2x2h(second, 0, status, 1u)) return; + jxr_t2x2h(first, 0, overflow); + jxr_t2x2h(second, 0, overflow); low[0] = first[0]; low[1] = first[2]; low[2] = first[1]; low[3] = first[3]; low[4] = second[0]; low[5] = second[1]; low[6] = second[2]; low[7] = second[3]; } if (plane.scale_after_first_transform != 0u) { for (uint block = 0; block < block_count; ++block) { - if (!jxr_narrow(long(low[block]) * 2l, low[block], status, 1u)) return; + low[block] = jxr_add(low[block], low[block], overflow); } } + if (overflow) { + jxr_fail(status, 1u); + return; + } const uint local_x = metadata.coded_x - plane.macroblock_origin_x; const uint local_y = metadata.coded_y - plane.macroblock_origin_y; const uint base_x = local_x * plane.block_columns; diff --git a/crates/jxr-metal/src/kernels/highpass_transform.metal b/crates/jxr-metal/src/kernels/highpass_transform.metal index 7e9234a..d22fd12 100644 --- a/crates/jxr-metal/src/kernels/highpass_transform.metal +++ b/crates/jxr-metal/src/kernels/highpass_transform.metal @@ -1,3 +1,17 @@ +// Accumulates one HP prediction chain in normative order: each block adds its +// raw coefficient to the already-predicted value of its predecessor, so every +// partial sum (and its overflow check) matches a serial macroblock traversal. +inline int jxr_predicted_hp(device const int *coefficients, uint first_block, uint last_block, + uint block_step, uint coefficient, thread bool &overflow) { + int value = coefficients[first_block * 16u + coefficient]; + for (uint block = first_block + block_step; block <= last_block; block += block_step) + value = jxr_add(coefficients[block * 16u + coefficient], value, overflow); + return value; +} + +// One thread reconstructs one 4x4 block. Consecutive threads cover the blocks +// of consecutive macroblocks, so SIMD groups stay full for every chroma +// layout and no threadgroup memory or barrier is needed. inline void jxr_highpass_second_transform_one( device const int *packed, device const JxrMacroblockAbi *macroblocks, @@ -5,69 +19,56 @@ inline void jxr_highpass_second_transform_one( device int *samples, device atomic_uint *status, JxrPlaneAbi plane, - threadgroup int *high, - uint group_id, - uint tid, - uint threads) { - if (group_id >= plane.macroblock_count || jxr_failed(status)) return; - const uint metadata_index = plane.macroblock_offset + group_id; - const JxrMacroblockAbi metadata = macroblocks[metadata_index]; + uint thread_index) { const uint block_count = plane.block_columns * plane.block_rows; - const uint coefficient_count = block_count * 16u; - for (uint index = tid; index < coefficient_count; index += threads) { - if (metadata.bands < 2u || (index & 15u) == 0u) { - high[index] = 0; - } else { - high[index] = packed[metadata.coefficient_offset + index]; + const uint macroblock = thread_index / block_count; + const uint block = thread_index - macroblock * block_count; + if (macroblock >= plane.macroblock_count || jxr_failed(status)) return; + const JxrMacroblockAbi metadata = macroblocks[plane.macroblock_offset + macroblock]; + const uint column = block % plane.block_columns; + const uint row = block / plane.block_columns; + int coefficients[16]; + for (uint coefficient = 0; coefficient < 16; ++coefficient) coefficients[coefficient] = 0; + if (metadata.bands >= 2u) { + device const int *high = packed + metadata.coefficient_offset; + for (uint coefficient = 1; coefficient < 16; ++coefficient) + coefficients[coefficient] = high[block * 16u + coefficient]; + bool prediction_overflow = false; + if (metadata.hp_prediction == 1u && column != 0u) { + for (uint coefficient = 4; coefficient <= 12; coefficient += 4) + coefficients[coefficient] = jxr_predicted_hp( + high, block - column, block, 1u, coefficient, prediction_overflow); + } else if (metadata.hp_prediction == 2u && row != 0u) { + for (uint coefficient = 1; coefficient <= 3; ++coefficient) + coefficients[coefficient] = jxr_predicted_hp( + high, column, block, plane.block_columns, coefficient, prediction_overflow); } - } - threadgroup_barrier(mem_flags::mem_threadgroup); - if (tid == 0u && metadata.bands >= 2u) { - if (metadata.hp_prediction == 1u) { - for (uint block = 1; block < block_count; ++block) { - if ((block % plane.block_columns) != 0u) { - for (uint coefficient = 4; coefficient <= 12; coefficient += 4) { - const uint destination = block * 16u + coefficient; - int predicted; - if (!jxr_add(high[destination], high[destination - 16u], - predicted, status, 3u)) break; - high[destination] = predicted; - } - } - } - } else if (metadata.hp_prediction == 2u) { - for (uint block = plane.block_columns; block < block_count; ++block) { - for (uint coefficient = 1; coefficient <= 3; ++coefficient) { - const uint destination = block * 16u + coefficient; - const uint source = destination - plane.block_columns * 16u; - int predicted; - if (!jxr_add(high[destination], high[source], predicted, status, 3u)) break; - high[destination] = predicted; - } - } + if (prediction_overflow) { + jxr_fail(status, 3u); + return; } } - threadgroup_barrier(mem_flags::mem_threadgroup | mem_flags::mem_device); - if (jxr_failed(status)) return; - for (uint block = tid; block < block_count; block += threads) { - int coefficients[16]; - coefficients[0] = 0; - for (uint coefficient = 1; coefficient < 16; ++coefficient) { - if (!jxr_mul(high[block * 16u + coefficient], metadata.quantizer_high_pass, - coefficients[coefficient], status, 4u)) return; - } - const uint local_x = metadata.coded_x - plane.macroblock_origin_x; - const uint local_y = metadata.coded_y - plane.macroblock_origin_y; - const uint low_x = local_x * plane.block_columns + block % plane.block_columns; - const uint low_y = local_y * plane.block_rows + block / plane.block_columns; - coefficients[0] = low_plane[plane.low_offset + low_y * plane.low_width + low_x]; - if (!jxr_inverse_transform(coefficients, status, 4u)) return; - const uint output_x = local_x * plane.block_columns * 4u + (block % plane.block_columns) * 4u; - const uint output_y = local_y * plane.block_rows * 4u + (block / plane.block_columns) * 4u; - for (uint row = 0; row < 4; ++row) - for (uint column = 0; column < 4; ++column) - samples[plane.sample_offset + (output_y + row) * plane.sample_width + output_x + column] = - coefficients[row * 4u + column]; + bool overflow = false; + for (uint coefficient = 1; coefficient < 16; ++coefficient) + coefficients[coefficient] = + jxr_mul(coefficients[coefficient], metadata.quantizer_high_pass, overflow); + const uint local_x = metadata.coded_x - plane.macroblock_origin_x; + const uint local_y = metadata.coded_y - plane.macroblock_origin_y; + const uint low_x = local_x * plane.block_columns + column; + const uint low_y = local_y * plane.block_rows + row; + coefficients[0] = low_plane[plane.low_offset + low_y * plane.low_width + low_x]; + jxr_inverse_transform(coefficients, overflow); + if (overflow) { + jxr_fail(status, 4u); + return; + } + const uint output_x = low_x * 4u; + const uint output_y = low_y * 4u; + for (uint y = 0; y < 4; ++y) { + device int4 *destination = reinterpret_cast( + samples + plane.sample_offset + (output_y + y) * plane.sample_width + output_x); + *destination = int4(coefficients[y * 4u], coefficients[y * 4u + 1u], + coefficients[y * 4u + 2u], coefficients[y * 4u + 3u]); } } @@ -78,12 +79,9 @@ kernel void jxr_highpass_second_transform( device int *samples [[buffer(3)]], device atomic_uint *status [[buffer(4)]], constant JxrPlaneAbi &plane [[buffer(5)]], - uint group_id [[threadgroup_position_in_grid]], - uint tid [[thread_index_in_threadgroup]], - uint threads [[threads_per_threadgroup]]) { - threadgroup int high[256]; + uint gid [[thread_position_in_grid]]) { jxr_highpass_second_transform_one( - packed, macroblocks, low_plane, samples, status, plane, high, group_id, tid, threads); + packed, macroblocks, low_plane, samples, status, plane, gid); } kernel void jxr_highpass_second_transform_batch( @@ -94,13 +92,9 @@ kernel void jxr_highpass_second_transform_batch( device atomic_uint *statuses [[buffer(4)]], device const JxrPlaneAbi *planes [[buffer(5)]], constant JxrBatchDispatchAbi &batch [[buffer(6)]], - uint3 group_id [[threadgroup_position_in_grid]], - uint3 tid [[thread_position_in_threadgroup]], - uint3 threads [[threads_per_threadgroup]]) { - if (group_id.y >= batch.image_count || group_id.z >= batch.plane_count) return; - threadgroup int high[256]; - const JxrPlaneAbi plane = planes[group_id.y * batch.plane_count + group_id.z]; + uint3 gid [[thread_position_in_grid]]) { + if (gid.y >= batch.image_count || gid.z >= batch.plane_count) return; + const JxrPlaneAbi plane = planes[gid.y * batch.plane_count + gid.z]; jxr_highpass_second_transform_one( - packed, macroblocks, low_plane, samples, statuses + group_id.y, - plane, high, group_id.x, tid.x, threads.x); + packed, macroblocks, low_plane, samples, statuses + gid.y, plane, gid.x); } diff --git a/crates/jxr-metal/src/kernels/output.metal b/crates/jxr-metal/src/kernels/output.metal index 480b99b..e146b88 100644 --- a/crates/jxr-metal/src/kernels/output.metal +++ b/crates/jxr-metal/src/kernels/output.metal @@ -12,17 +12,19 @@ inline int4 jxr_centering(uint code) { } } -inline bool jxr_weighted(int first, int fw, int second, int sw, thread int &result, - device atomic_uint *status) { - long value = long(first) * long(fw) + long(second) * long(sw) + 4l; - return jxr_narrow(value >> 3, result, status, 16u); +// Exact `(first * fw + second * sw + 4) >> 3` for non-negative weights that +// sum to eight. Splitting each operand into `8q + r` keeps every intermediate +// inside `int`, and the weighted average of two `int` values always fits. +inline int jxr_weighted(int first, int fw, int second, int sw) { + return fw * (first >> 3) + sw * (second >> 3) + + ((fw * (first & 7) + sw * (second & 7) + 4) >> 3); } -inline bool jxr_upsample_pair(int previous, int current, int next, uint centering, - thread int pair[2], device atomic_uint *status) { +inline void jxr_upsample_pair(int previous, int current, int next, uint centering, + thread int pair[2]) { int4 h = jxr_centering(centering); - return jxr_weighted(previous, h.z, current, h.w, pair[0], status) - && jxr_weighted(current, h.x, next, h.y, pair[1], status); + pair[0] = jxr_weighted(previous, h.z, current, h.w); + pair[1] = jxr_weighted(current, h.x, next, h.y); } inline int jxr_clamped_plane(device const int *samples, JxrSamplePlaneAbi plane, int x, int y) { @@ -31,94 +33,85 @@ inline int jxr_clamped_plane(device const int *samples, JxrSamplePlaneAbi plane, return samples[plane.sample_offset + uint(local_y) * plane.width + uint(local_x)]; } -inline bool jxr_chroma_sample(device const int *samples, JxrSamplePlaneAbi plane, - uint full_x, uint full_y, constant JxrOutputAbi ¶ms, - thread int &result, device atomic_uint *status) { - if (params.chroma_sampling == 3u) { - result = jxr_read_plane(samples, plane, full_x, full_y); - return true; - } +inline int jxr_chroma_sample(device const int *samples, JxrSamplePlaneAbi plane, + uint full_x, uint full_y, constant JxrOutputAbi ¶ms) { + if (params.chroma_sampling == 3u) return jxr_read_plane(samples, plane, full_x, full_y); int chroma_x = int(full_x >> 1u); + int pair[2]; if (params.chroma_sampling == 2u) { - int pair[2]; - if (!jxr_upsample_pair( + jxr_upsample_pair( jxr_clamped_plane(samples, plane, chroma_x - 1, int(full_y)), jxr_clamped_plane(samples, plane, chroma_x, int(full_y)), jxr_clamped_plane(samples, plane, chroma_x + 1, int(full_y)), - params.chroma_centering_x, pair, status)) return false; - result = pair[full_x & 1u]; - return true; + params.chroma_centering_x, pair); + return pair[full_x & 1u]; } int chroma_y = int(full_y >> 1u); int vertical[3]; for (int column = -1; column <= 1; ++column) { - int pair[2]; - if (!jxr_upsample_pair( + jxr_upsample_pair( jxr_clamped_plane(samples, plane, chroma_x + column, chroma_y - 1), jxr_clamped_plane(samples, plane, chroma_x + column, chroma_y), jxr_clamped_plane(samples, plane, chroma_x + column, chroma_y + 1), - params.chroma_centering_y, pair, status)) return false; + params.chroma_centering_y, pair); vertical[column + 1] = pair[full_y & 1u]; } - int horizontal[2]; - if (!jxr_upsample_pair(vertical[0], vertical[1], vertical[2], - params.chroma_centering_x, horizontal, status)) return false; - result = horizontal[full_x & 1u]; - return true; + jxr_upsample_pair(vertical[0], vertical[1], vertical[2], params.chroma_centering_x, pair); + return pair[full_x & 1u]; } -inline bool jxr_primary_values(device const int *samples, device const JxrSamplePlaneAbi *planes, - uint x, uint y, constant JxrOutputAbi ¶ms, - thread int values[4], device atomic_uint *status) { +inline void jxr_primary_values(device const int *samples, device const JxrSamplePlaneAbi *planes, + uint x, uint y, constant JxrOutputAbi ¶ms, thread int values[4]) { for (uint i = 0; i < 4; ++i) values[i] = 0; values[0] = jxr_read_plane(samples, planes[0], x, y); - if (params.component_count == 1u) return true; + if (params.component_count == 1u) return; if (params.internal_color == 1u) { - if (!jxr_chroma_sample(samples, planes[1], x, y, params, values[1], status) - || !jxr_chroma_sample(samples, planes[2], x, y, params, values[2], status)) return false; + values[1] = jxr_chroma_sample(samples, planes[1], x, y, params); + values[2] = jxr_chroma_sample(samples, planes[2], x, y, params); } else { for (uint i = 1; i < min(params.component_count, 4u); ++i) values[i] = jxr_read_plane(samples, planes[i], x, y); } - return true; } -inline bool jxr_converted(device const int *samples, device const JxrSamplePlaneAbi *planes, +inline void jxr_converted(device const int *samples, device const JxrSamplePlaneAbi *planes, uint x, uint y, constant JxrOutputAbi ¶ms, - thread int values[4], device atomic_uint *status) { + thread int values[4], thread bool &overflow) { int input[4]; - if (!jxr_primary_values(samples, planes, x, y, params, input, status)) return false; + jxr_primary_values(samples, planes, x, y, params, input); for (uint i = 0; i < 4; ++i) values[i] = input[i]; if (params.internal_color == 0u && params.output_color == 2u) { values[1] = input[0]; values[2] = input[0]; } else if (params.internal_color == 1u && (params.output_color == 2u || params.output_color == 6u)) { - int temporary; - int green; - int red; - if (!jxr_sub(0, input[1], temporary, status, 16u) - || !jxr_sub(input[0], temporary >> 1, green, status, 16u) - || !jxr_add(temporary, green, red, status, 16u) - || !jxr_sub(red, (input[2] >> 1) + (input[2] & 1), red, status, 16u) - || !jxr_add(input[2], red, values[2], status, 16u)) return false; + const int temporary = jxr_sub(0, input[1], overflow); + const int green = jxr_sub(input[0], temporary >> 1, overflow); + int red = jxr_add(temporary, green, overflow); + red = jxr_sub(red, (input[2] >> 1) + (input[2] & 1), overflow); + values[2] = jxr_add(input[2], red, overflow); values[0] = red; values[1] = green; if (params.bit_depth >= 8u && params.bit_depth <= 10u && params.red_blue_not_swapped == 0u) { int swap = values[0]; values[0] = values[2]; values[2] = swap; } } else if (params.internal_color == 5u && params.output_color == 3u) { - int black; - int magenta; - int cyan; - if (!jxr_add(input[3], input[0] >> 1, black, status, 16u) - || !jxr_sub(black, input[0], magenta, status, 16u) - || !jxr_sub(magenta, input[1] >> 1, magenta, status, 16u) - || !jxr_add(input[1], magenta, cyan, status, 16u) - || !jxr_add(cyan, input[2] >> 1, cyan, status, 16u) - || !jxr_sub(cyan, input[2], values[2], status, 16u)) return false; + const int black = jxr_add(input[3], input[0] >> 1, overflow); + int magenta = jxr_sub(black, input[0], overflow); + magenta = jxr_sub(magenta, input[1] >> 1, overflow); + int cyan = jxr_add(input[1], magenta, overflow); + cyan = jxr_add(cyan, input[2] >> 1, overflow); + values[2] = jxr_sub(cyan, input[2], overflow); values[0] = cyan; values[1] = magenta; values[3] = black; } else if (params.internal_color == 5u && params.output_color == 4u) { values[0] = input[1]; values[1] = input[2]; values[2] = input[3]; values[3] = input[0]; } - return true; +} + +// Loads and color-converts one pixel. Every channel store selects from this +// result instead of repeating chroma upsampling and conversion per channel. +inline void jxr_load_pixel(device const int *samples, device const JxrSamplePlaneAbi *planes, + uint x, uint y, constant JxrOutputAbi ¶ms, + thread int converted[4], thread bool &overflow) { + for (uint i = 0; i < 4; ++i) converted[i] = 0; + if (params.output_color != 7u) jxr_converted(samples, planes, x, y, params, converted, overflow); } inline int jxr_base_bias(uint depth, uint shift_bits) { @@ -132,8 +125,19 @@ inline int jxr_base_bias(uint depth, uint shift_bits) { } } -inline bool jxr_scale(int sample, uint component, bool alpha, constant JxrOutputAbi ¶ms, - thread int &result, device atomic_uint *status) { +// Matches the CPU `checked_shl` plus round-trip test. +inline int jxr_shl(int value, uint bits, thread bool &overflow) { + if (bits >= 32u) { + overflow = true; + return 0; + } + const int shifted = as_type(as_type(value) << bits); + overflow |= (shifted >> bits) != value; + return shifted; +} + +inline int jxr_scale(int sample, uint component, bool alpha, constant JxrOutputAbi ¶ms, + thread bool &overflow) { uint depth = params.bit_depth; uint scaled = alpha ? params.alpha_scaled : params.scaled; uint shift_bits = alpha ? params.alpha_shift_bits : params.shift_bits; @@ -141,34 +145,27 @@ inline bool jxr_scale(int sample, uint component, bool alpha, constant JxrOutput if (!alpha && params.output_color == 3u) bias = component < 3u ? bias >> 1 : -(bias >> 1); if (!alpha && params.output_color == 6u) bias = 0; uint scale = scaled * 3u; - int shifted_bias; - if (!jxr_narrow(long(bias) << scale, shifted_bias, status, 16u)) return false; - int biased; - if (!jxr_add(sample, shifted_bias, biased, status, 16u)) return false; + // |bias| <= 32768 and scale <= 3, so the scaled bias cannot overflow. + int biased = jxr_add(sample, bias * int(1u << scale), overflow); int rounding = scaled == 0u ? 0 : ((depth == 0u || depth == 3u) ? 4 : 3); - int rounded; - if (!jxr_add(biased, rounding, rounded, status, 16u)) return false; + int rounded = jxr_add(biased, rounding, overflow); uint downshift = scale + ((depth == 10u && component != 1u) ? 1u : 0u); int downscaled = rounded >> downshift; - if (depth == 3u || depth == 4u || depth == 6u) - return jxr_narrow(long(downscaled) << shift_bits, result, status, 16u); - result = downscaled; - return true; + if (depth == 3u || depth == 4u || depth == 6u) return jxr_shl(downscaled, shift_bits, overflow); + return downscaled; } inline uint jxr_f16_bits(int sample, bool alpha, constant JxrOutputAbi ¶ms, - device atomic_uint *status) { - int scaled; - if (!jxr_scale(sample, 0u, alpha, params, scaled, status)) return 0u; + thread bool &overflow) { + int scaled = jxr_scale(sample, 0u, alpha, params, overflow); uint sign = scaled < 0 ? 0x8000u : 0u; - uint magnitude = uint(min(abs(long(scaled)), 32767l)); - return sign | magnitude; + uint magnitude = scaled < 0 ? 0u - as_type(scaled) : as_type(scaled); + return sign | min(magnitude, 32767u); } inline uint jxr_f32_bits(int sample, bool alpha, constant JxrOutputAbi ¶ms, - device atomic_uint *status) { - int scaled; - if (!jxr_scale(sample, 0u, alpha, params, scaled, status)) return 0u; + thread bool &overflow) { + int scaled = jxr_scale(sample, 0u, alpha, params, overflow); uint length = alpha ? params.alpha_mantissa_length : params.mantissa_length; int exponent_bias = as_type(alpha ? params.alpha_exponent_bias_bits : params.exponent_bias_bits); uint sign = scaled < 0 ? 0x80000000u : 0u; @@ -182,13 +179,15 @@ inline uint jxr_f32_bits(int sample, bool alpha, constant JxrOutputAbi ¶ms, if (mantissa < implicit) exponent = 0l; else mantissa ^= implicit; mantissa <<= 23u - length; if (exponent < 0l || exponent > 255l || mantissa > 0x7ffffful) { - jxr_fail(status, 16u); return 0u; + overflow = true; + return 0u; } return sign | (uint(exponent) << 23u) | uint(mantissa); } +// Callers pass `maximum <= 65535`, so the rounded product fits in 32 bits. inline uint jxr_unsigned_premultiply(uint value, uint alpha, uint maximum) { - return uint((ulong(min(value, maximum)) * ulong(min(alpha, maximum)) + ulong(maximum / 2u)) / ulong(maximum)); + return (min(value, maximum) * min(alpha, maximum) + maximum / 2u) / maximum; } inline int jxr_signed_premultiply(int value, int alpha, int maximum) { @@ -197,59 +196,64 @@ inline int jxr_signed_premultiply(int value, int alpha, int maximum) { return int(value < 0 ? -result : result); } -inline bool jxr_component(device const int *samples, device const JxrSamplePlaneAbi *planes, - uint channel, uint x, uint y, constant JxrOutputAbi ¶ms, - thread int &sample, thread bool &alpha, device atomic_uint *status) { +inline bool jxr_is_padding(uint channel, constant JxrOutputAbi ¶ms) { + return (params.channel_layout == 12u || params.channel_layout == 13u) && channel == 3u; +} + +inline int jxr_channel_sample(device const int *samples, device const JxrSamplePlaneAbi *planes, + uint channel, uint x, uint y, constant JxrOutputAbi ¶ms, + thread const int converted[4], thread bool &alpha) { uint primary_count = params.alpha_plane == UINT_MAX ? params.channels : params.channels - 1u; alpha = params.alpha_plane != UINT_MAX && channel == primary_count; - if (alpha) { - sample = jxr_read_plane(samples, planes[params.alpha_plane], x, y); - return true; - } - if (params.output_color == 7u) { - sample = jxr_read_plane(samples, planes[channel], x, y); - return true; - } - if ((params.channel_layout == 12u || params.channel_layout == 13u) && channel == 3u) { - sample = 0; - alpha = false; - return true; - } - int values[4]; - if (!jxr_converted(samples, planes, x, y, params, values, status)) return false; + if (alpha) return jxr_read_plane(samples, planes[params.alpha_plane], x, y); + if (params.output_color == 7u) return jxr_read_plane(samples, planes[channel], x, y); + if (jxr_is_padding(channel, params)) return 0; uint source_channel = channel; if ((params.channel_layout == 6u || params.channel_layout == 7u || params.channel_layout == 13u) && channel < 3u) source_channel = 2u - channel; - sample = values[source_channel]; - return true; + return converted[source_channel]; } -inline bool jxr_formatted_integer(device const int *samples, device const JxrSamplePlaneAbi *planes, - uint channel, uint x, uint y, constant JxrOutputAbi ¶ms, - thread int &value, device atomic_uint *status) { - if ((params.channel_layout == 12u || params.channel_layout == 13u) && channel == 3u) { - value = 0; - return true; - } - int sample; bool alpha; - if (!jxr_component(samples, planes, channel, x, y, params, sample, alpha, status) - || !jxr_scale(sample, channel, alpha, params, value, status)) return false; +// Scaled alpha used to premultiply integer color channels, or zero when unused. +inline int jxr_premultiply_alpha(device const int *samples, device const JxrSamplePlaneAbi *planes, + uint x, uint y, constant JxrOutputAbi ¶ms, + thread bool &overflow) { + if (params.premultiply_alpha == 0u || params.alpha_plane == UINT_MAX) return 0; + return jxr_scale(jxr_read_plane(samples, planes[params.alpha_plane], x, y), 0u, true, params, overflow); +} + +inline int jxr_formatted_integer(device const int *samples, device const JxrSamplePlaneAbi *planes, + uint channel, uint x, uint y, constant JxrOutputAbi ¶ms, + thread const int converted[4], int alpha_value, + thread bool &overflow) { + if (jxr_is_padding(channel, params)) return 0; + bool alpha; + int sample = jxr_channel_sample(samples, planes, channel, x, y, params, converted, alpha); + int value = jxr_scale(sample, channel, alpha, params, overflow); if (params.premultiply_alpha != 0u && !alpha && params.alpha_plane != UINT_MAX) { - int alpha_sample = jxr_read_plane(samples, planes[params.alpha_plane], x, y); - int alpha_value; - if (!jxr_scale(alpha_sample, 0u, true, params, alpha_value, status)) return false; if (params.bit_depth == 1u) value = int(jxr_unsigned_premultiply(uint(clamp(value,0,255)), uint(clamp(alpha_value,0,255)), 255u)); else if (params.bit_depth == 2u || params.bit_depth == 3u) value = int(jxr_unsigned_premultiply(uint(clamp(value,0,65535)), uint(clamp(alpha_value,0,65535)), 65535u)); else if (params.bit_depth == 4u) value = jxr_signed_premultiply(value, alpha_value, 32767); else if (params.bit_depth == 6u) value = jxr_signed_premultiply(value, alpha_value, INT_MAX); } - return true; + return value; } inline uint jxr_output_index(JxrSurfacePlaneAbi surface, uint x, uint y, uint channel, uint bytes) { return surface.byte_offset + y * surface.row_stride_bytes + (x * surface.channels + channel) * bytes; } +// Planar surfaces store one scaled sample per thread from their own plane. +inline int jxr_planar_value(device const int *samples, device const JxrSamplePlaneAbi *planes, + uint2 gid, constant JxrOutputAbi ¶ms, thread bool &overflow) { + uint source_plane = params.output_plane < 3u ? params.output_plane : params.alpha_plane; + bool chroma = params.output_plane == 1u || params.output_plane == 2u; + uint x = params.crop_x / (chroma ? 2u : 1u) + gid.x; + uint y = params.crop_y / ((params.chroma_sampling == 1u && chroma) ? 2u : 1u) + gid.y; + return jxr_scale(jxr_read_plane(samples, planes[source_plane], x, y), params.output_plane, + params.output_plane >= 3u, params, overflow); +} + kernel void jxr_output_u8(device const int *samples [[buffer(0)]], device const JxrSamplePlaneAbi *planes [[buffer(1)]], device const JxrSurfacePlaneAbi *surfaces [[buffer(2)]], @@ -257,22 +261,23 @@ kernel void jxr_output_u8(device const int *samples [[buffer(0)]], constant JxrOutputAbi ¶ms [[buffer(5)]], uint2 gid [[thread_position_in_grid]]) { JxrSurfacePlaneAbi surface = surfaces[params.output_plane]; if (gid.x >= surface.width || gid.y >= surface.height || jxr_failed(status)) return; + bool overflow = false; if (params.output_plane_count > 1u) { - uint source_plane = params.output_plane < 3u ? params.output_plane : params.alpha_plane; - uint x = params.crop_x / (params.output_plane == 1u || params.output_plane == 2u ? 2u : 1u) + gid.x; - uint y_divisor = (params.chroma_sampling == 1u && (params.output_plane == 1u || params.output_plane == 2u)) ? 2u : 1u; - uint y = params.crop_y / y_divisor + gid.y; - int scaled; - bool alpha = params.output_plane >= 3u; - if (!jxr_scale(jxr_read_plane(samples, planes[source_plane], x, y), params.output_plane, alpha, params, scaled, status)) return; - output[jxr_output_index(surface, gid.x, gid.y, 0u, 1u)] = uchar(clamp(scaled, 0, 255)); + int value = jxr_planar_value(samples, planes, gid, params, overflow); + if (overflow) { jxr_fail(status, 16u); return; } + output[jxr_output_index(surface, gid.x, gid.y, 0u, 1u)] = uchar(clamp(value, 0, 255)); return; } uint x = params.crop_x + gid.x, y = params.crop_y + gid.y; + int converted[4]; + jxr_load_pixel(samples, planes, x, y, params, converted, overflow); + int alpha_value = jxr_premultiply_alpha(samples, planes, x, y, params, overflow); for (uint channel = 0; channel < params.channels; ++channel) { - int value; if (!jxr_formatted_integer(samples, planes, channel, x, y, params, value, status)) return; - output[jxr_output_index(surface, gid.x, gid.y, channel, 1u)] = uchar(clamp(value,0,255)); + int value = jxr_formatted_integer(samples, planes, channel, x, y, params, converted, + alpha_value, overflow); + output[jxr_output_index(surface, gid.x, gid.y, channel, 1u)] = uchar(clamp(value, 0, 255)); } + if (overflow) jxr_fail(status, 16u); } #define JXR_INTEGER_STORE(NAME, TYPE, BYTES, MINIMUM, MAXIMUM) \ @@ -282,23 +287,25 @@ kernel void NAME(device const int *samples [[buffer(0)]], device const JxrSample uint2 gid [[thread_position_in_grid]]) { \ JxrSurfacePlaneAbi surface = surfaces[params.output_plane]; \ if (gid.x >= surface.width || gid.y >= surface.height || jxr_failed(status)) return; \ - uint x = params.crop_x + gid.x, y = params.crop_y + gid.y; \ - uint channels = surface.channels; \ - uint source_plane = params.output_plane < 3u ? params.output_plane : params.alpha_plane; \ + bool overflow = false; \ if (params.output_plane_count > 1u) { \ - x = params.crop_x / (params.output_plane == 1u || params.output_plane == 2u ? 2u : 1u) + gid.x; \ - uint yd = (params.chroma_sampling == 1u && (params.output_plane == 1u || params.output_plane == 2u)) ? 2u : 1u; \ - y = params.crop_y / yd + gid.y; channels = 1u; \ + int value = jxr_planar_value(samples, planes, gid, params, overflow); \ + if (overflow) { jxr_fail(status, 16u); return; } \ + *reinterpret_cast(output + jxr_output_index(surface, gid.x, gid.y, 0u, BYTES)) = \ + TYPE(clamp(value, MINIMUM, MAXIMUM)); \ + return; \ } \ - for (uint channel = 0; channel < channels; ++channel) { \ - int value; \ - if (params.output_plane_count > 1u) { \ - bool alpha = params.output_plane >= 3u; \ - if (!jxr_scale(jxr_read_plane(samples, planes[source_plane], x, y), params.output_plane, alpha, params, value, status)) return; \ - } else if (!jxr_formatted_integer(samples, planes, channel, x, y, params, value, status)) return; \ - device TYPE *destination = reinterpret_cast(output + jxr_output_index(surface, gid.x, gid.y, channel, BYTES)); \ - *destination = TYPE(clamp(value, MINIMUM, MAXIMUM)); \ + uint x = params.crop_x + gid.x, y = params.crop_y + gid.y; \ + int converted[4]; \ + jxr_load_pixel(samples, planes, x, y, params, converted, overflow); \ + int alpha_value = jxr_premultiply_alpha(samples, planes, x, y, params, overflow); \ + for (uint channel = 0; channel < surface.channels; ++channel) { \ + int value = jxr_formatted_integer(samples, planes, channel, x, y, params, converted, \ + alpha_value, overflow); \ + *reinterpret_cast(output + jxr_output_index(surface, gid.x, gid.y, channel, BYTES)) = \ + TYPE(clamp(value, MINIMUM, MAXIMUM)); \ } \ + if (overflow) jxr_fail(status, 16u); \ } JXR_INTEGER_STORE(jxr_output_u16, ushort, 2u, 0, (params.bit_depth == 2u ? 1023 : 65535)) @@ -311,60 +318,120 @@ kernel void jxr_output_f16(device const int *samples [[buffer(0)]], device const uint2 gid [[thread_position_in_grid]]) { JxrSurfacePlaneAbi surface = surfaces[0]; if (gid.x >= surface.width || gid.y >= surface.height || jxr_failed(status)) return; - uint x=params.crop_x+gid.x,y=params.crop_y+gid.y; - uint alpha_bits = params.alpha_plane == UINT_MAX ? 0u : jxr_f16_bits(jxr_read_plane(samples,planes[params.alpha_plane],x,y),true,params,status); - for(uint c=0;c(output+jxr_output_index(surface,gid.x,gid.y,c,2u))=0;continue;} int sample; bool alpha; if(!jxr_component(samples,planes,c,x,y,params,sample,alpha,status))return; - uint bits=jxr_f16_bits(sample,alpha,params,status); - if(params.premultiply_alpha!=0u&&!alpha){ uint sign=bits&0x8000u; bits=sign|jxr_unsigned_premultiply(bits&0x7fffu,(alpha_bits&0x8000u)==0u?alpha_bits&0x7fffu:0u,0x7fffu); } - *reinterpret_cast(output+jxr_output_index(surface,gid.x,gid.y,c,2u))=ushort(bits); } + uint x = params.crop_x + gid.x, y = params.crop_y + gid.y; + bool overflow = false; + uint alpha_bits = params.alpha_plane == UINT_MAX + ? 0u : jxr_f16_bits(jxr_read_plane(samples, planes[params.alpha_plane], x, y), true, params, overflow); + int converted[4]; + jxr_load_pixel(samples, planes, x, y, params, converted, overflow); + for (uint c = 0; c < params.channels; ++c) { + device ushort *destination = reinterpret_cast(output + jxr_output_index(surface, gid.x, gid.y, c, 2u)); + if (jxr_is_padding(c, params)) { *destination = 0; continue; } + bool alpha; + uint bits = jxr_f16_bits(jxr_channel_sample(samples, planes, c, x, y, params, converted, alpha), alpha, params, overflow); + if (params.premultiply_alpha != 0u && !alpha) { + uint sign = bits & 0x8000u; + bits = sign | jxr_unsigned_premultiply(bits & 0x7fffu, (alpha_bits & 0x8000u) == 0u ? alpha_bits & 0x7fffu : 0u, 0x7fffu); + } + *destination = ushort(bits); + } + if (overflow) jxr_fail(status, 16u); } kernel void jxr_output_f32(device const int *samples [[buffer(0)]], device const JxrSamplePlaneAbi *planes [[buffer(1)]], device const JxrSurfacePlaneAbi *surfaces [[buffer(2)]], device uchar *output [[buffer(3)]], device atomic_uint *status [[buffer(4)]], constant JxrOutputAbi ¶ms [[buffer(5)]], uint2 gid [[thread_position_in_grid]]) { - JxrSurfacePlaneAbi surface=surfaces[0]; if(gid.x>=surface.width||gid.y>=surface.height||jxr_failed(status))return; - uint x=params.crop_x+gid.x,y=params.crop_y+gid.y; float alpha_value=1.0f; - if(params.alpha_plane!=UINT_MAX) alpha_value=clamp(as_type(jxr_f32_bits(jxr_read_plane(samples,planes[params.alpha_plane],x,y),true,params,status)),0.0f,1.0f); - for(uint c=0;c(output+jxr_output_index(surface,gid.x,gid.y,c,4u))=0.0f;continue;}int sample;bool alpha;if(!jxr_component(samples,planes,c,x,y,params,sample,alpha,status))return; - uint bits=jxr_f32_bits(sample,alpha,params,status);float value=as_type(bits);if(params.premultiply_alpha!=0u&&!alpha)value*=alpha_value; - *reinterpret_cast(output+jxr_output_index(surface,gid.x,gid.y,c,4u))=value;} + JxrSurfacePlaneAbi surface = surfaces[0]; + if (gid.x >= surface.width || gid.y >= surface.height || jxr_failed(status)) return; + uint x = params.crop_x + gid.x, y = params.crop_y + gid.y; + bool overflow = false; + float alpha_value = 1.0f; + if (params.alpha_plane != UINT_MAX) + alpha_value = clamp(as_type(jxr_f32_bits(jxr_read_plane(samples, planes[params.alpha_plane], x, y), true, params, overflow)), 0.0f, 1.0f); + int converted[4]; + jxr_load_pixel(samples, planes, x, y, params, converted, overflow); + for (uint c = 0; c < params.channels; ++c) { + device float *destination = reinterpret_cast(output + jxr_output_index(surface, gid.x, gid.y, c, 4u)); + if (jxr_is_padding(c, params)) { *destination = 0.0f; continue; } + bool alpha; + float value = as_type(jxr_f32_bits(jxr_channel_sample(samples, planes, c, x, y, params, converted, alpha), alpha, params, overflow)); + if (params.premultiply_alpha != 0u && !alpha) value *= alpha_value; + *destination = value; + } + if (overflow) jxr_fail(status, 16u); } kernel void jxr_output_bits(device const int *samples [[buffer(0)]], device const JxrSamplePlaneAbi *planes [[buffer(1)]], device const JxrSurfacePlaneAbi *surfaces [[buffer(2)]], device uchar *output [[buffer(3)]], device atomic_uint *status [[buffer(4)]], constant JxrOutputAbi ¶ms [[buffer(5)]], uint2 gid [[thread_position_in_grid]]) { - JxrSurfacePlaneAbi surface=surfaces[0]; uint row_bytes=(surface.width+7u)/8u; - if(gid.x>=row_bytes||gid.y>=surface.height||jxr_failed(status))return; uchar byte=0; - for(uint bit=0;bit<8u;++bit){uint px=gid.x*8u+bit;if(px>=surface.width)break;int value; - if(!jxr_scale(jxr_read_plane(samples,planes[0],params.crop_x+px,params.crop_y+gid.y),0u,false,params,value,status))return; - uint packed=uint(clamp(value,0,1));if(params.bit_black!=0u)packed=1u-packed;byte|=uchar(packed<<(7u-bit));} - output[surface.byte_offset+gid.y*surface.row_stride_bytes+gid.x]=byte; + JxrSurfacePlaneAbi surface = surfaces[0]; + uint row_bytes = (surface.width + 7u) / 8u; + if (gid.x >= row_bytes || gid.y >= surface.height || jxr_failed(status)) return; + bool overflow = false; + uchar byte = 0; + for (uint bit = 0; bit < 8u; ++bit) { + uint px = gid.x * 8u + bit; + if (px >= surface.width) break; + int value = jxr_scale(jxr_read_plane(samples, planes[0], params.crop_x + px, params.crop_y + gid.y), 0u, false, params, overflow); + uint packed = uint(clamp(value, 0, 1)); + if (params.bit_black != 0u) packed = 1u - packed; + byte |= uchar(packed << (7u - bit)); + } + if (overflow) { jxr_fail(status, 16u); return; } + output[surface.byte_offset + gid.y * surface.row_stride_bytes + gid.x] = byte; } kernel void jxr_output_packed16(device const int *samples [[buffer(0)]], device const JxrSamplePlaneAbi *planes [[buffer(1)]], device const JxrSurfacePlaneAbi *surfaces [[buffer(2)]], device uchar *output [[buffer(3)]], device atomic_uint *status [[buffer(4)]], constant JxrOutputAbi ¶ms [[buffer(5)]], uint2 gid [[thread_position_in_grid]]) { - JxrSurfacePlaneAbi surface=surfaces[0];if(gid.x>=surface.width||gid.y>=surface.height||jxr_failed(status))return; - int values[4];if(!jxr_converted(samples,planes,params.crop_x+gid.x,params.crop_y+gid.y,params,values,status))return;uint packed=0u; - for(uint c=0;c<3u;++c){int value;if(!jxr_scale(values[c],c,false,params,value,status))return;uint maximum=params.bit_depth==10u&&c==1u?63u:31u; - uint shift=params.bit_depth==10u?(c==0u?11u:(c==1u?5u:0u)):(2u-c)*5u;packed|=uint(clamp(value,0,int(maximum)))<(output+surface.byte_offset+gid.y*surface.row_stride_bytes+gid.x*2u)=ushort(packed); + JxrSurfacePlaneAbi surface = surfaces[0]; + if (gid.x >= surface.width || gid.y >= surface.height || jxr_failed(status)) return; + bool overflow = false; + int values[4]; + jxr_converted(samples, planes, params.crop_x + gid.x, params.crop_y + gid.y, params, values, overflow); + uint packed = 0u; + for (uint c = 0; c < 3u; ++c) { + int value = jxr_scale(values[c], c, false, params, overflow); + uint maximum = params.bit_depth == 10u && c == 1u ? 63u : 31u; + uint shift = params.bit_depth == 10u ? (c == 0u ? 11u : (c == 1u ? 5u : 0u)) : (2u - c) * 5u; + packed |= uint(clamp(value, 0, int(maximum))) << shift; + } + if (overflow) { jxr_fail(status, 16u); return; } + *reinterpret_cast(output + surface.byte_offset + gid.y * surface.row_stride_bytes + gid.x * 2u) = ushort(packed); } kernel void jxr_output_packed32(device const int *samples [[buffer(0)]], device const JxrSamplePlaneAbi *planes [[buffer(1)]], device const JxrSurfacePlaneAbi *surfaces [[buffer(2)]], device uchar *output [[buffer(3)]], device atomic_uint *status [[buffer(4)]], constant JxrOutputAbi ¶ms [[buffer(5)]], uint2 gid [[thread_position_in_grid]]) { - JxrSurfacePlaneAbi surface=surfaces[0];if(gid.x>=surface.width||gid.y>=surface.height||jxr_failed(status))return;int values[4]; - if(!jxr_converted(samples,planes,params.crop_x+gid.x,params.crop_y+gid.y,params,values,status))return;uint packed=0u; - if(params.output_color==6u){int scaled[3];uint exponent=0u;uint mantissa[3];uint local_exp[3]; - for(uint c=0;c<3u;++c){if(!jxr_scale(values[c],c,false,params,scaled[c],status))return;if(scaled[c]<=0){mantissa[c]=0;local_exp[c]=0;} - else if((scaled[c]>>7)>1){mantissa[c]=uint((scaled[c]&127)+128);local_exp[c]=uint(scaled[c]>>7);}else{mantissa[c]=uint(scaled[c]);local_exp[c]=1u;}exponent=max(exponent,local_exp[c]);} - for(uint c=0;c<3u;++c)if(exponent>local_exp[c]){uint d=exponent-local_exp[c];mantissa[c]=d>=31u?0u:uint((2u*mantissa[c]+1u)>>(d+1u));} - packed=(min(mantissa[0],255u))|(min(mantissa[1],255u)<<8u)|(min(mantissa[2],255u)<<16u)|(min(exponent,255u)<<24u); - }else{for(uint c=0;c<3u;++c){int value;if(!jxr_scale(values[c],c,false,params,value,status))return;packed|=uint(clamp(value,0,1023))<<((2u-c)*10u);}} - *reinterpret_cast(output+surface.byte_offset+gid.y*surface.row_stride_bytes+gid.x*4u)=packed; + JxrSurfacePlaneAbi surface = surfaces[0]; + if (gid.x >= surface.width || gid.y >= surface.height || jxr_failed(status)) return; + bool overflow = false; + int values[4]; + jxr_converted(samples, planes, params.crop_x + gid.x, params.crop_y + gid.y, params, values, overflow); + uint packed = 0u; + if (params.output_color == 6u) { + int scaled[3]; uint exponent = 0u; uint mantissa[3]; uint local_exp[3]; + for (uint c = 0; c < 3u; ++c) { + scaled[c] = jxr_scale(values[c], c, false, params, overflow); + if (scaled[c] <= 0) { mantissa[c] = 0; local_exp[c] = 0; } + else if ((scaled[c] >> 7) > 1) { mantissa[c] = uint((scaled[c] & 127) + 128); local_exp[c] = uint(scaled[c] >> 7); } + else { mantissa[c] = uint(scaled[c]); local_exp[c] = 1u; } + exponent = max(exponent, local_exp[c]); + } + for (uint c = 0; c < 3u; ++c) + if (exponent > local_exp[c]) { + uint d = exponent - local_exp[c]; + mantissa[c] = d >= 31u ? 0u : uint((2u * mantissa[c] + 1u) >> (d + 1u)); + } + packed = (min(mantissa[0], 255u)) | (min(mantissa[1], 255u) << 8u) | (min(mantissa[2], 255u) << 16u) | (min(exponent, 255u) << 24u); + } else { + for (uint c = 0; c < 3u; ++c) + packed |= uint(clamp(jxr_scale(values[c], c, false, params, overflow), 0, 1023)) << ((2u - c) * 10u); + } + if (overflow) { jxr_fail(status, 16u); return; } + *reinterpret_cast(output + surface.byte_offset + gid.y * surface.row_stride_bytes + gid.x * 4u) = packed; } diff --git a/crates/jxr-metal/src/kernels/overlap.metal b/crates/jxr-metal/src/kernels/overlap.metal index 86cd46f..dddf37f 100644 --- a/crates/jxr-metal/src/kernels/overlap.metal +++ b/crates/jxr-metal/src/kernels/overlap.metal @@ -1,197 +1,166 @@ -inline bool jxr_inverse_rotate(thread int v[2], device atomic_uint *status, uint code) { - int temporary; - return jxr_add(v[1], 1, temporary, status, code) - && jxr_sub(v[0], temporary >> 1, v[0], status, code) - && jxr_add(v[0], 1, temporary, status, code) - && jxr_add(v[1], temporary >> 1, v[1], status, code); +inline void jxr_inverse_rotate(thread int v[2], thread bool &overflow) { + v[0] = jxr_sub(v[0], jxr_add(v[1], 1, overflow) >> 1, overflow); + v[1] = jxr_add(v[1], jxr_add(v[0], 1, overflow) >> 1, overflow); } -inline bool jxr_inverse_scale(thread int v[2], device atomic_uint *status, uint code) { - int temporary; - return jxr_add(v[0], v[1], v[0], status, code) - && jxr_sub(v[0] >> 1, v[1], v[1], status, code) - && jxr_narrow(long(v[1]) * 3l, temporary, status, code) - && jxr_add(v[0], temporary >> 3, v[0], status, code) - && jxr_narrow(long(v[0]) * 3l, temporary, status, code) - && jxr_add(v[1], temporary >> 4, v[1], status, code) - && jxr_add(v[1], v[0] >> 7, v[1], status, code) - && jxr_sub(v[1], v[0] >> 10, v[1], status, code); +inline void jxr_inverse_scale(thread int v[2], thread bool &overflow) { + v[0] = jxr_add(v[0], v[1], overflow); + v[1] = jxr_sub(v[0] >> 1, v[1], overflow); + v[0] = jxr_add(v[0], jxr_mul3(v[1], overflow) >> 3, overflow); + v[1] = jxr_add(v[1], jxr_mul3(v[0], overflow) >> 4, overflow); + v[1] = jxr_add(v[1], v[0] >> 7, overflow); + v[1] = jxr_sub(v[1], v[0] >> 10, overflow); } -inline bool jxr_inverse_hadamard_post(thread int v[4], device atomic_uint *status, uint code) { - int temporary; - return jxr_sub(v[1], v[2], v[1], status, code) - && jxr_mul3_round(v[3], 4, 3, temporary, status, code) - && jxr_add(v[0], temporary, v[0], status, code) - && jxr_sub(v[3], v[1] >> 1, v[3], status, code) - && jxr_sub(v[0], v[1], temporary, status, code) - && jxr_sub(temporary >> 1, v[2], v[2], status, code) - && (temporary = v[2], v[2] = v[3], v[3] = temporary, true) - && jxr_sub(v[0], v[3], v[0], status, code) - && jxr_add(v[1], v[2], v[1], status, code); +inline void jxr_inverse_hadamard_post(thread int v[4], thread bool &overflow) { + v[1] = jxr_sub(v[1], v[2], overflow); + v[0] = jxr_add(v[0], jxr_mul3_round(v[3], 4, 3, overflow), overflow); + v[3] = jxr_sub(v[3], v[1] >> 1, overflow); + v[2] = jxr_sub(jxr_sub(v[0], v[1], overflow) >> 1, v[2], overflow); + const int swapped = v[2]; + v[2] = v[3]; + v[3] = swapped; + v[0] = jxr_sub(v[0], v[3], overflow); + v[1] = jxr_add(v[1], v[2], overflow); } -inline bool jxr_inverse_todd_odd_post(thread int v[4], device atomic_uint *status, uint code) { - int first; - int second; - int temporary; - return jxr_add(v[3], v[0], v[3], status, code) - && jxr_sub(v[2], v[1], v[2], status, code) - && (first = v[3] >> 1, second = v[2] >> 1, true) - && jxr_sub(v[0], first, v[0], status, code) - && jxr_add(v[1], second, v[1], status, code) - && jxr_mul3_round(v[1], 6, 3, temporary, status, code) - && jxr_sub(v[0], temporary, v[0], status, code) - && jxr_mul3_round(v[0], 2, 2, temporary, status, code) - && jxr_add(v[1], temporary, v[1], status, code) - && jxr_mul3_round(v[1], 4, 3, temporary, status, code) - && jxr_sub(v[0], temporary, v[0], status, code) - && jxr_sub(v[1], second, v[1], status, code) - && jxr_add(v[0], first, v[0], status, code) - && jxr_add(v[2], v[1], v[2], status, code) - && jxr_sub(v[3], v[0], v[3], status, code); +inline void jxr_inverse_todd_odd_post(thread int v[4], thread bool &overflow) { + v[3] = jxr_add(v[3], v[0], overflow); + v[2] = jxr_sub(v[2], v[1], overflow); + const int first = v[3] >> 1; + const int second = v[2] >> 1; + v[0] = jxr_sub(v[0], first, overflow); + v[1] = jxr_add(v[1], second, overflow); + v[0] = jxr_sub(v[0], jxr_mul3_round(v[1], 6, 3, overflow), overflow); + v[1] = jxr_add(v[1], jxr_mul3_round(v[0], 2, 2, overflow), overflow); + v[0] = jxr_sub(v[0], jxr_mul3_round(v[1], 4, 3, overflow), overflow); + v[1] = jxr_sub(v[1], second, overflow); + v[0] = jxr_add(v[0], first, overflow); + v[2] = jxr_add(v[2], v[1], overflow); + v[3] = jxr_sub(v[3], v[0], overflow); } -inline bool jxr_overlap4(thread int v[4], device atomic_uint *status, uint code) { - int temporary; - int pair[2]; - return jxr_add(v[0], v[3], v[0], status, code) - && jxr_add(v[1], v[2], v[1], status, code) - && jxr_add(v[0], 1, temporary, status, code) - && jxr_sub(v[3], temporary >> 1, v[3], status, code) - && jxr_add(v[1], 1, temporary, status, code) - && jxr_sub(v[2], temporary >> 1, v[2], status, code) - && (pair[0] = v[0], pair[1] = v[3], true) - && jxr_inverse_scale(pair, status, code) - && (v[0] = pair[0], v[3] = pair[1], pair[0] = v[1], pair[1] = v[2], true) - && jxr_inverse_scale(pair, status, code) - && (v[1] = pair[0], v[2] = pair[1], true) - && jxr_mul3_round(v[3], 4, 3, temporary, status, code) - && jxr_add(v[0], temporary, v[0], status, code) - && jxr_mul3_round(v[2], 4, 3, temporary, status, code) - && jxr_add(v[1], temporary, v[1], status, code) - && jxr_sub(v[3], v[0] >> 1, v[3], status, code) - && jxr_sub(v[2], v[1] >> 1, v[2], status, code) - && jxr_add(v[0], v[3], v[0], status, code) - && jxr_add(v[1], v[2], v[1], status, code) - && jxr_narrow(-long(v[3]), v[3], status, code) - && jxr_narrow(-long(v[2]), v[2], status, code) - && (pair[0] = v[2], pair[1] = v[3], true) - && jxr_inverse_rotate(pair, status, code) - && (v[2] = pair[0], v[3] = pair[1], true) - && jxr_add(v[0], 1, temporary, status, code) - && jxr_add(v[3], temporary >> 1, v[3], status, code) - && jxr_add(v[1], 1, temporary, status, code) - && jxr_add(v[2], temporary >> 1, v[2], status, code) - && jxr_sub(v[0], v[3], v[0], status, code) - && jxr_sub(v[1], v[2], v[1], status, code); +inline void jxr_overlap4(thread int v[4], thread bool &overflow) { + v[0] = jxr_add(v[0], v[3], overflow); + v[1] = jxr_add(v[1], v[2], overflow); + v[3] = jxr_sub(v[3], jxr_add(v[0], 1, overflow) >> 1, overflow); + v[2] = jxr_sub(v[2], jxr_add(v[1], 1, overflow) >> 1, overflow); + int pair[2] = { v[0], v[3] }; + jxr_inverse_scale(pair, overflow); + v[0] = pair[0]; v[3] = pair[1]; + pair[0] = v[1]; pair[1] = v[2]; + jxr_inverse_scale(pair, overflow); + v[1] = pair[0]; v[2] = pair[1]; + v[0] = jxr_add(v[0], jxr_mul3_round(v[3], 4, 3, overflow), overflow); + v[1] = jxr_add(v[1], jxr_mul3_round(v[2], 4, 3, overflow), overflow); + v[3] = jxr_sub(v[3], v[0] >> 1, overflow); + v[2] = jxr_sub(v[2], v[1] >> 1, overflow); + v[0] = jxr_add(v[0], v[3], overflow); + v[1] = jxr_add(v[1], v[2], overflow); + v[3] = jxr_neg(v[3], overflow); + v[2] = jxr_neg(v[2], overflow); + pair[0] = v[2]; pair[1] = v[3]; + jxr_inverse_rotate(pair, overflow); + v[2] = pair[0]; v[3] = pair[1]; + v[3] = jxr_add(v[3], jxr_add(v[0], 1, overflow) >> 1, overflow); + v[2] = jxr_add(v[2], jxr_add(v[1], 1, overflow) >> 1, overflow); + v[0] = jxr_sub(v[0], v[3], overflow); + v[1] = jxr_sub(v[1], v[2], overflow); } -inline bool jxr_overlap2(thread int v[2], device atomic_uint *status, uint code) { - int temporary; - return jxr_add(v[0], 2, temporary, status, code) - && jxr_add(v[1], temporary >> 2, v[1], status, code) - && jxr_add(v[1], 1, temporary, status, code) - && jxr_add(v[0], temporary >> 1, v[0], status, code) - && jxr_add(v[0], v[1] >> 5, v[0], status, code) - && jxr_add(v[0], v[1] >> 9, v[0], status, code) - && jxr_add(v[0], v[1] >> 13, v[0], status, code) - && jxr_add(v[0], 2, temporary, status, code) - && jxr_add(v[1], temporary >> 2, v[1], status, code); +inline void jxr_overlap2(thread int v[2], thread bool &overflow) { + v[1] = jxr_add(v[1], jxr_add(v[0], 2, overflow) >> 2, overflow); + v[0] = jxr_add(v[0], jxr_add(v[1], 1, overflow) >> 1, overflow); + v[0] = jxr_add(v[0], v[1] >> 5, overflow); + v[0] = jxr_add(v[0], v[1] >> 9, overflow); + v[0] = jxr_add(v[0], v[1] >> 13, overflow); + v[1] = jxr_add(v[1], jxr_add(v[0], 2, overflow) >> 2, overflow); } -inline bool jxr_overlap2x2(thread int v[4], device atomic_uint *status, uint code) { - int temporary; - return jxr_add(v[0], v[3], v[0], status, code) - && jxr_add(v[1], v[2], v[1], status, code) - && jxr_add(v[0], 1, temporary, status, code) - && jxr_sub(v[3], temporary >> 1, v[3], status, code) - && jxr_add(v[1], 1, temporary, status, code) - && jxr_sub(v[2], temporary >> 1, v[2], status, code) - && jxr_add(v[0], 2, temporary, status, code) - && jxr_add(v[1], temporary >> 2, v[1], status, code) - && jxr_add(v[1], 1, temporary, status, code) - && jxr_add(v[0], temporary >> 1, v[0], status, code) - && jxr_add(v[0], v[1] >> 5, v[0], status, code) - && jxr_add(v[0], v[1] >> 9, v[0], status, code) - && jxr_add(v[0], v[1] >> 13, v[0], status, code) - && jxr_add(v[0], 2, temporary, status, code) - && jxr_add(v[1], temporary >> 2, v[1], status, code) - && jxr_add(v[0], 1, temporary, status, code) - && jxr_add(v[3], temporary >> 1, v[3], status, code) - && jxr_add(v[1], 1, temporary, status, code) - && jxr_add(v[2], temporary >> 1, v[2], status, code) - && jxr_sub(v[0], v[3], v[0], status, code) - && jxr_sub(v[1], v[2], v[1], status, code); +inline void jxr_overlap2x2(thread int v[4], thread bool &overflow) { + v[0] = jxr_add(v[0], v[3], overflow); + v[1] = jxr_add(v[1], v[2], overflow); + v[3] = jxr_sub(v[3], jxr_add(v[0], 1, overflow) >> 1, overflow); + v[2] = jxr_sub(v[2], jxr_add(v[1], 1, overflow) >> 1, overflow); + v[1] = jxr_add(v[1], jxr_add(v[0], 2, overflow) >> 2, overflow); + v[0] = jxr_add(v[0], jxr_add(v[1], 1, overflow) >> 1, overflow); + v[0] = jxr_add(v[0], v[1] >> 5, overflow); + v[0] = jxr_add(v[0], v[1] >> 9, overflow); + v[0] = jxr_add(v[0], v[1] >> 13, overflow); + v[1] = jxr_add(v[1], jxr_add(v[0], 2, overflow) >> 2, overflow); + v[3] = jxr_add(v[3], jxr_add(v[0], 1, overflow) >> 1, overflow); + v[2] = jxr_add(v[2], jxr_add(v[1], 1, overflow) >> 1, overflow); + v[0] = jxr_sub(v[0], v[3], overflow); + v[1] = jxr_sub(v[1], v[2], overflow); } -inline bool jxr_group4x4(thread int values[16], uint4 indices, uint kind, - device atomic_uint *status, uint code) { +inline void jxr_group4x4(thread int values[16], uint4 indices, uint kind, thread bool &overflow) { int group[4] = { values[indices.x], values[indices.y], values[indices.z], values[indices.w] }; - bool ok = kind == 0u ? jxr_t2x2h(group, 0, status, code) - : (kind == 1u ? jxr_inverse_todd_odd_post(group, status, code) - : jxr_inverse_hadamard_post(group, status, code)); - if (!ok) return false; + if (kind == 0u) jxr_t2x2h(group, 0, overflow); + else if (kind == 1u) jxr_inverse_todd_odd_post(group, overflow); + else jxr_inverse_hadamard_post(group, overflow); values[indices.x] = group[0]; values[indices.y] = group[1]; values[indices.z] = group[2]; values[indices.w] = group[3]; - return true; } -inline bool jxr_overlap4x4(thread int values[16], device atomic_uint *status, uint code) { +inline void jxr_overlap4x4(thread int values[16], thread bool &overflow) { const uint4 groups[4] = { uint4(0,3,12,15), uint4(1,2,13,14), uint4(4,7,8,11), uint4(5,6,9,10) }; - for (uint i = 0; i < 4; ++i) if (!jxr_group4x4(values, groups[i], 0, status, code)) return false; + for (uint i = 0; i < 4; ++i) jxr_group4x4(values, groups[i], 0, overflow); const uint2 rotations[4] = { uint2(13,12), uint2(9,8), uint2(7,3), uint2(6,2) }; for (uint i = 0; i < 4; ++i) { int pair[2] = { values[rotations[i].x], values[rotations[i].y] }; - if (!jxr_inverse_rotate(pair, status, code)) return false; + jxr_inverse_rotate(pair, overflow); values[rotations[i].x] = pair[0]; values[rotations[i].y] = pair[1]; } - if (!jxr_group4x4(values, uint4(10,11,14,15), 1, status, code)) return false; + jxr_group4x4(values, uint4(10,11,14,15), 1, overflow); const uint2 scales[4] = { uint2(0,15), uint2(1,14), uint2(4,11), uint2(5,10) }; for (uint i = 0; i < 4; ++i) { int pair[2] = { values[scales[i].x], values[scales[i].y] }; - if (!jxr_inverse_scale(pair, status, code)) return false; + jxr_inverse_scale(pair, overflow); values[scales[i].x] = pair[0]; values[scales[i].y] = pair[1]; } - for (uint i = 0; i < 4; ++i) if (!jxr_group4x4(values, groups[i], 2, status, code)) return false; - return true; + for (uint i = 0; i < 4; ++i) jxr_group4x4(values, groups[i], 2, overflow); } inline void jxr_apply_overlap(device int *samples, JxrOverlapWorkAbi work, device atomic_uint *status) { + bool overflow = false; if (work.kind == 0u) { int values[16]; for (uint y = 0; y < 4; ++y) for (uint x = 0; x < 4; ++x) values[y * 4 + x] = samples[work.first + y * work.second + x]; - if (!jxr_overlap4x4(values, status, 2u)) return; + jxr_overlap4x4(values, overflow); + if (overflow) { jxr_fail(status, 2u); return; } for (uint y = 0; y < 4; ++y) for (uint x = 0; x < 4; ++x) samples[work.first + y * work.second + x] = values[y * 4 + x]; } else if (work.kind == 1u) { int values[4] = { samples[work.first], samples[work.first + work.second], samples[work.first + work.second * 2u], samples[work.first + work.second * 3u] }; - if (!jxr_overlap4(values, status, 2u)) return; + jxr_overlap4(values, overflow); + if (overflow) { jxr_fail(status, 2u); return; } for (uint i = 0; i < 4; ++i) samples[work.first + work.second * i] = values[i]; } else if (work.kind == 2u || work.kind == 3u) { int values[4] = { samples[work.first], samples[work.first + 1u], samples[work.first + work.second], samples[work.first + work.second + 1u] }; - bool ok = work.kind == 2u ? jxr_overlap4(values, status, 2u) - : jxr_overlap2x2(values, status, 2u); - if (!ok) return; + if (work.kind == 2u) jxr_overlap4(values, overflow); + else jxr_overlap2x2(values, overflow); + if (overflow) { jxr_fail(status, 2u); return; } samples[work.first] = values[0]; samples[work.first + 1u] = values[1]; samples[work.first + work.second] = values[2]; samples[work.first + work.second + 1u] = values[3]; } else if (work.kind == 4u) { int values[2] = { samples[work.first], samples[work.second] }; - if (!jxr_overlap2(values, status, 2u)) return; + jxr_overlap2(values, overflow); + if (overflow) { jxr_fail(status, 2u); return; } samples[work.first] = values[0]; samples[work.second] = values[1]; } else if (work.kind == 5u) { - int result; - if (jxr_sub(samples[work.first], samples[work.second], result, status, 2u)) - samples[work.first] = result; + const int result = jxr_sub(samples[work.first], samples[work.second], overflow); + if (overflow) { jxr_fail(status, 2u); return; } + samples[work.first] = result; } else if (work.kind == 6u) { - int result; - if (jxr_add(samples[work.first], samples[work.second], result, status, 2u)) - samples[work.first] = result; + const int result = jxr_add(samples[work.first], samples[work.second], overflow); + if (overflow) { jxr_fail(status, 2u); return; } + samples[work.first] = result; } }