Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions crates/jxr-metal/src/abi.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
35 changes: 18 additions & 17 deletions crates/jxr-metal/src/encode.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -525,32 +525,33 @@ 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)?;
encoder.bind_buffer(2, buffers.low.buffer(), 0)?;
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<u64>,
block_rows: impl Into<u64>,
) -> Result<u64, MetalError> {
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<dyn MTLDevice>,
encoder: &ProtocolObject<dyn MTLComputeCommandEncoder>,
Expand Down
24 changes: 10 additions & 14 deletions crates/jxr-metal/src/encode/batch.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)?;
Expand All @@ -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),
),
Expand Down
171 changes: 81 additions & 90 deletions crates/jxr-metal/src/kernels/common.metal
Original file line number Diff line number Diff line change
Expand Up @@ -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<int>(as_type<uint>(a) + as_type<uint>(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<int>(as_type<uint>(a) - as_type<uint>(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<int>(0u - as_type<uint>(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<uint>(value) : as_type<uint>(value);
const uint product = magnitude * factor;
overflow |= mulhi(magnitude, factor) != 0u || product > (negative ? 0x80000000u : 0x7fffffffu);
return as_type<int>(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<int>(as_type<uint>(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);
}
31 changes: 16 additions & 15 deletions crates/jxr-metal/src/kernels/first_transform.metal
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
Loading
Loading