Skip to main content

j2k_cuda/
codec.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2
3use j2k::{
4    DeviceDecodePlan, DeviceDecodeRequest, J2kCodec as CpuCodec, J2kContext as CpuJ2kContext,
5    J2kDecodeWarning, J2kDecoder as CpuDecoder, J2kScratchPool as CpuJ2kScratchPool,
6};
7use j2k_core::{
8    checked_surface_len, submit_ready_device, BackendRequest, Downscale, ImageCodec, PixelFormat,
9    ReadySubmission, Rect, TileBatchDecode, TileBatchDecodeDevice, TileBatchDecodeManyDevice,
10    TileBatchDecodeSubmit, TileRegionScaledDecodeJob, TileRegionScaledDeviceDecodeRequest,
11    DEFAULT_MAX_HOST_ALLOCATION_BYTES,
12};
13
14use crate::{
15    allocation::{try_collect_results_exact, try_vec_filled},
16    routing::{auto_cuda_available, auto_repeated_decode_uses_cuda, inputs_repeat_one_slice},
17    runtime::{validate_surface_request, wrap_surface},
18};
19use crate::{CudaSession, Error, J2kDecoder, Surface};
20
21/// Marker type implementing tile-batch CUDA surface decode traits.
22#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
23pub struct Codec;
24
25struct RegionScaledSurfaceRequest<'a> {
26    ctx: &'a mut CpuJ2kContext,
27    session: &'a mut CudaSession,
28    pool: &'a mut CpuJ2kScratchPool,
29    input: &'a [u8],
30    fmt: PixelFormat,
31    roi: Rect,
32    scale: Downscale,
33    backend: BackendRequest,
34}
35
36#[doc(hidden)]
37impl ImageCodec for Codec {
38    type Error = Error;
39    type Warning = J2kDecodeWarning;
40    type Pool = crate::J2kScratchPool;
41}
42
43impl Codec {
44    fn supports_cuda_batch_format(fmt: PixelFormat) -> bool {
45        matches!(
46            fmt,
47            PixelFormat::Gray8
48                | PixelFormat::Gray16
49                | PixelFormat::GrayI16
50                | PixelFormat::Rgb8
51                | PixelFormat::Rgba8
52                | PixelFormat::Rgb16
53                | PixelFormat::Rgba16
54        )
55    }
56
57    #[cfg(feature = "cuda-runtime")]
58    fn decode_tiles_to_cuda_batch(
59        inputs: &[&[u8]],
60        fmt: PixelFormat,
61        session: &mut CudaSession,
62    ) -> Result<Vec<Surface>, Error> {
63        J2kDecoder::decode_batch_to_device_with_session(inputs, fmt, session)
64    }
65
66    #[cfg(not(feature = "cuda-runtime"))]
67    fn decode_tiles_to_cuda_batch(
68        _inputs: &[&[u8]],
69        _fmt: PixelFormat,
70        _session: &mut CudaSession,
71    ) -> Result<Vec<Surface>, Error> {
72        Err(Error::CudaUnavailable)
73    }
74
75    fn decode_tile_to_surface_impl(
76        ctx: &mut CpuJ2kContext,
77        session: &mut CudaSession,
78        pool: &mut CpuJ2kScratchPool,
79        input: &[u8],
80        fmt: PixelFormat,
81        backend: BackendRequest,
82    ) -> Result<Surface, Error> {
83        validate_surface_request(backend)?;
84        if matches!(backend, BackendRequest::Cuda) {
85            let mut decoder = J2kDecoder::new(input)?;
86            return decoder.decode_to_device_with_session(fmt, session);
87        }
88        let dims = CpuDecoder::inspect(input)?.dimensions;
89        let (mut out, stride) = allocate_cpu_surface(dims, fmt)?;
90        CpuCodec::decode_tile(ctx, pool, input, &mut out, stride, fmt)?;
91        wrap_surface(out, dims, fmt, backend, session)
92    }
93
94    fn decode_tile_region_to_surface_impl(
95        ctx: &mut CpuJ2kContext,
96        session: &mut CudaSession,
97        pool: &mut CpuJ2kScratchPool,
98        input: &[u8],
99        fmt: PixelFormat,
100        roi: Rect,
101        backend: BackendRequest,
102    ) -> Result<Surface, Error> {
103        validate_surface_request(backend)?;
104        if matches!(backend, BackendRequest::Cuda) {
105            let mut decoder = J2kDecoder::new(input)?;
106            return decoder.decode_region_to_device_with_session(fmt, roi, session);
107        }
108        let dims = DeviceDecodePlan::for_image(
109            CpuDecoder::inspect(input)?.dimensions,
110            DeviceDecodeRequest::Region { roi },
111        )?
112        .output_dims();
113        let (mut out, stride) = allocate_cpu_surface(dims, fmt)?;
114        CpuCodec::decode_tile_region(ctx, pool, input, &mut out, stride, fmt, roi)?;
115        wrap_surface(out, dims, fmt, backend, session)
116    }
117
118    fn decode_tile_scaled_to_surface_impl(
119        ctx: &mut CpuJ2kContext,
120        session: &mut CudaSession,
121        pool: &mut CpuJ2kScratchPool,
122        input: &[u8],
123        fmt: PixelFormat,
124        scale: Downscale,
125        backend: BackendRequest,
126    ) -> Result<Surface, Error> {
127        validate_surface_request(backend)?;
128        if matches!(backend, BackendRequest::Cuda) {
129            let mut decoder = J2kDecoder::new(input)?;
130            return decoder.decode_scaled_to_device_with_session(fmt, scale, session);
131        }
132        let dims = DeviceDecodePlan::for_image(
133            CpuDecoder::inspect(input)?.dimensions,
134            DeviceDecodeRequest::Scaled { scale },
135        )?
136        .output_dims();
137        let (mut out, stride) = allocate_cpu_surface(dims, fmt)?;
138        CpuCodec::decode_tile_scaled(ctx, pool, input, &mut out, stride, fmt, scale)?;
139        wrap_surface(out, dims, fmt, backend, session)
140    }
141
142    fn decode_tile_region_scaled_to_surface_impl(
143        request: RegionScaledSurfaceRequest<'_>,
144    ) -> Result<Surface, Error> {
145        let RegionScaledSurfaceRequest {
146            ctx,
147            session,
148            pool,
149            input,
150            fmt,
151            roi,
152            scale,
153            backend,
154        } = request;
155        validate_surface_request(backend)?;
156        if matches!(backend, BackendRequest::Cuda) {
157            let mut decoder = J2kDecoder::new(input)?;
158            return decoder.decode_region_scaled_to_device_with_session(fmt, roi, scale, session);
159        }
160        let dims = DeviceDecodePlan::for_image(
161            CpuDecoder::inspect(input)?.dimensions,
162            DeviceDecodeRequest::RegionScaled { roi, scale },
163        )?
164        .output_dims();
165        let (mut out, stride) = allocate_cpu_surface(dims, fmt)?;
166        CpuCodec::decode_tile_region_scaled(
167            ctx,
168            pool,
169            fmt,
170            TileRegionScaledDecodeJob {
171                input,
172                out: &mut out,
173                stride,
174                roi,
175                scale,
176            },
177        )?;
178        wrap_surface(out, dims, fmt, backend, session)
179    }
180}
181
182fn allocate_cpu_surface(dims: (u32, u32), fmt: PixelFormat) -> Result<(Vec<u8>, usize), Error> {
183    let (stride, len) = checked_surface_len(
184        dims,
185        fmt.bytes_per_pixel(),
186        DEFAULT_MAX_HOST_ALLOCATION_BYTES,
187        "j2k CUDA CPU fallback surface",
188    )?;
189    Ok((
190        try_vec_filled(len, 0u8, "j2k CUDA CPU fallback surface")?,
191        stride,
192    ))
193}
194
195#[doc(hidden)]
196impl TileBatchDecodeSubmit for Codec {
197    type Context = CpuJ2kContext;
198    type Session = CudaSession;
199    type DeviceSurface = Surface;
200    type SubmittedSurface = ReadySubmission<Surface, Error>;
201
202    fn submit_tile_to_device(
203        ctx: &mut Self::Context,
204        session: &mut Self::Session,
205        pool: &mut Self::Pool,
206        input: &[u8],
207        fmt: PixelFormat,
208        backend: BackendRequest,
209    ) -> Result<Self::SubmittedSurface, Self::Error> {
210        validate_surface_request(backend)?;
211        Ok(submit_ready_device(session, |session| {
212            Self::decode_tile_to_surface_impl(ctx, session, pool, input, fmt, backend)
213        }))
214    }
215
216    fn submit_tile_region_to_device(
217        ctx: &mut Self::Context,
218        session: &mut Self::Session,
219        pool: &mut Self::Pool,
220        input: &[u8],
221        fmt: PixelFormat,
222        roi: Rect,
223        backend: BackendRequest,
224    ) -> Result<Self::SubmittedSurface, Self::Error> {
225        validate_surface_request(backend)?;
226        Ok(submit_ready_device(session, |session| {
227            Self::decode_tile_region_to_surface_impl(ctx, session, pool, input, fmt, roi, backend)
228        }))
229    }
230
231    fn submit_tile_scaled_to_device(
232        ctx: &mut Self::Context,
233        session: &mut Self::Session,
234        pool: &mut Self::Pool,
235        input: &[u8],
236        fmt: PixelFormat,
237        scale: Downscale,
238        backend: BackendRequest,
239    ) -> Result<Self::SubmittedSurface, Self::Error> {
240        validate_surface_request(backend)?;
241        Ok(submit_ready_device(session, |session| {
242            Self::decode_tile_scaled_to_surface_impl(ctx, session, pool, input, fmt, scale, backend)
243        }))
244    }
245
246    fn submit_tile_region_scaled_to_device(
247        ctx: &mut Self::Context,
248        session: &mut Self::Session,
249        pool: &mut Self::Pool,
250        request: TileRegionScaledDeviceDecodeRequest<'_>,
251    ) -> Result<Self::SubmittedSurface, Self::Error> {
252        let TileRegionScaledDeviceDecodeRequest {
253            input,
254            fmt,
255            roi,
256            scale,
257            backend,
258        } = request;
259        validate_surface_request(backend)?;
260        Ok(submit_ready_device(session, |session| {
261            Self::decode_tile_region_scaled_to_surface_impl(RegionScaledSurfaceRequest {
262                ctx,
263                session,
264                pool,
265                input,
266                fmt,
267                roi,
268                scale,
269                backend,
270            })
271        }))
272    }
273}
274
275#[doc(hidden)]
276impl TileBatchDecodeDevice for Codec {
277    type Context = CpuJ2kContext;
278    type DeviceSurface = Surface;
279}
280
281#[doc(hidden)]
282impl TileBatchDecodeManyDevice for Codec {
283    type Context = CpuJ2kContext;
284    type DeviceSurface = Surface;
285
286    fn decode_tiles_to_device(
287        ctx: &mut Self::Context,
288        pool: &mut Self::Pool,
289        inputs: &[&[u8]],
290        fmt: PixelFormat,
291        backend: BackendRequest,
292    ) -> Result<Vec<Self::DeviceSurface>, Self::Error> {
293        validate_surface_request(backend)?;
294        if inputs.is_empty() {
295            return Ok(Vec::new());
296        }
297
298        let mut session = CudaSession::default();
299        if matches!(backend, BackendRequest::Cuda) && Self::supports_cuda_batch_format(fmt) {
300            return Self::decode_tiles_to_cuda_batch(inputs, fmt, &mut session);
301        }
302        if backend == BackendRequest::Auto
303            && Self::supports_cuda_batch_format(fmt)
304            && inputs_repeat_one_slice(inputs)
305        {
306            let support = CpuDecoder::inspect_support(inputs[0])?;
307            if auto_repeated_decode_uses_cuda(
308                support.info.dimensions,
309                support.info.components,
310                fmt,
311                support.transfer_syntax,
312                support.payload_kind,
313                inputs.len(),
314            ) && auto_cuda_available(&mut session)?
315            {
316                return Self::decode_tiles_to_cuda_batch(inputs, fmt, &mut session);
317            }
318        }
319
320        try_collect_results_exact(
321            inputs.iter().map(|input| {
322                Self::decode_tile_to_surface_impl(ctx, &mut session, pool, input, fmt, backend)
323            }),
324            "j2k CUDA decode batch surfaces",
325        )
326    }
327}
328
329#[cfg(all(test, feature = "cuda-runtime"))]
330mod tests {
331    use j2k_core::{BackendRequest, PixelFormat, TileBatchDecodeManyDevice};
332    use j2k_test_support::{cuda_runtime_required, htj2k_rgb8_pattern_fixture};
333
334    use super::{Codec, CpuJ2kContext, CpuJ2kScratchPool};
335    use crate::decoder::{
336        testing_cuda_htj2k_batch_decode_calls, testing_reset_cuda_htj2k_batch_decode_calls,
337    };
338    use crate::{Error, SurfaceResidency};
339
340    #[test]
341    fn explicit_cuda_rgb_many_decode_uses_batch_api_once() {
342        testing_reset_cuda_htj2k_batch_decode_calls();
343        let fixture = rgb8_htj2k_fixture(32, 32);
344        let inputs = [fixture.as_slice(), fixture.as_slice()];
345        let mut ctx = CpuJ2kContext::default();
346        let mut pool = CpuJ2kScratchPool::new();
347
348        let result = Codec::decode_tiles_to_device(
349            &mut ctx,
350            &mut pool,
351            &inputs,
352            PixelFormat::Rgb8,
353            BackendRequest::Cuda,
354        );
355
356        assert_eq!(testing_cuda_htj2k_batch_decode_calls(), 1);
357        match result {
358            Ok(surfaces) => {
359                assert_eq!(surfaces.len(), inputs.len());
360                for surface in surfaces {
361                    assert_eq!(surface.residency(), SurfaceResidency::CudaResidentDecode);
362                    assert_eq!(surface.as_host_bytes(), None);
363                }
364            }
365            Err(Error::CudaUnavailable) => {
366                assert!(!cuda_runtime_required());
367            }
368            Err(error) => panic!("unexpected strict CUDA RGB batch error: {error}"),
369        }
370    }
371
372    fn rgb8_htj2k_fixture(width: u32, height: u32) -> Vec<u8> {
373        htj2k_rgb8_pattern_fixture(width, height, 17)
374    }
375}