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 a 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 的 jitvmappmap 組合使用。


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 進行安裝。