Member of Technical Staff - Inference-TPU

RadixArk

• $200K — $400K *
Information Technology
Less than 5 years of experience
Job Overview by Ladders

Qualifications

  • 3+ years experience building production ML systems with JAX, XLA, or TPU frameworks
  • Bachelor's or Master's degree in Computer Science, Electrical Engineering, or similar
  • Deep understanding of JAX/XLA internals like HLO and partitioning
  • Strong performance tuning instincts at compiler and runtime layers
  • Experience with distributed inference or training frameworks
  • Proficiency in Python for writing high-performance production code
  • Familiarity with Pallas or ability to learn quickly

Responsibilities

  • Build high-performance inference and training systems using JAX/XLA/Pallas
  • Push large-model workloads to the limits on TPU v4, v5e, and v5p
  • Optimize end-to-end latency for LLM serving on TPU infrastructure
  • Design efficient SPMD strategies for distributed inference and training
  • Profile and optimize XLA compilation pipelines and transformations
  • Collaborate with kernel engineers and compiler teams for performance optimization
  • Contribute to open-source TPU optimization projects
  • Create testing frameworks for numerical correctness and performance detection

Benefits

  • Significant founding team equity
  • Comprehensive health benefits
  • Flexible work arrangements
Full Job Description
About the Role

RadixArk is looking for a Member of Technical Staff - TPU Systems to build high-performance inference and training systems using JAX, XLA, and Pallas. You'll push model workloads to their limits on TPU hardware, working on SGLang-JAX and other critical infrastructure that enables efficient deployment of frontier models on Google's tensor processing units.
Requirements
  • 3+ years experience building production ML systems utilizing JAX/Torch, XLA, or TPU-focused frameworks.
  • Bachelor's or Master's degree in Computer Science, Electrical Engineering, or equivalent industry experience
  • Deep understanding of XLA internals preferred: HLO, MLIR, operator fusion, SPMD partitioning, and sharding strategies.
  • Strong performance tuning instincts across compiler and runtime layers
  • Experience with distributed inference systems (e.g. SGLang, vLLM) or training frameworks (e.g. Miles, Alpa, Pathways)
  • Proficiency in Python with demonstrated ability to write high-performance, production-quality code
  • Experience writing custom GPU/TPU/AI Accelerator kernels. Familiarity with Pallas for kernel development is strongly preferred.
Responsibilities
  • Build high-performance inference and training systems using JAX/XLA/Pallas, including SGLang-JAX
  • Push large-model workloads to the limits on the newest TPU hardwares
  • Optimize end-to-end latency and throughput for LLM serving on TPU infrastructure
  • Design and implement SPMD strategies for efficient distributed inference and training
  • Design and implement Pallas kernels for operations that require customized low level control for best performance
  • Profile and optimize XLA compilation pipelines and HLO graph transformations
  • Collaborate with kernel engineers and compiler teams to achieve performance wins across the stack
  • Contribute to open-source projects with TPU optimization guides, benchmarks, and architectural insights


Compensation

Depending on background, skills, and experience, the expected annual salary range for this position is $200,000 - $400,000 USD + equity.

Similar Jobs

More Jobs at RadixArk

More Information Technology Jobs

Find similar Member of Technical Staff - Inference-TPU jobs: