diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 74a6eea..aff15a7 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -35,7 +35,7 @@ jobs: run: | python -m venv .venv .venv/bin/python -m pip install --upgrade pip - .venv/bin/pip install numpy pytest ruff 'maturin>=1.8,<2' + .venv/bin/pip install numpy pytest ruff 'maturin>=1.8,<2' 'opencv-python-headless>=4.13,<5' 'Pillow>=11.3' echo "$GITHUB_WORKSPACE/.venv/bin" >> "$GITHUB_PATH" echo "VIRTUAL_ENV=$GITHUB_WORKSPACE/.venv" >> "$GITHUB_ENV" - name: Configure prebuilt FFmpeg @@ -51,8 +51,8 @@ jobs: "$prefix/bin/ffmpeg" -version - name: Static checks run: | - ruff check src/tensorcodec tests scripts - ruff format --check src/tensorcodec tests scripts + ruff check src/tensorcodec tests scripts benchmarks/image_codecs.py + ruff format --check src/tensorcodec tests scripts benchmarks/image_codecs.py cargo fmt --manifest-path native/Cargo.toml --check cargo clippy --manifest-path native/Cargo.toml --locked -- -D warnings - name: Build extension against prebuilt FFmpeg diff --git a/.github/workflows/macos-wheels.yml b/.github/workflows/macos-wheels.yml index 3976ae6..2db8d2d 100644 --- a/.github/workflows/macos-wheels.yml +++ b/.github/workflows/macos-wheels.yml @@ -33,7 +33,7 @@ jobs: run: | brew install nasm pkg-config coreutils meson ninja uv venv - uv pip install numpy pytest 'maturin>=1.8,<2' delocate twine + uv pip install numpy pytest 'opencv-python-headless>=4.13,<5' 'Pillow>=11.3' 'maturin>=1.8,<2' delocate twine echo "$PWD/.venv/bin" >> "$GITHUB_PATH" echo "LIBCLANG_PATH=$(xcode-select -p)/Toolchains/XcodeDefault.xctoolchain/usr/lib" >> "$GITHUB_ENV" echo "$HOME/.pixi/envs/ffmpeg/bin" >> "$GITHUB_PATH" diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index 00b2543..b85973d 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -87,7 +87,7 @@ jobs: sudo apt-get update sudo apt-get install -y ffmpeg uv venv - uv pip install pytest numpy twine dist/*.whl + uv pip install pytest numpy twine 'opencv-python-headless>=4.13,<5' 'Pillow>=11.3' dist/*.whl - name: Generate baseline runtime fixtures run: uv run --no-sync python scripts/check_wheel_runtime.py generate .runtime-fixtures - name: Test on glibc 2.17 with Python 3.10 and 3.13 diff --git a/README.md b/README.md index 9de207a..e4612d3 100644 --- a/README.md +++ b/README.md @@ -3,6 +3,7 @@ # TensorCodec CPU video/audio decoding with TorchCodec-style APIs and NumPy output. +Optional image decoding and JPEG/PNG encoding reuse OpenCV.

CI @@ -14,7 +15,7 @@ CPU video/audio decoding with TorchCodec-style APIs and NumPy output. License: Apache-2.0

-[Quick start](#quick-start) · [Features](#features) · [Package size](#package-size) · [Compatibility](docs/compatibility.md) +[Quick start](#quick-start) · [Features](#features) · [Package size](#package-size) · [Compatibility](docs/compatibility.md) · [Image codecs](docs/images.md) diff --git a/benchmarks/image_codecs.py b/benchmarks/image_codecs.py new file mode 100644 index 0000000..eac0de4 --- /dev/null +++ b/benchmarks/image_codecs.py @@ -0,0 +1,65 @@ +"""Compare RGB decoding on a directory of JPEG/PNG/WebP files; run inside a uv venv.""" + +import argparse +import json +import os +import statistics +import time +from pathlib import Path + +import cv2 +import numpy as np +import torch +import torchcodec +from torchcodec.decoders import decode_image as reference + +from tensorcodec.decoders import decode_image + + +def latency(fn): + for _ in range(3): + fn() + start = time.perf_counter() + fn() + count = max(2, min(50, round(0.025 / (time.perf_counter() - start)))) + times = [] + for _ in range(5): + start = time.perf_counter() + for _ in range(count): + fn() + times.append((time.perf_counter() - start) * 1000 / count) + return statistics.median(times) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("corpus", type=Path) + args = parser.parse_args() + if hasattr(os, "sched_getaffinity"): + os.sched_setaffinity(0, {min(os.sched_getaffinity(0))}) + cv2.setNumThreads(1) + torch.set_num_threads(1) + print(json.dumps({"opencv": cv2.__version__, "torchcodec": torchcodec.__version__})) + for path in sorted(args.corpus.iterdir()): + if not path.is_file(): + continue + data = path.read_bytes() + if not (data.startswith((b"\xff\xd8\xff", b"\x89PNG")) or data[8:12] == b"WEBP"): + continue + actual, expected = decode_image(data), reference(data).numpy() + np.testing.assert_array_equal(actual, expected, err_msg=path.name) + print( + json.dumps( + { + "file": path.name, + "max_error": 0, + "tensorcodec_ms": latency(lambda data=data: decode_image(data)), + "torchcodec_ms": latency(lambda data=data: reference(data)), + } + ), + flush=True, + ) + + +if __name__ == "__main__": + main() diff --git a/docs/images.md b/docs/images.md new file mode 100644 index 0000000..4900da4 --- /dev/null +++ b/docs/images.md @@ -0,0 +1,79 @@ +# Optional CPU image codecs + +Image decoding and JPEG/PNG encoding use the user's OpenCV installation through +Python. NumPy remains the only required Python dependency; importing TensorCodec +does not import OpenCV or Pillow. No native image libraries or build steps are +added to TensorCodec. Pillow is used only to generate independent test inputs. + +Use an existing `cv2 >= 4.13` installation, or install one explicitly: + +```sh +uv pip install 'tensorcodec[images]' +``` + +The extra selects `opencv-python-headless`. Do not install multiple OpenCV wheel +variants in the same environment. OpenCV adds its own wheel size and dependencies; +it is not part of TensorCodec's advertised base wheel size. + +```python +from tensorcodec.decoders import decode_image, decode_jpeg +from tensorcodec.encoders import JpegEncoder, PngEncoder + +rgb = decode_image('input.webp') # uint8 CHW, RGB +frames = decode_image('animation.gif') # NCHW for multiple frames +batch = decode_jpeg(['one.jpg', 'two.jpg']) # list, possibly different sizes +PngEncoder(rgb).to_file('output.png', compression_level=6) +encoded = JpegEncoder(rgb).to_tensor(quality=90) # 1-D uint8 NumPy array +``` + +## Contract and limits + +- `decode_image`, `decode_jpeg`, `decode_png`, `decode_webp`, `decode_gif`, + `decode_avif` accept paths, bytes, bytearray or 1-D uint8 arrays. +- `mode` accepts case-insensitive `UNCHANGED`, `GRAY`, `GRAY_ALPHA`, `RGB` + (default), `RGB_ALPHA`/`RGBA`, or `ImageReadMode` values. +- `output_dtype` accepts uint8 (default), uint16 or `"auto"`. PNG preserves native + 8/16-bit precision with `"auto"`; explicit conversions scale the integer range. +- Multiple frames return NCHW; still images return CHW. Animated WebP keeps NCHW + even with one frame. Frame timings and loop counts are not returned. +- JPEG/PNG/WebP EXIF orientation and AVIF primary-item rotation/mirror are applied. + AVIF track-specific transforms are outside this adapter's contract. +- APNG, high-bit-depth AVIF, CMYK JPEG `UNCHANGED`, GPU decoding and nondefault + AVIF `num_threads` raise explicit errors. OpenCV has no per-call AVIF thread + setting; the adapter never modifies global OpenCV thread settings. +- AVIF color conversion follows OpenCV. Dropping alpha preserves straight RGB; + TorchCodec 0.17.0 premultiplies AVIF RGB in that case. Pixel identity with + TorchCodec is not promised across formats, builds or codec versions. +- HEIC is unsupported. There is no separate libheif or other decoder fallback. +- Other formats depend on the installed OpenCV build. Missing dependencies, + unsupported codecs and decode failures raise; no alternate decoder is tried. +- Encoders accept nonempty CHW uint8 arrays with 1 or 3 channels. Both provide + `to_file`, `to_file_like` and `to_tensor`; JPEG quality is 1–100 (default 75), + PNG compression level is 0–9 (default 6). Encoded bytes need not match TorchCodec. + +## Validation and performance + +Tests cover known PNG samples, independent Pillow decoding of encoder outputs, +orientation, animations, malformed input and TorchCodec 0.17.0 comparisons. +Pillow is a test dependency, not a runtime backend. + +On one Linux CPU, the adapter matched TorchCodec exactly for RGB output on 54 +JPEG/PNG/WebP inputs: photograph, graphics and seeded noise at 224 square, +640×480 and 1920×1080. OpenCV was 4.13.0.92. Timings used one pinned CPU, +one OpenCV/Torch thread, three warmups and the median of five batches. + +| Encoding | Adapter / TorchCodec latency, geometric mean | +| --- | ---: | +| JPEG 4:2:0 | 1.05 | +| JPEG 4:4:4 | 1.03 | +| Progressive JPEG | 1.02 | +| PNG | 1.28 | +| Lossy WebP | 1.07 | +| Lossless WebP | 1.11 | + +These corpus-specific results favor simplicity over specialized native backends; +OpenCV is not universally fastest. Encoding speed has not been benchmarked. +With the development and oracle dependencies installed, run +`python benchmarks/image_codecs.py CORPUS_DIRECTORY` inside a uv virtual +environment to measure in-memory decoding and exact RGB agreement. The script +prints versions and per-file results; disk reads are outside timed sections. diff --git a/pyproject.toml b/pyproject.toml index e3ce749..d502f21 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -17,12 +17,15 @@ classifiers = [ "Topic :: Multimedia :: Video", ] +[project.optional-dependencies] +images = ["opencv-python-headless>=4.13,<5"] + [project.urls] Repository = "https://github.com/MilkClouds/tensorcodec" Issues = "https://github.com/MilkClouds/tensorcodec/issues" [dependency-groups] -dev = ["pytest>=8", "ruff>=0.9", "maturin>=1.8,<2"] +dev = ["pytest>=8", "ruff>=0.9", "maturin>=1.8,<2", "opencv-python-headless>=4.13,<5", "Pillow>=11.3"] oracle = ["torch==2.14.1", "torchcodec==0.17.0"] [build-system] diff --git a/src/tensorcodec/_opencv.py b/src/tensorcodec/_opencv.py new file mode 100644 index 0000000..b20f34e --- /dev/null +++ b/src/tensorcodec/_opencv.py @@ -0,0 +1,14 @@ +"""Lazy access to the user's OpenCV installation.""" + + +def opencv(): + try: + import cv2 + except ImportError as exc: + raise ImportError( + "Image codecs require OpenCV >= 4.13; install tensorcodec[images] " + "or use an existing compatible cv2 installation" + ) from exc + if tuple(int(part) for part in cv2.__version__.split(".")[:2]) < (4, 13): + raise ImportError("Image codecs require OpenCV >= 4.13") + return cv2 diff --git a/src/tensorcodec/decoders/__init__.py b/src/tensorcodec/decoders/__init__.py index 7c13c8e..9de58eb 100644 --- a/src/tensorcodec/decoders/__init__.py +++ b/src/tensorcodec/decoders/__init__.py @@ -1,4 +1,26 @@ from tensorcodec._metadata import AudioStreamMetadata, VideoStreamMetadata from tensorcodec.decoders._decoder import AudioDecoder, CpuFallbackStatus, VideoDecoder +from tensorcodec.decoders._images import ( + ImageReadMode, + decode_avif, + decode_gif, + decode_image, + decode_jpeg, + decode_png, + decode_webp, +) -__all__ = ["AudioDecoder", "AudioStreamMetadata", "CpuFallbackStatus", "VideoDecoder", "VideoStreamMetadata"] +__all__ = [ + "AudioDecoder", + "AudioStreamMetadata", + "CpuFallbackStatus", + "ImageReadMode", + "VideoDecoder", + "VideoStreamMetadata", + "decode_avif", + "decode_gif", + "decode_image", + "decode_jpeg", + "decode_png", + "decode_webp", +] diff --git a/src/tensorcodec/decoders/_image_orientation.py b/src/tensorcodec/decoders/_image_orientation.py new file mode 100644 index 0000000..0a57955 --- /dev/null +++ b/src/tensorcodec/decoders/_image_orientation.py @@ -0,0 +1,165 @@ +"""JPEG/PNG/WebP EXIF and AVIF primary-item orientation.""" + +import numpy as np + + +def _boxes(data, start=0, end=None): + end = len(data) if end is None else end + while start < end: + if start + 8 > end: + raise RuntimeError("truncated AVIF box") + size, kind = int.from_bytes(data[start : start + 4], "big"), data[start + 4 : start + 8] + header = 8 + if size == 1: + if start + 16 > end: + raise RuntimeError("truncated extended AVIF box") + size, header = int.from_bytes(data[start + 8 : start + 16], "big"), 16 + elif size == 0: + size = end - start + if size < header or start + size > end: + raise RuntimeError("invalid AVIF box size") + yield kind, start + header, start + size + start += size + + +def _avif_orientation(data): + properties, associations, primary = [], [], None + for kind, begin, end in _boxes(data): + if kind != b"meta": + continue + for child, lo, hi in _boxes(data, begin + 4, end): + if child == b"pitm": + if lo + 4 > hi: + raise RuntimeError("truncated AVIF primary item") + width = 2 if data[lo] == 0 else 4 + if lo + 4 + width > hi: + raise RuntimeError("truncated AVIF primary item") + primary = int.from_bytes(data[lo + 4 : lo + 4 + width], "big") + if child != b"iprp": + continue + for prop, p, q in _boxes(data, lo, hi): + if prop == b"ipco": + properties = list(_boxes(data, p, q)) + elif prop == b"ipma": + if p + 8 > q: + raise RuntimeError("truncated AVIF property associations") + width = 2 if data[p] == 0 else 4 + entry_width = 2 if int.from_bytes(data[p + 1 : p + 4], "big") & 1 else 1 + count, cursor = int.from_bytes(data[p + 4 : p + 8], "big"), p + 8 + for _ in range(count): + if cursor + width + 1 > q: + raise RuntimeError("truncated AVIF item association") + item = int.from_bytes(data[cursor : cursor + width], "big") + n = data[cursor + width] + cursor += width + 1 + if cursor + n * entry_width > q: + raise RuntimeError("truncated AVIF property indices") + indices = [ + int.from_bytes(data[i : i + entry_width], "big") & ((1 << (entry_width * 8 - 1)) - 1) + for i in range(cursor, cursor + n * entry_width, entry_width) + ] + associations.append((item, indices)) + cursor += n * entry_width + angle, axis = 0, None + for item, indices in associations: + if item != primary: + continue + for index in indices: + if index == 0: + continue + if index > len(properties): + raise RuntimeError("invalid AVIF property index") + kind, start, end = properties[index - 1] + if kind in (b"irot", b"imir"): + if start == end: + raise RuntimeError("empty AVIF orientation property") + if kind == b"irot": + angle = data[start] & 3 + else: + axis = data[start] & 1 + # ISO/IEC 23008-12 item rotation/mirror mapped to TIFF's eight orientations. + return ((1, 4, 2), (8, 5, 7), (3, 2, 4), (6, 7, 5))[angle][0 if axis is None else axis + 1] + + +def _tiff_orientation(data): + if len(data) < 8 or data[:2] not in (b"II", b"MM"): + return 1 + order = "little" if data[:2] == b"II" else "big" + if int.from_bytes(data[2:4], order) != 42: + return 1 + offset = int.from_bytes(data[4:8], order) + if offset + 2 > len(data): + return 1 + count = int.from_bytes(data[offset : offset + 2], order) + for index in range(count): + entry = offset + 2 + index * 12 + if entry + 12 > len(data): + return 1 + if int.from_bytes(data[entry : entry + 2], order) == 274: + if ( + int.from_bytes(data[entry + 2 : entry + 4], order) != 3 + or int.from_bytes(data[entry + 4 : entry + 8], order) != 1 + ): + return 1 + value = int.from_bytes(data[entry + 8 : entry + 10], order) + return value if 1 <= value <= 8 else 1 + return 1 + + +def orientation(data, codec): + if codec == "avif": + return _avif_orientation(data) + if codec == "webp": + offset = 12 + while offset + 8 <= len(data): + size = int.from_bytes(data[offset + 4 : offset + 8], "little") + if offset + 8 + size > len(data): + break + if data[offset : offset + 4] == b"EXIF": + payload = data[offset + 8 : offset + 8 + size] + return _tiff_orientation(payload.removeprefix(b"Exif\0\0")) + offset += 8 + size + (size & 1) + elif codec == "png": + offset = 8 + while offset + 12 <= len(data): + size = int.from_bytes(data[offset : offset + 4], "big") + if offset + 12 + size > len(data): + break + if data[offset + 4 : offset + 8] == b"eXIf": + return _tiff_orientation(data[offset + 8 : offset + 8 + size]) + offset += size + 12 + elif codec == "jpeg": + offset = 2 + while offset < len(data) and data[offset] == 0xFF: + while offset < len(data) and data[offset] == 0xFF: + offset += 1 + if offset >= len(data) or data[offset] in (0xDA, 0xD9): + break + marker, offset = data[offset], offset + 1 + size = int.from_bytes(data[offset : offset + 2], "big") + if size < 2 or offset + size > len(data): + break + payload = data[offset + 2 : offset + size] + if marker == 0xE1 and payload.startswith(b"Exif\0\0"): + return _tiff_orientation(payload[6:]) + offset += size + return 1 + + +def apply_orientation(images, value): + # NHWC throughout; the public wrapper moves channels only after orientation. + if value == 2: + return images[:, :, ::-1] + if value == 3: + return images[:, ::-1, ::-1] + if value == 4: + return images[:, ::-1] + if value == 5: + return images.swapaxes(1, 2) + if value == 6: + return np.rot90(images, -1, axes=(1, 2)) + if value == 7: + return images.swapaxes(1, 2)[:, ::-1, ::-1] + if value == 8: + return np.rot90(images, 1, axes=(1, 2)) + return images diff --git a/src/tensorcodec/decoders/_images.py b/src/tensorcodec/decoders/_images.py new file mode 100644 index 0000000..328371d --- /dev/null +++ b/src/tensorcodec/decoders/_images.py @@ -0,0 +1,301 @@ +"""CPU image decoding with TorchCodec's function API and NumPy output.""" + +from enum import Enum +from pathlib import Path + +import numpy as np + +from tensorcodec._opencv import opencv +from tensorcodec.decoders._image_orientation import apply_orientation, orientation + + +class ImageReadMode(Enum): + UNCHANGED = 0 + GRAY = 1 + GRAY_ALPHA = 2 + RGB = 3 + RGB_ALPHA = 4 + RGBA = RGB_ALPHA + + +def _mode(value): + if isinstance(value, ImageReadMode): + return value + if isinstance(value, str): + try: + return ImageReadMode[value.upper()] + except KeyError: + raise ValueError(f"Invalid image mode: {value!r}") from None + raise TypeError("mode must be a string or ImageReadMode") + + +def _dtype(value): + if isinstance(value, str) and value == "auto": + return value + try: + dtype = np.dtype(value) + except (TypeError, ValueError): + raise ValueError("output_dtype must be uint8, uint16 or 'auto'") from None + if dtype not in (np.dtype("uint8"), np.dtype("uint16")): + raise ValueError("output_dtype must be uint8, uint16 or 'auto'") + return dtype + + +def _bytes(source): + if isinstance(source, (str, Path)): + return Path(source).read_bytes() + if isinstance(source, (bytes, bytearray)): + return bytes(source) + if isinstance(source, np.ndarray): + if source.dtype != np.uint8 or source.ndim != 1: + raise ValueError("encoded image arrays must be one-dimensional uint8") + return source.tobytes() + raise TypeError("source must be a path, bytes or a one-dimensional uint8 NumPy array") + + +def _format(data): + if data.startswith(b"\xff\xd8\xff"): + return "jpeg" + if data.startswith(b"\x89PNG\r\n\x1a\n"): + return "png" + if data[:6] in (b"GIF87a", b"GIF89a"): + return "gif" + if data[:4] == b"RIFF" and data[8:12] == b"WEBP": + return "webp" + if data[4:8] == b"ftyp": + size = int.from_bytes(data[:4], "big") + brands = [data[8:12], *[data[i : i + 4] for i in range(16, min(size, len(data)), 4)]] + if any(b in (b"avif", b"avis") for b in brands): + return "avif" + if any(b in (b"heic", b"heix", b"heim", b"heis", b"hevc", b"hevx", b"mif1", b"msf1") for b in brands): + return "heic" + raise ValueError("Unsupported or unrecognized image format") + + +def _png_channels(data): + channels = None + offset = 8 + while offset + 12 <= len(data): + size = int.from_bytes(data[offset : offset + 4], "big") + kind = data[offset + 4 : offset + 8] + if offset + size + 12 > len(data): + raise RuntimeError("truncated PNG chunk") + if kind == b"IHDR" and size == 13: + channels = {0: 1, 2: 3, 3: 3, 4: 2, 6: 4}.get(data[offset + 17]) + elif kind == b"tRNS": + channels = 2 if channels == 1 else 4 + elif kind == b"acTL": + raise NotImplementedError("animated PNG decoding is unsupported") + elif kind == b"IDAT": + break + offset += size + 12 + if channels is None: + raise RuntimeError("PNG has no valid IHDR") + return channels + + +def _jpeg_components(data): + offset = 2 + while offset < len(data): + if data[offset] != 0xFF: + break + while offset < len(data) and data[offset] == 0xFF: + offset += 1 + if offset >= len(data): + break + marker, offset = data[offset], offset + 1 + if marker in (0xDA, 0xD9): + break + length = int.from_bytes(data[offset : offset + 2], "big") + if length < 2 or offset + length > len(data): + break + if marker in (0xC0, 0xC1, 0xC2, 0xC3, 0xC5, 0xC6, 0xC7, 0xC9, 0xCA, 0xCB, 0xCD, 0xCE, 0xCF): + if length < 8: + break + return data[offset + 7] + offset += length + raise RuntimeError("JPEG has no valid frame header") + + +def _check_webp(data): + animated = False + offset = 12 + while offset + 8 <= len(data): + size = int.from_bytes(data[offset + 4 : offset + 8], "little") + if data[offset : offset + 4] in (b"ANIM", b"ANMF"): + animated = True + if offset + 8 + size > len(data): + raise RuntimeError("truncated WebP chunk") + offset += 8 + size + (size & 1) + return animated + + +def _color(images, mode, codec): + channels = images.shape[-1] + if mode is ImageReadMode.UNCHANGED: + return images + if channels == mode.value: + return images + gray_source = channels < 3 + color = images[..., :1] if gray_source else images[..., :3] + alpha = images[..., -1:] if channels in (2, 4) else np.full_like(images[..., :1], np.iinfo(images.dtype).max) + if mode in (ImageReadMode.GRAY, ImageReadMode.GRAY_ALPHA): + if not gray_source: + if codec == "png": + # TorchCodec requests libpng coefficients 0.2989 and 0.587. + color = ( + ( + (color.astype(np.uint64) * np.array([9794, 19234, 3740], np.uint64)).sum( + axis=-1, keepdims=True + ) + + (16384 if images.dtype == np.uint16 else 0) + ) + >> 15 + ).astype(images.dtype) + else: + color = ( + np.rint( + (color.astype(np.float32) * np.array([0.2989, 0.587, 0.114], np.float32)).sum( + axis=-1, keepdims=True + ) + ) + .clip(0, np.iinfo(images.dtype).max) + .astype(images.dtype) + ) + elif gray_source: + color = np.repeat(color, 3, axis=-1) + return ( + np.concatenate((color, alpha), axis=-1) + if mode in (ImageReadMode.GRAY_ALPHA, ImageReadMode.RGB_ALPHA) + else color + ) + + +def _image(source, codec, mode, output_dtype): + mode, dtype = _mode(mode), _dtype(output_dtype) + data = _bytes(source) + found = _format(data) + if codec is not None and found != codec: + raise RuntimeError(f"expected {codec}, got {found}") + codec = found + orient = orientation(data, codec) + animated = False + channels = _png_channels(data) if codec == "png" else None + if codec == "heic": + raise NotImplementedError("HEIC decoding is not supported by the OpenCV image adapter") + cv = opencv() + flags = cv.IMREAD_UNCHANGED + if codec == "jpeg": + if _jpeg_components(data) == 4 and mode is ImageReadMode.UNCHANGED: + raise NotImplementedError("OpenCV cannot preserve CMYK JPEG channels") + if mode in (ImageReadMode.GRAY, ImageReadMode.GRAY_ALPHA): + flags = cv.IMREAD_GRAYSCALE | cv.IMREAD_IGNORE_ORIENTATION + if codec == "webp": + animated = _check_webp(data) + encoded = np.frombuffer(data, np.uint8).reshape(1, -1) + try: + if codec in ("gif", "webp", "avif"): + ok, frames = cv.imdecodemulti(encoded, flags) + if not ok or not frames: + raise RuntimeError(f"OpenCV could not decode {codec}; check its codec build support") + else: + frame = cv.imdecode(encoded, flags) + if frame is None: + raise RuntimeError(f"OpenCV could not decode {codec}") + frames = [frame] + except cv.error as exc: + raise RuntimeError(f"OpenCV failed to decode {codec}: {exc}") from exc + converted = [] + for frame in frames: + if frame.ndim == 2: + frame = frame[..., None] + elif frame.shape[-1] == 3: + frame = cv.cvtColor(frame, cv.COLOR_BGR2RGB) + elif frame.shape[-1] == 4: + frame = cv.cvtColor(frame, cv.COLOR_BGRA2RGBA) + if channels == 2 and frame.shape[-1] == 4: + frame = frame[..., [0, 3]] + if codec == "avif" and frame.dtype == np.uint16: + raise NotImplementedError("high-bit-depth AVIF is not supported by this OpenCV adapter") + if codec == "png" and data[25] in (0, 2): + # OpenCV drops grayscale tRNS; UNCHANGED retains original non-palette channels. + base = 1 if data[25] == 0 else 3 + frame = frame[..., :base] + if mode in (ImageReadMode.GRAY_ALPHA, ImageReadMode.RGB_ALPHA): + offset = 8 + while offset + 12 <= len(data): + size = int.from_bytes(data[offset : offset + 4], "big") + if data[offset + 4 : offset + 8] == b"tRNS": + key = np.frombuffer(data[offset + 8 : offset + 8 + size], dtype=">u2") + if len(key) != base: + raise RuntimeError("invalid PNG transparency key") + if data[24] < 8: + key = key * (255 // ((1 << data[24]) - 1)) + alpha = np.where( + np.all(frame == key, axis=-1, keepdims=True), 0, np.iinfo(frame.dtype).max + ).astype(frame.dtype) + frame = np.concatenate((frame, alpha), axis=-1) + break + offset += size + 12 + converted.append(frame) + if any(f.shape != converted[0].shape or f.dtype != converted[0].dtype for f in converted): + raise RuntimeError("image frames have different shapes or bit depths") + images = converted[0][None] if len(converted) == 1 else np.stack(converted) + images = _color(images, mode, codec) + if dtype != "auto" and images.dtype != dtype: + if dtype == np.uint16: + images = images.astype(np.uint16) * 257 + else: + images = np.rint(images.astype(np.float32) / 257).clip(0, 255).astype(np.uint8) + images = apply_orientation(images, orient).transpose(0, 3, 1, 2) + keep_batch = codec == "webp" and animated + return images[0] if len(images) == 1 and not keep_batch else images + + +def decode_image(source, *, mode="RGB", output_dtype=np.uint8): + """Detect encoded content and return a CHW image or NCHW animation on CPU. + + Sources are paths, bytes or 1-D uint8 arrays. Modes: UNCHANGED, GRAY, + GRAY_ALPHA, RGB, RGB_ALPHA (case-insensitive strings or ImageReadMode). + output_dtype is uint8, uint16 or 'auto'; integer conversion scales the range. + Requires optional OpenCV >= 4.13. HEIC and animated PNG are unsupported. + """ + return _image(source, None, mode, output_dtype) + + +def decode_jpeg(source, *, mode="RGB", output_dtype=np.uint8, device="cpu"): + """Decode a JPEG to CHW, or a list/tuple of sources to a list of CHW arrays. + + Only CPU decoding is supported. Batches may have different image dimensions. + """ + if str(device) != "cpu": + raise ValueError("only CPU image decoding is supported") + if isinstance(source, (list, tuple)): + _mode(mode) + _dtype(output_dtype) + return [_image(item, "jpeg", mode, output_dtype) for item in source] + return _image(source, "jpeg", mode, output_dtype) + + +def decode_png(source, *, mode="RGB", output_dtype=np.uint8): + """Decode a PNG to CHW; 'auto' preserves native 8/16-bit sample precision.""" + return _image(source, "png", mode, output_dtype) + + +def decode_webp(source, *, mode="RGB", output_dtype=np.uint8): + """Decode WebP to CHW/NCHW. Animations use OpenCV compositing.""" + return _image(source, "webp", mode, output_dtype) + + +def decode_gif(source, *, mode="RGB", output_dtype=np.uint8): + """Decode GIF to CHW for one frame or NCHW for multiple frames.""" + return _image(source, "gif", mode, output_dtype) + + +def decode_avif(source, *, mode="RGB", output_dtype=np.uint8, num_threads=1): + """Decode AVIF to CHW/NCHW. Only the default num_threads=1 is accepted.""" + if not isinstance(num_threads, int) or isinstance(num_threads, bool) or num_threads < 1: + raise ValueError("num_threads must be a positive integer") + if num_threads != 1: + raise NotImplementedError("OpenCV exposes no per-call AVIF thread control") + return _image(source, "avif", mode, output_dtype) diff --git a/src/tensorcodec/encoders/__init__.py b/src/tensorcodec/encoders/__init__.py new file mode 100644 index 0000000..cc30477 --- /dev/null +++ b/src/tensorcodec/encoders/__init__.py @@ -0,0 +1,5 @@ +"""CPU image encoders with NumPy input.""" + +from tensorcodec.encoders._images import JpegEncoder, PngEncoder + +__all__ = ["JpegEncoder", "PngEncoder"] diff --git a/src/tensorcodec/encoders/_images.py b/src/tensorcodec/encoders/_images.py new file mode 100644 index 0000000..98a96ae --- /dev/null +++ b/src/tensorcodec/encoders/_images.py @@ -0,0 +1,59 @@ +"""TorchCodec-style image encoding delegated to OpenCV.""" + +from pathlib import Path + +import numpy as np + +from tensorcodec._opencv import opencv + + +class _ImageEncoder: + def __init__(self, img): + if not isinstance(img, np.ndarray): + raise TypeError("img must be a NumPy array") + if img.dtype != np.uint8 or img.ndim != 3 or img.shape[0] not in (1, 3) or 0 in img.shape: + raise ValueError("img must be a nonempty CHW uint8 array with 1 or 3 channels") + self.img = img + + def _encode(self, extension, parameter, value, low, high): + if isinstance(value, bool) or not isinstance(value, int) or not low <= value <= high: + raise ValueError(f"encoding parameter must be an integer in [{low}, {high}]") + cv = opencv() + pixels = self.img[0] if self.img.shape[0] == 1 else self.img.transpose(1, 2, 0)[..., ::-1] + try: + ok, encoded = cv.imencode(extension, np.ascontiguousarray(pixels), [getattr(cv, parameter), value]) + except cv.error as exc: + raise RuntimeError(f"OpenCV failed to encode {extension}: {exc}") from exc + if not ok: + raise RuntimeError(f"OpenCV could not encode {extension}") + return encoded.reshape(-1) + + def to_file(self, dest, **kwargs): + """Write encoded bytes to a filesystem path.""" + Path(dest).write_bytes(self.to_tensor(**kwargs).tobytes()) + + def to_file_like(self, dest, **kwargs): + """Write encoded bytes to a binary file-like object.""" + data = self.to_tensor(**kwargs).tobytes() + offset = 0 + while offset < len(data): + written = dest.write(data[offset:]) + if not isinstance(written, int) or written <= 0 or written > len(data) - offset: + raise OSError("file-like object did not report a valid write length") + offset += written + + +class JpegEncoder(_ImageEncoder): + """Encode a CHW uint8 grayscale/RGB image on CPU.""" + + def to_tensor(self, *, quality=75): + """Return encoded JPEG bytes as a one-dimensional uint8 NumPy array.""" + return self._encode(".jpg", "IMWRITE_JPEG_QUALITY", quality, 1, 100) + + +class PngEncoder(_ImageEncoder): + """Encode a CHW uint8 grayscale/RGB image on CPU.""" + + def to_tensor(self, *, compression_level=6): + """Return encoded PNG bytes as a one-dimensional uint8 NumPy array.""" + return self._encode(".png", "IMWRITE_PNG_COMPRESSION", compression_level, 0, 9) diff --git a/tests/test_image_encoders.py b/tests/test_image_encoders.py new file mode 100644 index 0000000..46ea1f5 --- /dev/null +++ b/tests/test_image_encoders.py @@ -0,0 +1,118 @@ +"""Image encoder round trips, stream writes and optional dependency boundaries.""" + +import subprocess +import sys +from io import BytesIO + +import numpy as np +import pytest +from PIL import Image + +from tensorcodec.decoders import decode_image +from tensorcodec.encoders import JpegEncoder, PngEncoder + + +@pytest.mark.parametrize("encoder", [JpegEncoder, PngEncoder]) +@pytest.mark.parametrize("channels", [1, 3]) +def test_encoder_outputs(encoder, channels, tmp_path): + pixels = np.full((channels, 13, 17), 71, np.uint8) + if channels == 3: + pixels[:] = np.array([23, 91, 177])[:, None, None] + obj = encoder(pixels) + result = obj.to_tensor() + assert result.ndim == 1 and result.dtype == np.uint8 + independent = np.array(Image.open(BytesIO(result.tobytes()))) + expected = pixels[0] if channels == 1 else pixels.transpose(1, 2, 0) + np.testing.assert_allclose(independent.astype(int), expected.astype(int), atol=2 if encoder is JpegEncoder else 0) + path, stream = tmp_path / "image.bin", BytesIO() + obj.to_file(path) + obj.to_file_like(stream) + assert path.read_bytes() == stream.getvalue() == result.tobytes() + np.testing.assert_allclose(decode_image(result, mode="UNCHANGED").astype(int), pixels.astype(int), atol=2) + + +def test_png_noncontiguous_input_and_partial_writes(): + pixels = np.arange(3 * 13 * 17, dtype=np.uint8).reshape(3, 13, 17)[:, ::-1, ::2] + original = pixels.copy() + encoder = PngEncoder(pixels) + + class Partial(BytesIO): + def write(self, data): + return super().write(data[:7]) + + stream = Partial() + encoder.to_file_like(stream, compression_level=9) + np.testing.assert_array_equal(decode_image(stream.getvalue()), original) + np.testing.assert_array_equal(pixels, original) + + +@pytest.mark.parametrize( + "encoder,key,values", + [ + (JpegEncoder, "quality", [0, 101, True, 1.5]), + (PngEncoder, "compression_level", [-1, 10, True, 1.5]), + ], +) +def test_encoder_parameter_errors(encoder, key, values): + obj = encoder(np.zeros((3, 2, 2), np.uint8)) + for value in values: + with pytest.raises(ValueError): + obj.to_tensor(**{key: value}) + + +@pytest.mark.parametrize( + "img", + [ + np.zeros((2, 2), np.uint8), + np.zeros((4, 2, 2), np.uint8), + np.zeros((3, 0, 2), np.uint8), + np.zeros((3, 2, 2), np.uint16), + ], +) +def test_encoder_input_errors(img): + with pytest.raises(ValueError): + PngEncoder(img) + + +def test_nonprogressing_writer(): + class Writer: + def write(self, data): + return None + + with pytest.raises(OSError, match="write length"): + PngEncoder(np.zeros((3, 2, 2), np.uint8)).to_file_like(Writer()) + + +def test_cv2_is_lazy_and_optional(): + subprocess.run( + [ + sys.executable, + "-c", + """ +import sys +import tensorcodec.decoders +import tensorcodec.encoders +assert 'cv2' not in sys.modules +assert 'PIL' not in sys.modules +sys.modules['cv2'] = None +import numpy as np +try: + tensorcodec.encoders.PngEncoder(np.zeros((3, 2, 2), np.uint8)).to_tensor() +except ImportError as exc: + assert 'tensorcodec[images]' in str(exc) +else: + raise AssertionError('missing OpenCV was silently bypassed') +""", + ], + check=True, + ) + + +def test_no_silent_decode_fallback(monkeypatch): + import cv2 + + from tests.test_images import png_bytes + + monkeypatch.setattr(cv2, "imdecode", lambda *args: None) + with pytest.raises(RuntimeError, match="OpenCV"): + decode_image(png_bytes(np.zeros((2, 2, 3), np.uint8))) diff --git a/tests/test_images.py b/tests/test_images.py new file mode 100644 index 0000000..d990d19 --- /dev/null +++ b/tests/test_images.py @@ -0,0 +1,511 @@ +"""Image API contracts from known PNG samples and independent FFmpeg encodings.""" + +import struct +import zlib +from io import BytesIO +from pathlib import Path + +import numpy as np +import pytest +from PIL import Image + +from tests.utils import as_numpy, run_ffmpeg + + +def chunk(kind, data): + return struct.pack(">I", len(data)) + kind + data + struct.pack(">I", zlib.crc32(kind + data)) + + +def png_bytes(pixels): + height, width, channels = pixels.shape + depth = pixels.dtype.itemsize * 8 + color = {1: 0, 2: 4, 3: 2, 4: 6}[channels] + + raw = pixels.astype(">u2" if depth == 16 else np.uint8).tobytes() + stride = width * channels * (depth // 8) + scanlines = b"".join(b"\0" + raw[i : i + stride] for i in range(0, len(raw), stride)) + return ( + b"\x89PNG\r\n\x1a\n" + + chunk(b"IHDR", struct.pack(">IIBBBBB", width, height, depth, color, 0, 0, 0)) + + chunk(b"IDAT", zlib.compress(scanlines)) + + chunk(b"IEND", b"") + ) + + +def dtype_for(backend, name): + if name == "auto": + return name + if backend.__name__.startswith("torchcodec"): + import torch + + return getattr(torch, name) + return getattr(np, name) + + +@pytest.fixture(scope="module") +def encoded_images(tmp_path_factory): + root = tmp_path_factory.mktemp("images") + pixels = np.arange(16 * 24 * 3, dtype=np.uint8).reshape(16, 24, 3) + (root / "source.png").write_bytes(png_bytes(pixels)) + for codec, options in { + "jpeg": ["-c:v", "mjpeg", "-q:v", "1", "-pix_fmt", "yuvj444p"], + "gif": [], + "avif": ["-c:v", "libaom-av1", "-still-picture", "1", "-crf", "0", "-cpu-used", "8"], + }.items(): + run_ffmpeg("-i", root / "source.png", "-frames:v", 1, *options, "-threads", 1, root / f"image.{codec}") + Image.fromarray(pixels).save(root / "image.webp", lossless=True) + (root / "second.png").write_bytes(png_bytes(255 - pixels)) + run_ffmpeg( + "-framerate", + 2, + "-pattern_type", + "glob", + "-i", + str(root / "s*.png"), + "-threads", + 1, + root / "animated.gif", + ) + return root + + +@pytest.mark.parametrize("channels", [1, 2, 3, 4]) +@pytest.mark.parametrize("depth", [8, 16]) +def test_png_native_samples(backend, channels, depth): + dtype = np.uint8 if depth == 8 else np.uint16 + pixels = np.arange(5 * 7 * channels, dtype=dtype).reshape(5, 7, channels) + if depth == 16: + pixels = pixels * 311 + decoded = backend.decode_png(png_bytes(pixels), mode="UNCHANGED", output_dtype="auto") + np.testing.assert_array_equal(as_numpy(decoded), pixels.transpose(2, 0, 1)) + assert as_numpy(decoded).dtype == dtype + + +@pytest.mark.parametrize("mode,channels", [("RGB", 3), ("GRAY", 1), ("RGB_ALPHA", 4), ("GRAY_ALPHA", 2)]) +def test_png_color_modes(backend, mode, channels): + pixels = np.array([[[255, 0, 0, 17], [0, 255, 0, 93], [0, 0, 255, 201]]], np.uint8) + result = as_numpy(backend.decode_image(png_bytes(pixels), mode=mode.lower())) + assert result.shape == (channels, 1, 3) + if "ALPHA" in mode: + np.testing.assert_array_equal(result[-1], pixels[..., 3]) + if mode.startswith("RGB"): + np.testing.assert_array_equal(result[:3], pixels[..., :3].transpose(2, 0, 1)) + else: + np.testing.assert_array_equal(result[0], [[76, 149, 29]]) + + +def test_output_dtype_scales_values(backend): + pixels = np.array([[[0], [1], [128], [255]]], np.uint8) + result = backend.decode_png(png_bytes(pixels), mode="GRAY", output_dtype=dtype_for(backend, "uint16")) + np.testing.assert_array_equal(as_numpy(result), pixels.transpose(2, 0, 1).astype(np.uint16) * 257) + pixels = np.array([[[0], [128], [129], [32768], [65535]]], np.uint16) + result = backend.decode_png(png_bytes(pixels), mode="GRAY") + np.testing.assert_array_equal(as_numpy(result), np.rint(pixels.transpose(2, 0, 1) / 257).astype(np.uint8)) + + +@pytest.mark.parametrize("kind", ["bytes", "bytearray", "path", "str", "array"]) +def test_image_sources_and_content_detection(backend, tmp_path, kind): + pixels = np.arange(27, dtype=np.uint8).reshape(3, 3, 3) + data = png_bytes(pixels) + path = tmp_path / "misleading.jpg" + path.write_bytes(data) + sources = {"bytes": data, "bytearray": bytearray(data), "path": path, "str": str(path)} + if kind == "array": + if backend.__name__.startswith("torchcodec"): + import torch + + source = torch.frombuffer(bytearray(data), dtype=torch.uint8) + else: + source = np.frombuffer(data, dtype=np.uint8) + else: + source = sources[kind] + np.testing.assert_array_equal(as_numpy(backend.decode_image(source)), pixels.transpose(2, 0, 1)) + + +@pytest.mark.parametrize("codec", ["jpeg", "webp", "gif", "avif"]) +def test_format_functions_and_dispatch(backend, encoded_images, codec): + path = encoded_images / f"image.{codec}" + direct = as_numpy(getattr(backend, f"decode_{codec}")(path)) + assert direct.shape == (3, 16, 24) + assert direct.dtype == np.uint8 + np.testing.assert_array_equal(direct, as_numpy(backend.decode_image(path.read_bytes()))) + + +def test_jpeg_batch(backend, encoded_images): + source = encoded_images / "image.jpeg" + images = backend.decode_jpeg([source, source.read_bytes()]) + assert isinstance(images, list) + assert len(images) == 2 + np.testing.assert_array_equal(as_numpy(images[0]), as_numpy(images[1])) + images[0][...] = 0 + assert as_numpy(images[1]).any() + assert backend.decode_jpeg([]) == [] + + +def test_gif_animation(backend, encoded_images): + frames = as_numpy(backend.decode_gif(encoded_images / "animated.gif")) + assert frames.shape == (2, 3, 16, 24) + assert not np.array_equal(frames[0], frames[1]) + + +@pytest.mark.parametrize("mode", ["RGB", "UNCHANGED", "GRAY", "GRAY_ALPHA", "RGB_ALPHA"]) +@pytest.mark.parametrize("codec", ["jpeg", "png", "webp", "gif", "avif"]) +def test_image_differential(oracle, encoded_images, codec, mode): + import tensorcodec.decoders as actual + + path = encoded_images / ("source.png" if codec == "png" else f"image.{codec}") + got = getattr(actual, f"decode_{codec}")(path, mode=mode) + expected = as_numpy(getattr(oracle, f"decode_{codec}")(path, mode=mode)) + assert got.shape == expected.shape + np.testing.assert_allclose(got.astype(np.int32), expected.astype(np.int32), atol=2, rtol=0) + + +def test_image_errors_and_cpu_boundary(): + from tensorcodec.decoders import decode_avif, decode_image, decode_jpeg, decode_png + + data = png_bytes(np.zeros((2, 3, 3), np.uint8)) + with pytest.raises(ValueError, match="mode"): + decode_png(data, mode="BGR") + with pytest.raises(ValueError, match="output_dtype"): + decode_png(data, output_dtype=np.float32) + with pytest.raises(ValueError, match="one-dimensional"): + decode_png(np.zeros((2, 3), np.uint8)) + with pytest.raises(TypeError, match="source"): + decode_image(object()) + with pytest.raises(ValueError, match="CPU"): + decode_jpeg(b"", device="cuda") + with pytest.raises(RuntimeError, match="expected jpeg"): + decode_jpeg(data) + with pytest.raises(ValueError, match="unrecognized"): + decode_image(b"not an image") + with pytest.raises(ValueError, match="num_threads"): + decode_avif(b"", num_threads=0) + with pytest.raises(FileNotFoundError): + decode_image(Path("/nonexistent/image.png")) + + +def test_heic_is_explicitly_unsupported(): + from tensorcodec.decoders import decode_image + + data = struct.pack(">I", 24) + b"ftypheic" + b"\0" * 4 + b"mif1heic" + with pytest.raises(NotImplementedError, match="HEIC"): + decode_image(data) + + +@pytest.fixture(scope="module") +def optional_images(tmp_path_factory): + root = tmp_path_factory.mktemp("optional-images") + pixels = np.zeros((16, 24, 4), np.uint8) + pixels[..., :3] = [30, 60, 90] + pixels[..., 3] = 255 + pixels[:8, :12, 3] = 0 + image = Image.fromarray(pixels) + second = image.copy() + second.paste((120, 50, 10, 128), (8, 4, 16, 12)) + image.save(root / "animation.webp", save_all=True, append_images=[second], lossless=True, duration=100) + return root + + +@pytest.mark.parametrize("mode", ["RGB", "UNCHANGED", "GRAY", "GRAY_ALPHA", "RGB_ALPHA"]) +def test_animated_webp_differential(oracle, optional_images, mode): + from tensorcodec.decoders import decode_image + + path = optional_images / "animation.webp" + got = decode_image(path, mode=mode) + expected = as_numpy(oracle.decode_image(path, mode=mode)) + assert got.shape == expected.shape + np.testing.assert_allclose(got.astype(np.int32), expected.astype(np.int32), atol=2, rtol=0) + + +def test_single_frame_animated_webp_keeps_batch_dimension(oracle, optional_images): + from tensorcodec.decoders import decode_webp + + data = (optional_images / "animation.webp").read_bytes() + chunks, offset, frames = [], 12, 0 + while offset + 8 <= len(data): + size = int.from_bytes(data[offset + 4 : offset + 8], "little") + end = offset + 8 + size + (size & 1) + if data[offset : offset + 4] == b"ANMF": + frames += 1 + if frames < 2: + chunks.append(data[offset:end]) + offset = end + body = b"WEBP" + b"".join(chunks) + encoded = b"RIFF" + struct.pack("u2").tobytes()) + data = ( + b"\x89PNG\r\n\x1a\n" + + chunk(b"IHDR", struct.pack(">IIBBBBB", 37, 29, 16, 6, 0, 0, 1)) + + chunk(b"IDAT", zlib.compress(raw)) + + chunk(b"IEND", b"") + ) + for mode in ("UNCHANGED", "RGB", "RGB_ALPHA", "GRAY", "GRAY_ALPHA"): + actual = decode_png(data, mode=mode, output_dtype="auto") + expected = as_numpy(oracle.decode_png(data, mode=mode, output_dtype="auto")) + np.testing.assert_array_equal(actual, expected) + + +@pytest.mark.parametrize("bits", [1, 2, 4]) +def test_low_bit_palette_png(oracle, bits): + from tensorcodec.decoders import decode_png + + image = Image.new("P", (37, 29)) + image.putpalette(list(range(256)) * 3) + image.putdata((np.arange(37 * 29) % (2**bits)).tolist()) + buffer = BytesIO() + image.save(buffer, format="PNG", bits=bits, transparency=0) + for mode in ("UNCHANGED", "RGB", "RGB_ALPHA", "GRAY", "GRAY_ALPHA"): + np.testing.assert_array_equal( + decode_png(buffer.getvalue(), mode=mode), as_numpy(oracle.decode_png(buffer.getvalue(), mode=mode)) + ) + + +@pytest.mark.parametrize("codec", ["WEBP", "GIF", "AVIF"]) +def test_image_animation_all_frames(oracle, codec): + from tensorcodec.decoders import decode_image + + rng = np.random.default_rng(42) + frames = [Image.fromarray(rng.integers(0, 256, (29, 37, 3), dtype=np.uint8)) for _ in range(3)] + output = BytesIO() + frames[0].save(output, format=codec, save_all=True, append_images=frames[1:], duration=100, max_threads=1) + actual = decode_image(output.getvalue()) + assert actual.shape == (3, 3, 29, 37) + if codec == "AVIF": + # AVIF YUV conversion belongs to the installed OpenCV/libavif build. + import cv2 + + ok, decoded = cv2.imdecodemulti( + np.frombuffer(output.getvalue(), np.uint8).reshape(1, -1), cv2.IMREAD_UNCHANGED + ) + assert ok + expected = np.stack([frame[..., ::-1].transpose(2, 0, 1) for frame in decoded]) + else: + expected = as_numpy(oracle.decode_image(output.getvalue())) + np.testing.assert_array_equal(actual, expected) + + +@pytest.mark.parametrize("bits", [1, 2, 4]) +def test_low_bit_grayscale_transparency(bits): + from tensorcodec.decoders import decode_png + + data = ( + b"\x89PNG\r\n\x1a\n" + + chunk(b"IHDR", struct.pack(">IIBBBBB", 2, 1, bits, 0, 0, 0, 0)) + + chunk(b"tRNS", struct.pack(">H", 1)) + + chunk(b"IDAT", zlib.compress(bytes([0, 1 << (8 - bits)]))) + + chunk(b"IEND", b"") + ) + gray = 255 // ((1 << bits) - 1) + for mode, channels in (("GRAY_ALPHA", 1), ("RGB_ALPHA", 3)): + result = decode_png(data, mode=mode) + np.testing.assert_array_equal(result[:-1], np.tile([[[gray, 0]]], (channels, 1, 1))) + np.testing.assert_array_equal(result[-1], [[0, 255]])