
About the Role
AssemblyAI is looking for a Senior Research Engineer to join the Research team, developing and improving the systems behind large-scale distributed training, data processing, and inference. The goal is to solve customer problems and improve products quickly through model development and measurement — raising experimental velocity and lowering the friction of running, measuring, and debugging experiments.
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 working closely with researchers, the infrastructure team, and production engineering to follow problems through to resolution.
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, stay skeptical of your own results, and treat unexplained improvements as problems rather than wins.
- Appetite for the whole pipeline. Experience expanding outward from JAX and TPU performance into data problems or evaluation blind spots.
- 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.
- Strong Python skills; C++ or Rust experience for kernel-level work is a plus.
- Excellent communication and a collaborative mindset with the ability to clearly explain complex tradeoffs.
Bonus
- Domain knowledge in Speech-to-Text: ASR architectures, audio processing, streaming inference.
Timezone overlap
UTC-8–-4
Open to
US · New York · United States
Sign in to track applications and earn points.