Skip to content
Closed
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
44 changes: 30 additions & 14 deletions crates/j2k-metal/src/engine/decode_dispatch/idwt.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,11 +6,35 @@ use crate::metal_types::prelude::*;
use super::{
checked_buffer_copy_into, commit_and_wait_metal, copied_slice_buffer, dispatch_2d_pipeline,
dispatch_3d_pipeline, hybrid_stage_signpost, label_compute_encoder, new_command_buffer,
new_compute_command_encoder, with_runtime, Buffer, CommandBufferRef, ComputeCommandEncoderRef,
DirectIdwtCommandBuffers, Error, J2kIdwtSingleDecompositionParams,
new_compute_command_encoder, new_shared_buffer, with_runtime, Buffer, CommandBufferRef,
ComputeCommandEncoderRef, DirectIdwtCommandBuffers, Error, J2kIdwtSingleDecompositionParams,
J2kRepeatedIdwtSingleDecompositionParams, J2kSingleDecompositionIdwtJob,
SIGNPOST_DECODE_HYBRID_IDWT_COMMAND_ENCODE,
};

#[cfg(target_os = "macos")]
fn checked_host_output_layout(
width: u32,
height: u32,
output_len: usize,
) -> Result<(usize, usize), Error> {
let required_len = (width as usize)
.checked_mul(height as usize)
.ok_or_else(|| Error::MetalKernel {
message: "J2K Metal IDWT output length overflow".to_string(),
})?;
let required_bytes = required_len
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| Error::MetalKernel {
message: "J2K Metal IDWT output byte length overflow".to_string(),
})?;
if output_len < required_len {
return Err(Error::MetalKernel {
message: "J2K Metal IDWT output slice is too small".to_string(),
});
}
Ok((required_len, required_bytes))
}
#[cfg(target_os = "macos")]
mod batched_irreversible;
#[cfg(target_os = "macos")]
Expand Down Expand Up @@ -39,12 +63,8 @@ pub(crate) fn decode_reversible53_single_decomposition_idwt(
output: &mut [f32],
) -> Result<(), Error> {
with_runtime(|runtime| {
let required_len = job.rect.width() as usize * job.rect.height() as usize;
if output.len() < required_len {
return Err(Error::MetalKernel {
message: "J2K Metal IDWT output slice is too small".to_string(),
});
}
let (required_len, required_bytes) =
checked_host_output_layout(job.rect.width(), job.rect.height(), output.len())?;

let params = J2kIdwtSingleDecompositionParams {
x0: job.rect.x0,
Expand All @@ -71,15 +91,11 @@ pub(crate) fn decode_reversible53_single_decomposition_idwt(
hh_height: job.hh.rect.height(),
};

let decoded = new_shared_buffer(&runtime.device, required_bytes)?;
let ll = copied_slice_buffer(&runtime.device, job.ll.coefficients)?;
let hl = copied_slice_buffer(&runtime.device, job.hl.coefficients)?;
let lh = copied_slice_buffer(&runtime.device, job.lh.coefficients)?;
let hh = copied_slice_buffer(&runtime.device, job.hh.coefficients)?;
let decoded = copied_slice_buffer(&runtime.device, output)?;
#[cfg(test)]
crate::engine::test_counters::record_idwt_host_overwritten_output_upload(
std::mem::size_of_val(output),
);

let command_buffer = new_command_buffer(&runtime.queue)?;

Expand Down Expand Up @@ -128,7 +144,7 @@ pub(crate) fn decode_reversible53_single_decomposition_idwt(
);
encoder.endEncoding();
commit_and_wait_metal(&command_buffer)?;
checked_buffer_copy_into(&decoded, 0, output, "IDWT output")?;
checked_buffer_copy_into(&decoded, 0, &mut output[..required_len], "IDWT output")?;
Ok(())
})
}
Expand Down
24 changes: 8 additions & 16 deletions crates/j2k-metal/src/engine/decode_dispatch/idwt/irreversible.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,11 +6,11 @@ use crate::metal_types::prelude::*;
use super::super::{
checked_buffer_copy_into, commit_and_wait_metal, copied_slice_buffer, dispatch_2d_pipeline,
dispatch_3d_pipeline, hybrid_stage_signpost, label_compute_encoder, new_command_buffer,
new_compute_command_encoder, with_runtime, Buffer, CommandBufferRef, ComputeCommandEncoderRef,
Error, J2kIdwt97StepParams, J2kIdwtSingleDecompositionParams, J2kSingleDecompositionIdwtJob,
SIGNPOST_DECODE_HYBRID_IDWT_COMMAND_ENCODE,
new_compute_command_encoder, new_shared_buffer, with_runtime, Buffer, CommandBufferRef,
ComputeCommandEncoderRef, Error, J2kIdwt97StepParams, J2kIdwtSingleDecompositionParams,
J2kSingleDecompositionIdwtJob, SIGNPOST_DECODE_HYBRID_IDWT_COMMAND_ENCODE,
};
use super::{IdwtSubBandBuffers, SingleIdwtDispatch};
use super::{checked_host_output_layout, IdwtSubBandBuffers, SingleIdwtDispatch};
use j2k_codec_math::dwt;

#[cfg(test)]
Expand Down Expand Up @@ -51,12 +51,8 @@ fn decode_irreversible97_staged_single_decomposition_idwt_with_high_pass(
high_pass: f32,
) -> Result<(), Error> {
with_runtime(|runtime| {
let required_len = job.rect.width() as usize * job.rect.height() as usize;
if output.len() < required_len {
return Err(Error::MetalKernel {
message: "J2K Metal IDWT output slice is too small".to_string(),
});
}
let (required_len, required_bytes) =
checked_host_output_layout(job.rect.width(), job.rect.height(), output.len())?;

let params = J2kIdwtSingleDecompositionParams {
x0: job.rect.x0,
Expand All @@ -83,15 +79,11 @@ fn decode_irreversible97_staged_single_decomposition_idwt_with_high_pass(
hh_height: job.hh.rect.height(),
};

let decoded = new_shared_buffer(&runtime.device, required_bytes)?;
let ll = copied_slice_buffer(&runtime.device, job.ll.coefficients)?;
let hl = copied_slice_buffer(&runtime.device, job.hl.coefficients)?;
let lh = copied_slice_buffer(&runtime.device, job.lh.coefficients)?;
let hh = copied_slice_buffer(&runtime.device, job.hh.coefficients)?;
let decoded = copied_slice_buffer(&runtime.device, output)?;
#[cfg(test)]
crate::engine::test_counters::record_idwt_host_overwritten_output_upload(
std::mem::size_of_val(output),
);
let command_buffer = new_command_buffer(&runtime.queue)?;
let encoder = new_compute_command_encoder(&command_buffer)?;
dispatch_irreversible97_single_decomposition_buffers_in_encoder_with_high_pass(
Expand All @@ -117,7 +109,7 @@ fn decode_irreversible97_staged_single_decomposition_idwt_with_high_pass(
encoder.endEncoding();
commit_and_wait_metal(&command_buffer)?;

checked_buffer_copy_into(&decoded, 0, output, "IDWT output")?;
checked_buffer_copy_into(&decoded, 0, &mut output[..required_len], "IDWT output")?;
Ok(())
})
}
Expand Down
54 changes: 51 additions & 3 deletions crates/j2k-metal/src/engine/direct_buffers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -92,9 +92,57 @@ pub(crate) fn checked_buffer_copy_into<T: GpuAbi>(
output: &mut [T],
context: &str,
) -> Result<(), Error> {
let values = checked_buffer_slice_at(buffer, byte_offset, output.len(), context)?;
output.copy_from_slice(&values);
Ok(())
let result = (|| {
let buffer_len = buffer.length();
let element_size = size_of::<T>();
if element_size == 0 {
return Err(MetalSupportError::BufferZeroSizedType { abi_name: T::NAME });
}
let byte_len =
output
.len()
.checked_mul(element_size)
.ok_or(MetalSupportError::BufferBounds {
offset_bytes: byte_offset,
byte_len: usize::MAX,
buffer_len,
})?;
let end = byte_offset
.checked_add(byte_len)
.ok_or(MetalSupportError::BufferBounds {
offset_bytes: byte_offset,
byte_len,
buffer_len,
})?;
if end > buffer_len {
return Err(MetalSupportError::BufferBounds {
offset_bytes: byte_offset,
byte_len,
buffer_len,
});
}
let align = core::mem::align_of::<T>();
if !byte_offset.is_multiple_of(align) {
return Err(MetalSupportError::BufferAlignment {
offset_bytes: byte_offset,
align,
});
}
if output.is_empty() {
return Ok(());
}
// SAFETY: These private readback callers use completed/synchronized,
// immutable Metal storage and a separate, exclusively borrowed host
// destination. Reuse the audited support read for storage and actual
// pointer alignment, then the existing checked byte-borrow boundary.
let bytes = unsafe {
support_checked_buffer_read::<T>(buffer, byte_offset)?;
completed_metal_buffer_bytes(buffer, byte_offset, byte_len)?
};
T::slice_as_bytes_mut(output).copy_from_slice(bytes);
Ok(())
})();
result.map_err(|error| buffer_access_error(context, error))
}

#[cfg(target_os = "macos")]
Expand Down
8 changes: 0 additions & 8 deletions crates/j2k-metal/src/engine/test_counters.rs
Original file line number Diff line number Diff line change
Expand Up @@ -295,14 +295,6 @@ pub(crate) fn idwt_host_transfer_counters_for_test() -> (usize, usize, usize) {
)
}

pub(crate) fn record_idwt_host_overwritten_output_upload(bytes: usize) {
IDWT_HOST_OVERWRITTEN_OUTPUT_UPLOAD_BYTES.set(
IDWT_HOST_OVERWRITTEN_OUTPUT_UPLOAD_BYTES
.get()
.saturating_add(bytes),
);
}

pub(crate) fn record_idwt_host_temporary_readback_vec(bytes: usize) {
IDWT_HOST_TEMPORARY_READBACK_VEC_ALLOCATIONS.set(
IDWT_HOST_TEMPORARY_READBACK_VEC_ALLOCATIONS
Expand Down
9 changes: 1 addition & 8 deletions crates/j2k-metal/src/idwt_host_slice_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -142,14 +142,7 @@ fn host_slice_idwt_output_transfer_accounting() {
assert!(decoder
.decode_single_decomposition_idwt(fixture.job(transform), &mut output)
.unwrap());
assert_eq!(
idwt_host_transfer_counters_for_test(),
(
output.len() * size_of::<f32>(),
1,
output.len() * size_of::<f32>()
)
);
assert_eq!(idwt_host_transfer_counters_for_test(), (0, 0, 0));
}
}

Expand Down
Loading