j2k-metal 0.7.5

Metal decoder and encode-stage adapter for j2k
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
// SPDX-License-Identifier: MIT OR Apache-2.0

use core::mem::size_of;
use std::sync::Arc;

use j2k_native::{DecodeSettings, DecoderContext, EncodeOptions, Image};
use metal::foreign_types::ForeignType;

use super::super::{
    default_metal_ht_chunk_limits, encode_metal_ht_batches_in_encoder,
    encode_repeated_metal_ht_batch_in_command_buffer, HtBatchInput,
};
use crate::compute::{
    checked_buffer_slice, commit_and_wait_metal, decode_prepared_ht_sub_band_group_on_cpu_profile,
    new_command_buffer, new_compute_command_encoder, new_shared_buffer,
    prepare_direct_grayscale_plan, validate_direct_status, wait_for_completion_metal,
    DirectStatusCheck, MetalRuntime, PreparedHtExecutionOwner,
};

#[test]
fn completed_ht_status_storage_returns_to_the_shared_pool() {
    if !j2k_test_support::metal_runtime_gate(module_path!()) {
        return;
    }
    let pixels = (0..64_u8).rev().collect::<Vec<_>>();
    let bytes = j2k_native::encode_htj2k(
        &pixels,
        8,
        8,
        1,
        8,
        false,
        &EncodeOptions {
            reversible: true,
            num_decomposition_levels: 1,
            ..EncodeOptions::default()
        },
    )
    .expect("encode status-pool fixture");
    let image = Image::new(&bytes, &DecodeSettings::default()).expect("fixture image");
    let mut context = DecoderContext::default();
    let direct = image
        .build_direct_grayscale_plan_with_context(&mut context)
        .expect("direct fixture plan");
    let prepared = prepare_direct_grayscale_plan(&direct).expect("prepared fixture plan");
    let group = prepared.ht_groups.first().expect("prepared HT group");
    let input = HtBatchInput {
        source_index: 0,
        payload: group.payload_source.as_ht_payload_source(),
        jobs: &group.jobs,
        output_base: 0,
        execution_owner: &group.execution_owner,
    };
    let runtime = MetalRuntime::new().expect("isolated Metal runtime");

    for submission in 0..2 {
        let output =
            new_shared_buffer(&runtime.device, group.total_coefficients * size_of::<f32>())
                .expect("status-pool output");
        let command_buffer = new_command_buffer(&runtime.queue).expect("status-pool command");
        let encoder = new_compute_command_encoder(&command_buffer).expect("status-pool encoder");
        let (retained, status) = encode_metal_ht_batches_in_encoder(
            &runtime,
            &encoder,
            &[input],
            &output,
            group.total_coefficients,
            default_metal_ht_chunk_limits(),
        )
        .expect("status-pool encode");
        encoder.end_encoding();
        commit_and_wait_metal(&command_buffer).expect("status-pool completion");
        validate_direct_status(&runtime, status).expect("status-pool validation");
        drop(retained);

        let diagnostics = runtime
            .buffer_pool_diagnostics()
            .expect("status-pool diagnostics");
        assert_eq!(
            diagnostics.shared.cached_buffers, 1,
            "submission {submission} must retire exactly one reusable HT status buffer"
        );
    }
}

