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