发表机构
Massachusetts Institute of Technology(麻省理工学院)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
vidax是一个开源的JAX/Flax推理引擎,统一了张量和序列并行,支持多种视频生成模型,在TPU上实现高效推理,并开源作为基线。
AI 中文摘要
开源视频生成模型几乎完全以PyTorch/CUDA参考实现的形式发布。这使得云TPU Pod缺乏生产就绪的推理路径,尽管它们提供了大型、成本效益高的加速器内存池,非常适合长序列时空注意力。我们提出了vidax,一个开源的JAX/Flax推理引擎和零拷贝PyTorch到JAX权重转换器,用于现代视频生成架构。vidax涵盖了多种时空模型——包括扩散Transformer、全模态混合Transformer、3D VAE、文本编码器和原生采样器——在执行路径中零PyTorch依赖。该框架在单个JAX分片网格上统一了1D张量并行与DeepSpeed-Ulysses序列并行,集成了TPU闪存注意力内核,并实现了逐层权重卸载以支持超过单设备内存的参考分辨率。我们在TPU v4-8硬件上基准测试了编译时间、延迟和峰值内存利用率,并记录了检查点转换过程中出现的实际数值错误。vidax作为JAX和TPU视频生成研究的基线开源发布。
英文摘要
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 zero-copy PyTorch-to-JAX weight translator for modern video generation architectures. vidax covers a diverse set of spatiotemporal models --- including Diffusion Transformers, omnimodal Mixture-of-Transformers, 3D VAEs, text encoders, and native samplers --- with zero PyTorch dependency in the execution path. The framework unifies 1D tensor parallelism with DeepSpeed-Ulysses sequence parallelism on a single JAX sharding mesh, integrates TPU flash-attention kernels, and implements per-layer weight offloading to support reference resolutions that exceed single-device memory. We benchmark compile times, latency, and peak memory utilization on TPU v4-8 hardware, and document real-world numerical bugs surfaced during checkpoint translation. vidax is released open-source as a baseline for JAX and TPU video generation research.