feat(flux2klein): Onboard FLUX.2-Klein 4B and 9B with Flax NNX, fast weight loading, causal splash attention, and end-to-end optimizations
- Implemented Flax NNX Transformer architecture (NNXFlux2KleinTransformer2DModel) with support for both 4B (5 double / 20 single layers) and 9B (8 double / 24 single layers) configurations. - Integrated Flax Qwen3 text encoder with 3-layer intermediate hidden states extraction (layers 9, 18, 27), custom causal splash attention, and proper sharding constraints. - Implemented FlaxAutoencoderKL VAE decoder with fused batch normalization unscaling and channel re-layout. - Added fused end-to-end denoising loop scan with Flow Match Euler scheduler. - Added concurrent AOT XLA compilation across Qwen3, Flux transformer, and VAE. - Implemented fast host-memory streaming weight converter for safetensors shards directly into NNX State PyTree. - Optimized splash attention block sizes and Ulysses context parallelism sharding. - Added comprehensive unit tests (nnx_flux2klein_test.py) and end-to-end smoke test suite (generate_flux2klein_smoke_test.py).
A
Amey Pasarkar committed
3cb44556538edab9f667e74bf037eebcd58dc121
Parent: df69836