ML Infra Engineer (TPU/Jax/Optimization)
Core
Design, implement, and maintain systems for large-scale model training, including scheduling, job management, checkpointing, and metrics/logging.
Role type
ML Infrastructure Engineer (TPU/JAX/Optimization)
Builds
Scalable training and inference infrastructure for large-scale model training
Domain
Machine Learning Infrastructure / Distributed Systems
Deliverable
infrastructure
Required skills
Software engineering fundamentals, Large-scale training experience, JAX, PyTorch, Distributed training, Cloud platform management, Performance debugging and optimization
Preferred skills
ML systems background, Hardware performance tuning, Robotics or multimodal model experience, System abstraction design
Technologies
JAX, PyTorch, TPU, GPU, Kubernetes, GCP, AWS, SLURM
Responsibilities
Design and maintain training/inference infrastructure systems, Scale distributed training across TPU and GPU clusters, Optimize memory usage and device utilization, Build abstractions for experiment management, Manage cloud compute resources and costs, Partner with researchers to translate needs into infra capabilities, Contribute to core JAX model and training code
Seniority
Mid-to-Senior, hands-on IC