About The Position

As an engineer on the ML Compute team, your work will include: Drive performance optimization for large-scale foundation model training on TPUs, focusing on efficiency, throughput, and scalability. Profile and optimize JAX/XLA workloads across compute, memory, communication, and compilation. Develop and optimize high-performance TPU kernels for critical ML operations such as attention and Mixture-of-Experts (MoE). Optimize distributed training techniques, sharding strategies, and collective communication over TPU interconnects (ICI/Fabric). Research and implement new techniques across the JAX, XLA, and TPU stack to improve end-to-end training performance. Develop performance profiling, benchmarking, and automated tuning capabilities for large-scale training workloads. Collaborate with cross-functional engineers to solve large-scale ML training challenges. Lead complex technical projects and mentor engineers in areas of your expertise. Cultivate a team centered on collaboration, technical excellence, and innovation.

Requirements

  • 6+ years of experience building or optimizing high-performance ML or distributed systems
  • Proficient in Python or other relevant programming languages
  • Strong understanding of distributed systems, parallel computing, and performance optimization
  • Experience profiling and optimizing compute-, memory-, or communication-intensive workloads
  • Ability to clearly communicate complex technical problems and collaborate with partners to develop solutions
  • Bachelor's degree in Computer Science, Engineering, or a related field

Nice To Haves

  • Advanced degree in Computer Science, Engineering, or a related field
  • Experience with accelerators such as TPU or GPU and understanding of accelerator architecture and performance characteristics
  • Experience with JAX, XLA, PyTorch or other ML compiler/runtime stacks
  • Experience developing or optimizing accelerator kernels using Pallas, Triton, CUDA, or similar technologies
  • Experience optimizing large-scale foundation model training and distributed communication

Responsibilities

  • Drive performance optimization for large-scale foundation model training on TPUs, focusing on efficiency, throughput, and scalability
  • Profile and optimize JAX/XLA workloads across compute, memory, communication, and compilation
  • Develop and optimize high-performance TPU kernels for critical ML operations such as attention and Mixture-of-Experts (MoE)
  • Optimize distributed training techniques, sharding strategies, and collective communication over TPU interconnects (ICI/Fabric)
  • Research and implement new techniques across the JAX, XLA, and TPU stack to improve end-to-end training performance
  • Develop performance profiling, benchmarking, and automated tuning capabilities for large-scale training workloads
  • Collaborate with cross-functional engineers to solve large-scale ML training challenges
  • Lead complex technical projects and mentor engineers in areas of your expertise
  • Cultivate a team centered on collaboration, technical excellence, and innovation
© 2026 Teal Labs, Inc
Privacy PolicyTerms of Service