#[test]
fn reused_ht_status_storage_is_overwritten_by_every_dispatched_job() {
    if !j2k_test_support::metal_runtime_gate(module_path!()) {
        return;
    }
    let image = Image::new(
        j2k_test_support::openhtj2k_refinement_fixture(),
        &DecodeSettings::default(),
    )
    .expect("refinement status-overwrite fixture image");
    let mut context = DecoderContext::default();
    let direct = image
        .build_direct_grayscale_plan_with_context(&mut context)
        .expect("direct fixture plan");
    let prepared = prepare_direct_grayscale_plan(&direct).expect("prepared fixture plan");
    let group = prepared.ht_groups.first().expect("prepared HT group");
    assert!(
        group.jobs.iter().any(|job| job.number_of_coding_passes > 1),
        "status-overwrite fixture must exercise refinement jobs"
    );
    let mut invalid_jobs = group.jobs.clone();
    for job in &mut invalid_jobs {
        job.num_bitplanes = 0;
    }
    let invalid_owner = Arc::new(PreparedHtExecutionOwner);
    let runtime = MetalRuntime::new().expect("isolated Metal runtime");

    let invalid_output =
        new_shared_buffer(&runtime.device, group.total_coefficients * size_of::<f32>())
            .expect("invalid status output");
    let invalid_command = new_command_buffer(&runtime.queue).expect("invalid status command");
    let invalid_encoder =
        new_compute_command_encoder(&invalid_command).expect("invalid status encoder");
    let (_, invalid_status) = encode_metal_ht_batches_in_encoder(
        &runtime,
        &invalid_encoder,
        &[HtBatchInput {
            source_index: 0,
            payload: group.payload_source.as_ht_payload_source(),
            jobs: &invalid_jobs,
            output_base: 0,
            execution_owner: &invalid_owner,
        }],
        &invalid_output,
        group.total_coefficients,
        default_metal_ht_chunk_limits(),
    )
    .expect("invalid status encode");
    invalid_encoder.end_encoding();
    commit_and_wait_metal(&invalid_command).expect("invalid status completion");
    let DirectStatusCheck::Ht {
        buffer: invalid_status_buffer,
        ..
    } = &invalid_status
    else {
        panic!("invalid distinct submission must retain HT status")
    };
    let invalid_status_ptr = invalid_status_buffer.as_ptr();
    assert!(
        validate_direct_status(&runtime, invalid_status).is_err(),
        "invalid first submission must seed every status slot with a failure"
    );

    let valid_output =
        new_shared_buffer(&runtime.device, group.total_coefficients * size_of::<f32>())
            .expect("valid status output");
    let valid_command = new_command_buffer(&runtime.queue).expect("valid status command");
    let valid_encoder = new_compute_command_encoder(&valid_command).expect("valid status encoder");
    let (_, valid_status) = encode_metal_ht_batches_in_encoder(
        &runtime,
        &valid_encoder,
        &[HtBatchInput {
            source_index: 0,
            payload: group.payload_source.as_ht_payload_source(),
            jobs: &group.jobs,
            output_base: 0,
            execution_owner: &group.execution_owner,
        }],
        &valid_output,
        group.total_coefficients,
        default_metal_ht_chunk_limits(),
    )
    .expect("valid status encode");
    valid_encoder.end_encoding();
    commit_and_wait_metal(&valid_command).expect("valid status completion");
    let DirectStatusCheck::Ht {
        buffer: valid_status_buffer,
        ..
    } = &valid_status
    else {
        panic!("valid distinct submission must retain HT status")
    };
    assert_eq!(
        valid_status_buffer.as_ptr(),
        invalid_status_ptr,
        "valid distinct dispatch must overwrite the recycled failure-status allocation"
    );
    validate_direct_status(&runtime, valid_status)
        .expect("every valid dispatch must overwrite its recycled failure status");
}

#[test]
fn reused_repeated_ht_status_storage_is_overwritten_by_every_dispatched_job() {
    if !j2k_test_support::metal_runtime_gate(module_path!()) {
        return;
    }
    let image = Image::new(
        j2k_test_support::openhtj2k_refinement_fixture(),
        &DecodeSettings::default(),
    )
    .expect("repeated refinement status fixture image");
    let mut context = DecoderContext::default();
    let direct = image
        .build_direct_grayscale_plan_with_context(&mut context)
        .expect("repeated refinement direct fixture plan");
    let prepared = prepare_direct_grayscale_plan(&direct).expect("prepared refinement plan");
    let group = prepared.ht_groups.first().expect("prepared HT group");
    assert!(
        group.jobs.iter().any(|job| job.number_of_coding_passes > 1),
        "repeated status-overwrite fixture must exercise refinement jobs"
    );
    let mut invalid_jobs = group.jobs.clone();
    for job in &mut invalid_jobs {
        job.num_bitplanes = 0;
    }
    let invalid_owner = Arc::new(PreparedHtExecutionOwner);
    let runtime = MetalRuntime::new().expect("isolated repeated Metal runtime");
    let count = 2;
    let output_words = group
        .total_coefficients
        .checked_mul(count)
        .expect("repeated status output word count");

    let invalid_output = new_shared_buffer(&runtime.device, output_words * size_of::<f32>())
        .expect("invalid repeated status output");
    let invalid_command =
        new_command_buffer(&runtime.queue).expect("invalid repeated status command");
    let (_, invalid_status) = encode_repeated_metal_ht_batch_in_command_buffer(
        &runtime,
        &invalid_command,
        HtBatchInput {
            source_index: 0,
            payload: group.payload_source.as_ht_payload_source(),
            jobs: &invalid_jobs,
            output_base: 0,
            execution_owner: &invalid_owner,
        },
        count,
        group.total_coefficients,
        &invalid_output,
        default_metal_ht_chunk_limits(),
    )
    .expect("invalid repeated status encode");
    commit_and_wait_metal(&invalid_command).expect("invalid repeated status completion");
    let DirectStatusCheck::Ht {
        buffer: invalid_status_buffer,
        ..
    } = &invalid_status
    else {
        panic!("invalid repeated submission must retain HT status")
    };
    let invalid_status_ptr = invalid_status_buffer.as_ptr();
    assert!(
        validate_direct_status(&runtime, invalid_status).is_err(),
        "invalid repeated submission must seed every status slot with a failure"
    );

    let valid_output = new_shared_buffer(&runtime.device, output_words * size_of::<f32>())
        .expect("valid repeated status output");
    let valid_command = new_command_buffer(&runtime.queue).expect("valid repeated status command");
    let (_, valid_status) = encode_repeated_metal_ht_batch_in_command_buffer(
        &runtime,
        &valid_command,
        HtBatchInput {
            source_index: 0,
            payload: group.payload_source.as_ht_payload_source(),
            jobs: &group.jobs,
            output_base: 0,
            execution_owner: &group.execution_owner,
        },
        count,
        group.total_coefficients,
        &valid_output,
        default_metal_ht_chunk_limits(),
    )
    .expect("valid repeated status encode");
    commit_and_wait_metal(&valid_command).expect("valid repeated status completion");
    let DirectStatusCheck::Ht {
        buffer: valid_status_buffer,
        ..
    } = &valid_status
    else {
        panic!("valid repeated submission must retain HT status")
    };
    assert_eq!(
        valid_status_buffer.as_ptr(),
        invalid_status_ptr,
        "valid repeated dispatch must overwrite the recycled failure-status allocation"
    );
    validate_direct_status(&runtime, valid_status)
        .expect("every valid repeated dispatch must overwrite recycled failure status");
}

