
ML Infra Engineer
Physical Intelligence2 years ago
Responsibilities
- Design, implement, and maintain infrastructure for large-scale model training and inference, including scheduling, job management, checkpointing, and metrics/logging.
- Scale distributed JAX training across TPU and GPU clusters and optimize memory usage, device utilization, throughput, and synchronization.
- Build abstractions for launching, monitoring, debugging, and reproducing experiments.
- Manage cloud-based GPU/TPU compute allocation and utilization while controlling costs.
- Translate researcher needs into infrastructure capabilities and guide best practices for large-scale training.
- Evolve JAX model and training code to support new architectures, modalities, and evaluation metrics.
Requirements
- Strong software engineering fundamentals and experience building ML training infrastructure or internal platforms.
- Hands-on experience with large-scale training in JAX and PyTorch.
- Familiarity with distributed training, multi-host setups, data loaders, and evaluation pipelines.
- Experience managing training workloads on cloud platforms and systems such as SLURM, Kubernetes, GCP TPU/GKE, or AWS.
- Ability to debug and optimize performance bottlenecks across the training stack.
- Strong cross-functional communication and ownership skills.