arXiv:2609.18077v1 Announce Type: cross Abstract: Open-source video generative models ship almost exclusively as PyTorch/CUDA reference implementations. This leaves Cloud TPU pods without a production-ready inference path, despite offering large, cost-effective accelerator memory pools ideal for long-sequence spatiotemporal attention. We present vidax, an open-source JAX/Flax inference engine and
vidax is an open-source JAX/Flax inference engine and weight translator designed to run modern video generative models on Cloud TPUs with zero PyTorch dependency. Submitted on September 16, 2026, it bridges the gap for TPU users by supporting five major model families: Wan, Cosmos, LTX, HunyuanVideo, and CogVideoX.
The framework enables reference-resolution video generation (up to 720p) on TPU v4-8 hardware by unifying Megatron-style tensor parallelism with DeepSpeed-Ulysses sequence parallelism. It utilizes Pallas flash-attention kernels and per-layer weight offloading to manage memory constraints, allowing models like the HunyuanVideo 13B to run within the ~30.75 GB HBM budget per chip.
Key capabilities include: Exact Weight Translation: A zero-copy translator converts PyTorch checkpoints to JAX with high numerical parity. Composable Parallelism: A 3-axis sharding mesh balances parameter residency and activation memory. Benchmark Performance: Per-step latencies range from 2.4 seconds (small models) to 299.6 seconds (largest models requiring offloading). Bug Taxonomy: The release documents specific numerical edge cases and precision issues encountered during cross-framework porting.
This material presents vidax, an open-source JAX/Flax inference framework designed to run video generative models on accelerator meshes, with a particular focus on Cloud TPU pods. The central motivation is an ecosystem gap: most open-source video generation models are released primarily as PyTorch/CUDA reference implementations, which leaves TPU-based deployments without a mature, production-ready inference path. This matters because TPU pods can provide very large, cost-effective accelerator memory pools, which are especially useful for video generation workloads where long-sequence spatiotemporal attention can be memory- and compute-intensive.
The key contribution is a unified JAX-based stack that makes it practical to port and serve video generative models on accelerator meshes rather than treating TPU inference as a secondary or experimental target. By leveraging JAX/Flax, vidax is positioned to take advantage of XLA-style compilation, composable parallelism, and mesh-aware execution, which are important for scaling attention and other video-model operations across many accelerators. The paper therefore addresses not just model compatibility, but the broader systems question of how to make open video generation models deployable on non-CUDA accelerator fabrics in a production-grade way.
This work matters because it lowers the barrier to running large video generative models on TPUs, where large pooled memory can be a practical advantage for long videos or high-resolution spatiotemporal modeling. More broadly, it pushes the video generation ecosystem beyond the default PyTorch/CUDA assumption, improving accessibility for researchers and practitioners who have access to TPU pods or other accelerator meshes. If vidax provides a reusable inference engine rather than a one-off model port, its impact extends to the broader open-source video model landscape by enabling a more portable, accelerator-agnostic deployment path.