Senior Research Engineer
About the Role
We're looking for a Senior Research Engineer to join our Research team, developing and improving the systems behind large-scale distributed training, data processing, and inference. Our goal as an organization is to solve customer problems and improve our products quickly through model development and measurement — and how fast we move depends on how quickly anyone here can run an experiment, measure it, and find out what's wrong. Raising that ceiling is the heart of this role. You'll be working inside the pipeline you're improving, not alongside it.
The ideal candidate has a deep understanding of modern deep learning systems, combined with strong engineering expertise across JAX and TPUs, layer-level optimization, large-scale distributed training, streaming, low-latency and asynchronous inference, inference compilers, and advanced parallelization techniques.
This is a cross-functional role. You'll work closely with our researchers, our infrastructure team, and production engineering — not as a handoff point, but as the person who learns enough of each domain to follow problems through to resolution. At times you'll train models, run evaluations, and analyze data yourself, both to deliver impact directly and to learn what's worth building to multiply the team's work. The bar is someone who understands the end-to-end impact they intend to make, measures it from the outset, and would rather find out they were wrong in a week than in a quarter. That discipline is what turns cross-functional ownership into an advantage.
The role is embedded within the Research team.
What You’ll Do
Raise the team's experimental velocity — make it faster to launch an experiment job, get a number back you can trust, and know what to try next.
Maintain and evolve our JAX training framework, keeping it scalable and efficient for large-scale distributed training runs on TPU.
Improve the data our models learn from: investigating quality issues, building the tooling to surface them, and turning what you find into measurable accuracy gains.
Analyze the accuracy of production models, build evaluation harnesses, and work out which improvements will matter most to customers.
Translate research prototypes into production-ready systems, refactoring and modernizing model architectures and infrastructure along the way.
Optimize production inference for speech language models, both from a serving architecture perspective and through advanced techniques such as quantization and speculative decoding.
Investigate and resolve performance bottlenecks across the stack, from low-level kernels (XLA, Pallas) to high-level system design.
Partner with researchers, infrastructure, and production engineering to trace problems to their real source and ship fixes that hold.
What You’ll Need
Expert-level proficiency with JAX and TPUs, including the surrounding ecosystem (Flax, Optax, the XLA compilation pipeline).
Measurement discipline. You define what success looks like before you start, you stay skeptical of your own results until they hold up, and you treat an unexplained improvement as a problem rather than a win.
Appetite for the whole pipeline. Your core strength might be JAX and TPU performance, but when a customer issue traces back to a data problem or an evaluation blind spot, you want to go find it yourself. The people who do well here went deep in one area first, then kept expanding outward.
Strong experience optimizing inference systems for production, ideally with LLMs or speech models.
Deep understanding of distributed training at scale, modern deep learning systems, and ML infrastructure best practices.
Familiarity with modern inference optimization techniques: continuous batching, KV-cache management, sharding strategies, quantization.
Enthusiasm for refactoring and improving existing systems — you thrive on making products and code faster and better.
Strong Python skills; C++ or Rust experience for kernel-level work is a plus.
Excellent communication and a collaborative mindset — you can clearly explain complex tradeoffs and prioritize high-impact work.
Bonus
Domain knowledge in Speech-to-Text: ASR architectures, audio processing, streaming inference.
Pay Transparency:
AssemblyAI strives to recruit and retain exceptional talent from diverse backgrounds while ensuring pay equity for our team. Our salary ranges are based on paying competitively for our size, stage, and industry, and are one part of many compensation, benefit, and other reward opportunities we provide.
There are many factors that go into salary determinations, including relevant experience, skill level, qualifications assessed during the interview process, and maintaining internal equity with peers on the team. The range shared below is a general expectation for the function as posted, but we are also open to considering candidates who may be more or less experienced than outlined in the job description. In this case, we will communicate any updates in the expected salary range.
The provided range is the expected salary for candidates in the U.S. Outside of those regions, there may be a change in the range which will be communicated to candidates throughout the interview process.
Salary range: $270,000 - $310,000