google-deepmind/dm_pix

PIX is an image processing library in JAX, for JAX.

What is PIX?

PIX 是一个轻量级图像处理库,构建在 JAX 之上,JAX 是驱动许多现代机器学习研究的高性能 NumPy 兼容框架。PIX 的所有函数(例如:翻转、缩放、色彩转换)在编写时都考虑到了编译与加速。它们可以使用 jax.jit 进行编译、使用 jax.vmap 进行向量化,或使用 jax.pmap 在多个设备上运行。在实践中,这意味着您可以将图像视为任何其他 JAX 阵列,并获得与神经网络-神经网络代码相同的 GPU/TPU 加速效果。


Quick start

import dm_pix as pix

# Load an image with any library that returns a NumPy array, e.g. Pillow, OpenCV, etc.
image = load_image()

# Simple left‑right flip – works on CPU, GPU or TPU.
flipped = pix.flip_left_right(image)

因为 pix.flip_left_right 是一个纯 JAX 函数,您也可以这样做:

import jax

# Compile once for maximum speed.
flipped = jax.jit(pix.flip_left_right)(image)

# Apply to a batch of images without writing a loop.
batch = image[None, ...]               # add aation leading batch dimension
flipped_batch = jax.vmap(pix.flip_left_right)(batch)

Installation

  1. Install JAX first – follow the official JAX installation guide and pick the version that matches your CUDA/TPU setup.
  2. Install PIX from PyPI:
    pip install dm-pix
    
    (PIX 本身是纯 Python;它只需要 JAX 运行时,这就是为什么 JAX 没有被列为自动依赖项。)

What can you do with PIX?

  • 基础几何变换:翻转、旋转、裁剪、缩放。
  • 色彩空间工具(例如:RGB↔HSV)。
  • 在加速器上高效运行的卷积风格滤镜。
  • 任何操作都可以与 JAX 的自动微分功能结合使用,因此例如,您可以优化图像处理流程作为大型学习系统的一部分。

所有示例都位于仓库的 examples/ 文件夹中,且刻意设计得很简单,用以说明函数如何与 JAX 的 jit, vmap, 和 pmap 组合使用。


Testing & contribution

  • 使用 pytest 运行测试套件(通过 pip install -e "[test]" 安装额外的测试依赖项)。
  • 欢迎任何形式的贡献;工作流程请参阅 CONTRIBUTING.md

When to use PIX?

如果您已经在使用 JAX 进行模型开发,并且需要快速、对 GPU 友善的图像预处理,且该预处理可以与您的模型代码一起编译和并行化,那么 PIX 提供了一套现成且经过良好测试的工具集。这能让您免于编写仅限 NumPy 的自定义代码,否则在迁移到加速器时会成为瓶颈。


TL;DR: PIX = JAX 原生图像处理,完全可 jit‑able 与并行化,在设置好 JAX 后,使用 pip install dm-pix 進行安裝。