2 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.
Physical Intelligence

About Physical Intelligence

201-500 employees
Contact me