About the Role
This role is embedded in production training, focusing on solving complex systems and performance challenges in large-scale model training. You will work directly with researchers, producing code, measurements, and system changes to enable better research and improve the performance, reliability, and numerical stability of training runs.
Responsibilities
- Improve the performance, reliability, and numerical stability of production training runs for large multimodal generative models
- Profile full training steps across model code, attention, kernels, data loading, encoders, communication, optimizer steps, checkpointing, and memory pressure
- Implement and validate GPU-level optimizations: fused kernels, attention paths, low-precision matmuls, quantization kernels, CUDA/Triton/CuTe/CUTLASS experiments, and no-compile alternatives
- Push lower-precision training forward, including FP8 / MXFP8 / FP4-style paths, weight and activation quantization, accumulation choices, convergence risk, and quality tradeoffs
- Work with researchers to translate architecture changes into efficient training implementations
- Debug distributed training failures: NaNs, loss spikes, silent numerical drift, memory leaks, stragglers, bad nodes, NCCL issues, and throughput cliffs
- Build benchmarking and profiling harnesses that make performance claims trustworthy
- Help the training team move quickly when an urgent bottleneck appears, while turning repeated failures into better abstractions and tools
Requirements
- Experience working deeply on large-scale training systems, ideally as part of a training group working closely with researchers
- Strong PyTorch fluency, including comfort reading and modifying low-level training code
- Experience with distributed training concepts such as FSDP, tensor/model/context/sequence parallelism, activation checkpointing, NCCL, and overlapping compute and communication
- Hands-on experience improving training throughput, memory footprint, or stability in real training runs
- Experience profiling GPU workloads with tools like Nsight Systems, Nsight Compute, torch profiler, trace viewers, or custom telemetry
- Practical GPU performance judgment
- Understanding of low-precision training and quantization tradeoffs: FP8, MXFP8, FP4/NVFP4-style formats, scaling, accumulation, numerical validation, and convergence risk
- Good research judgment: ability to partner with researchers on ablations, understand measurements, and tie optimization work to model-quality outcomes
- Comfortable operating in ambiguity
Skills
- PyTorch
- Distributed training
- FSDP
- Tensor parallelism
- Model parallelism
- Context parallelism
- Sequence parallelism
- Activation checkpointing
- NCCL
- GPU profiling
- Nsight Systems
- Nsight Compute
- Torch profiler
- Low-precision training
- Quantization
- FP8
- MXFP8
- FP4
- NVFP4
- CUDA
- Triton
- CuTe
- CUTLASS
Location
- Freiburg, Germany
- San Francisco
- Remote
Work Type
- Hybrid
- Remote
Experience Level
- Mid-level
- Senior
Salary/Compensations
- US $180,000 - $290,000 + equity
About the Company
- We're the team behind Latent Diffusion, Stable Diffusion, and FLUX—foundational technologies that changed how the world creates images and video.
- We’re creating the generative models that power how people make images and video—tools used by millions of creators, developers, and businesses worldwide.
- Our FLUX models are among the most advanced in the world, and we're just getting started.
- Headquartered in Freiburg, Germany with a growing presence in San Francisco, we’re scaling fast while staying true to what makes us different: research excellence, open science, and building technology that expands human creativity.
