google-deepmind/dm_pix

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

What is PIX?

PIXは、多くの現代的な機械学習研究を支える高性能なNumPy互換フレームワークであるJAXの上に構築された、軽量な画像処理ライブラリです。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の 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 でインストール。