« Tüm yayınlar

metaljax: JAX'ı Apple Silicon GPU'larında Çalıştıran Metal Backend

metaljax, JAX kodunu Apple Silicon GPU'larında MLX üzerinden çalıştıran açık kaynak bir PJRT eklentisi; JAX test paketinin %98,4'ünü geçiyor.

metaljax, JAX'ta hiçbir değişiklik yapmadan Apple Silicon GPU'larında çalışmayı sağlayan yeni bir PJRT eklentisi. Derlenmiş StableHLO programlarını MLX dizilerine yorumlayarak GPU üzerinde yürütüyor; jit, grad, vmap, lax.scan ve optax eğitim döngüleri sorunsuz çalışıyor, transformer eğitim adımları PyTorch'un MPS backend'ine yakın performans gösteriyor. Dikkat çekici nokta, projenin büyük ölçüde 'vibe coding' yani AI destekli yinelemeli geliştirme yöntemiyle yazılmış olması; buna rağmen JAX'ın resmi test paketinin (28.200 testten 27.779'u, %98,4) geçildiği ve her sürümde CPU backend'ine karşı doğruluk kontrolünden geçtiği belirtiliyor.

Mühendisler için önemi şu: JAX'ın şu an resmi bir Metal backend'i yok, bu yüzden Apple GPU kullanıcıları genelde PyTorch'un MPS desteğine bağımlı kalıyordu. Bilinen sınırlamalar hata değil, platform kısıtları: float64 desteklenmiyor (Metal GPU'larda f64 ALU yok, isteğe bağlı f32 emülasyonu var), pmap/shard_map tek cihazla sınırlı, denormal sayılar sıfıra yuvarlanıyor. Kalan eksikler (sıralı efekt token'ları, complex64'ün kutup noktalarındaki uç durumları, bazı PJRT yüzey API'leri) belgelenmiş ve CPU backend'ine göre denetlenmiş durumda.

Proje beta aşamasında; pip ile kurulabiliyor, JAX_PLATFORMS=metal ayarlanmadıkça varsayılan CPU backend'i değişmiyor, Apple Silicon, macOS 14+ ve Python 3.12+ gerektiriyor. Geliştirici, JAX ekosisteminde resmi bir Metal backend'i çıkarsa bu paketin kullanımdan kaldırılacağını belirtiyor.