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