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; } }