
ML Infra Engineer, Modeling
Physical Intelligence5 days ago
Responsibilities
- Design, implement, and maintain large-scale model training and inference infrastructure, including scheduling, job management, checkpointing, and metrics and logging.
- Scale JAX-based distributed training across TPU and GPU clusters with researchers.
- Profile and optimize memory usage, device utilization, throughput, and distributed synchronization.
- Build abstractions for launching, monitoring, debugging, and reproducing experiments.
- Translate research needs into infrastructure capabilities and guide training-at-scale best practices.
- 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 large-scale training experience in JAX, with PyTorch experience also relevant.
- Familiarity with distributed training, multi-host setups, data loaders, and evaluation pipelines.
- Experience managing training workloads using SLURM, Kubernetes, GCP TPU/GKE, or AWS.
- Ability to debug and optimize performance bottlenecks across the training stack.
- Strong cross-functional communication and ownership mindset.
- Bonus qualifications include experience with training compilers, runtime optimization, custom kernels, GPU/TPU performance tuning, robotics, multimodal models, large-scale foundation models, or flexible and reliable researcher-facing abstractions.