ML Infra Engineer, Modeling

Physical Intelligence
San Francisco, CA
On-site

Who this role is best for

A natural match if you have experience with large-scale ML training infrastructure and JAX.

Best fit for

  • Candidates with hands-on experience in JAX and distributed training systems
    — “Hands-on large-scale training experience in JAX (preferred), PyTorch.
  • Individuals who can balance researcher flexibility with system reliability
    — “Experience designing abstractions that balance researcher flexibility with system reliability.
  • Candidates with a background in robotics or foundation models
    — “Background in robotics, multimodal models, or large-scale foundation models.

Things to consider

  • The role requires handling TPU and GPU clusters in a high-leverage, cross-functional setting
    — “This is a hands-on, high-leverage role at the intersection of ML, software engineering, and scalable infrastructure.
  • Strong ownership mindset and communication skills are non-negotiable
    — “Strong cross-functional communication and ownership mindset.

How to stand out

  • Highlight experience with JAX training pipelines and distributed systems optimization
    — “building reusable and efficient JAX training pipelines.
  • Showcase your ability to debug and optimize training stack performance bottlenecks
    — “Ability to debug and optimize performance bottlenecks across the training stack.
  • Emphasize your track record of translating research needs into production-ready infrastructure
    — “Translate research needs into infra capabilities and guide best practices for training at scale.
Pace · SteadyCollaboration · HighAutonomy · MediumDecision Impact · Team

Derived from job-description analysis by Serendipath's career intelligence engine.

What success looks like

  • design and implementation of training systems
  • optimization of performance
  • abstractions for experiments
  • contribution to core training code
Typical background
strong software engineering fundamentalshands-on large-scale training experience

Skills & requirements

Required

Large-scale Training InfrastructureJAX Training PipelinesDistributed TrainingGpu/tpu Performance TuningTraining CompilersRuntime OptimizationCustom KernelsCloud Platforms

Preferred

RoboticsMultimodal ModelsLarge-scale Foundation Models

Stack & domain

JAXPyTorchDistributed TrainingTpu/gpu ClustersSlurmKubernetesGCP Tpu/gkeAWSGpu/tpu Performance TuningCommunicationProblem-solvingTeamworkML InfrastructureTraining SystemsLarge-scale Training

About the role

Original posting from Physical Intelligence via Ashby

Physical Intelligence is bringing general-purpose AI into the physical world. We are a group of engineers, scientists, roboticists, and company builders developing foundation models and learning algorithms to power the robots of today and the physically-actuated devices of the future.

In this role you will help scale and optimize our training systems and core model code. You’ll own critical infrastructure for large-scale training, from managing GPU/TPU compute and job orchestration to building reusable and efficient JAX training pipelines. You’ll work closely with researchers and model engineers to translate ideas into experiments—and those experiments into production training runs.

This is a hands-on, high-leverage role at the intersection of ML, software engineering, and scalable infrastructure.

The Team

The ML Infrastructure team supports and accelerates PI’s core modeling efforts by building the systems that make large-scale training reliable, reproducible, and fast. The team works closely with research, data, and platform engineers to ensure models can scale from prototype to production-grade training runs.

In This Role You Will

  • Own training/inference infrastructure: Design, implement, and maintain systems for large-scale model training, including scheduling, job management, checkpointing, and metrics/logging.
  • Scale distributed training: Work with researchers to scale JAX-based training across TPU and GPU clusters with minimal friction.
  • Optimize performance: Profile and improve memory usage, device utilization, throughput, and distributed synchronization.
  • Enable rapid iteration: Build abstractions for launching, monitoring, debugging, and reproducing experiments.
  • Partner with researchers: Translate research needs into infra capabilities and guide best practices for training at scale.
  • Contribute to core training code: Evolve JAX model and training code to support new architectures, modalities, and evaluation metrics.

What We Hope You’ll Bring

  • Strong software engineering fundamentals and experience building ML training infrastructure or internal platforms.
  • Hands-on large-scale training experience in JAX (preferred), PyTorch.
  • Familiarity with distributed training, multi-host setups, data loaders, and evaluation pipelines.
  • Experience managing training workloads on cloud platforms (e.g., SLURM, Kubernetes, GCP TPU/GKE, AWS).
  • Ability to debug and optimize performance bottlenecks across the training stack.
  • Strong cross-functional communication and ownership mindset.

Bonus Points If You Have

  • Deep ML systems background (e.g., training compilers, runtime optimization, custom kernels).
  • Experience operating close to hardware (GPU/TPU performance tuning).
  • Background in robotics, multimodal models, or large-scale foundation models.
  • Experience designing abstractions that balance researcher flexibility with system reliability.

Pursuant to the San Francisco Fair Chance Ordinance, we will consider for employment qualified applicants with arrest and conviction records.

Source: Physical Intelligence careers (Ashby)

Similar roles