#[test]
#[expect(
    clippy::too_many_lines,
    reason = "the overlap-lifetime scenario keeps both pending submissions and ownership assertions visible"
)]
fn overlapping_prepared_ht_submissions_keep_distinct_status_owners() {
    if !j2k_test_support::metal_runtime_gate(module_path!()) {
        return;
    }
    let pixels = (0..64_u8).collect::<Vec<_>>();
    let bytes = j2k_native::encode_htj2k(
        &pixels,
        8,
        8,
        1,
        8,
        false,
        &EncodeOptions {
            reversible: true,
            num_decomposition_levels: 1,
            ..EncodeOptions::default()
        },
    )
    .expect("encode overlapping-status fixture");
    let image = Image::new(&bytes, &DecodeSettings::default()).expect("fixture image");
    let mut context = DecoderContext::default();
    let direct = image
        .build_direct_grayscale_plan_with_context(&mut context)
        .expect("direct fixture plan");
    let prepared = prepare_direct_grayscale_plan(&direct).expect("prepared fixture plan");
    let group = prepared.ht_groups.first().expect("prepared HT group");
    let input = HtBatchInput {
        source_index: 0,
        payload: group.payload_source.as_ht_payload_source(),
        jobs: &group.jobs,
        output_base: 0,
        execution_owner: &group.execution_owner,
    };
    let cpu = decode_prepared_ht_sub_band_group_on_cpu_profile(group, None)
        .expect("overlapping CPU coefficient oracle");
    let runtime = MetalRuntime::new().expect("isolated Metal runtime");

    let first_output =
        new_shared_buffer(&runtime.device, group.total_coefficients * size_of::<f32>())
            .expect("first overlapping output");
    let first_command = new_command_buffer(&runtime.queue).expect("first overlapping command");
    let first_encoder =
        new_compute_command_encoder(&first_command).expect("first overlapping encoder");
    let (first_retained, first_status) = encode_metal_ht_batches_in_encoder(
        &runtime,
        &first_encoder,
        &[input],
        &first_output,
        group.total_coefficients,
        default_metal_ht_chunk_limits(),
    )
    .expect("first overlapping encode");
    first_encoder.end_encoding();
    first_command.commit();

    let second_output =
        new_shared_buffer(&runtime.device, group.total_coefficients * size_of::<f32>())
            .expect("second overlapping output");
    let second_command = new_command_buffer(&runtime.queue).expect("second overlapping command");
    let second_encoder =
        new_compute_command_encoder(&second_command).expect("second overlapping encoder");
    let (second_retained, second_status) = encode_metal_ht_batches_in_encoder(
        &runtime,
        &second_encoder,
        &[input],
        &second_output,
        group.total_coefficients,
        default_metal_ht_chunk_limits(),
    )
    .expect("second overlapping encode");
    second_encoder.end_encoding();
    second_command.commit();

    let DirectStatusCheck::Ht {
        buffer: first_buffer,
        ..
    } = &first_status
    else {
        panic!("first prepared HT submission must retain HT status")
    };
    let DirectStatusCheck::Ht {
        buffer: second_buffer,
        ..
    } = &second_status
    else {
        panic!("second prepared HT submission must retain HT status")
    };
    assert_ne!(
        first_buffer.as_ptr(),
        second_buffer.as_ptr(),
        "overlapping submissions must not alias in-flight status storage"
    );

    wait_for_completion_metal(&first_command).expect("first overlapping completion");
    wait_for_completion_metal(&second_command).expect("second overlapping completion");
    validate_direct_status(&runtime, first_status).expect("first overlapping status");
    validate_direct_status(&runtime, second_status).expect("second overlapping status");
    let first_coefficients = checked_buffer_slice::<f32>(
        &first_output,
        group.total_coefficients,
        "first overlapping HT output",
    )
    .expect("read first overlapping HT output");
    let second_coefficients = checked_buffer_slice::<f32>(
        &second_output,
        group.total_coefficients,
        "second overlapping HT output",
    )
    .expect("read second overlapping HT output");
    assert_eq!(first_coefficients, cpu);
    assert_eq!(second_coefficients, cpu);
    drop((first_retained, second_retained));
}