Real-time VLA inference on AMD hardware
Porting FlashRT to MI350X and MI300X GPUs.
- Published
- 2026-10-06T15:57Z
- Revised
- 2026-10-06
- Read
- 8 min / 1,635 words
- Topics
- VLA / Inference
- Authors
- Liang Su (Creator of FlashRT) & Cybernetic Physics
Inference speed is much more critical in robotics than in LLMs because a vision-language-action model (VLA) will always act on a world based on a frame taken before inference started. If inference takes 100 ms, the robot is always reacting to the world as it was 100 ms ago.
At 50 Hz, the policy makes a new decision every 20 milliseconds, so the VLA has to see the frame, think, and output an action in under 20 ms. FlashRT exists to make the entire inference call fit that budget. It’s an inference engine for small-batch, latency-sensitive models including VLAs like π0, π0.5 and GR00T.
Until this summer, FlashRT only ran on NVIDIA hardware. The past few weeks, we worked with the creator of FlashRT, Liang Su, to port FlashRT to AMD GPUs. FlashRT's AMD backend was originally introduced for the MI350X in PR #189, targeting AMD's CDNA4 architecture (gfx950). The next step was to do the same for the MI300X generation, which uses the previous CDNA3 architecture (gfx942) and is much more widely available for people to rent today. Our PR #214 extends that backend to MI300X and surprisingly the 16-bit path beat the 8-bit one, by a lot.
At first glance, the old chip (MI300X) and the new one (MI350X) look very similar. They both use ROCm, expose AMD Matrix Cores, support BF16 and FP8, and, as far as FlashRT is concerned, can run the same high-level model pipelines. However, that was no longer the case when we reached the optimized kernels FlashRT is built on top of. The two chips actually differ substantially (different FP8 format, different matrix instruction sizes), so the lower-level primitives couldn't simply be reused.
Firstly, CDNA3 and CDNA4 use different FP8 representations with CDNA3 using a variant called “FNUZ” and CDNA4 using the OCP layout. That means that the same byte decodes to half the value under FNUZ. Their BF16 MFMA instructions also have different shapes (16×16×16 on CDNA3, 16×16×32 on CDNA4) which changes how weights should be packed in memory. Finally, π0.5's action-decoder GEMMs are far smaller than the large GEMMs these GPUs are optimized around. In this post we go over how we addressed those differences.
How FlashRT works
The motivation behind FlashRT is low-latency inference for robotics models, which is crucial for running those models in production. By default, PyTorch runs each operation sequentially. For π0.5 that is about 2,500 times per call, and given that π0.5 is a fairly small model this adds unnecessary overhead. FlashRT, on the other hand, uses graph execution to pay that overhead once, when the model loads. By recording the sequence of operations once as a CUDA Graph (HIP Graph on AMD), it can then replay without the Python or memory allocation in the loop. It quantizes and packs the weights, it picks the fastest GEMM kernel for each layer and decides which attention code to use and then it runs the model once while recording every GPU operation in order. It then saves that recording as a static graph that it replays at every inference call.

The trade-off here is that all the choices (weight quantization and packing, GEMM kernel per layer, attention backend) have to be made up front and then frozen into the graph. To figure out whether we need a custom kernel or hipBLASLt for each layer, we profile on MI350X. Switching to another GPU generation does not guarantee that the same choices still give correctness and/or optimal performance.
The AMD stack
The AMD backend introduced in Liang's PR uses two libraries, hipBLASLt and AITER, plus custom kernels that use MFMA instructions. hipBLASLt is AMD's library for matrix multiplication (GEMM) on its GPUs, with kernels for both BF16 and FP8. It's the AMD equivalent of NVIDIA's cuBLASLt.
AITER provides fast attention kernels for AMD. Profiling one full inference step shows that attention takes up a big chunk of the total time: on the MI350X for π0.5, replacing AITER with the PyTorch SDPA fallback increased median latency from roughly 16.3 ms to 22.2 ms.
At the lowest level, our custom matrix kernels use MFMA, AMD's Matrix Fused Multiply-Add instruction family. For CUDA programmers, MFMA is closely related to what you'd encounter through NVIDIA MMA instructions. On NVIDIA GPUs, the special hardware just for multiplying matrices is called Tensor Cores. CUDA exposes that hardware to programmers through warp- and warpgroup-level operations such as MMA and WGMMA. AMD similarly has Matrix Cores and MFMA is to AMD's Matrix Cores what MMA is to NVIDIA's Tensor Cores.
An MFMA instruction such as __builtin_amdgcn_mfma_f32_16x16x16bf16_1k(...) is executed by all 64 threads (lanes) of a wavefront, and the hardware decides which part of the two input matrices each thread has to provide. On CDNA4 that part is 8 BF16 values per thread while on CDNA3 it is 4. We therefore have to pack the model weights in memory so that each lane's slice is contiguous and by extension we can’t reuse the CDNA4 kernel for CDNA3.

From CDNA4 to CDNA3
With PR #189, Liang introduced FlashRT's first AMD backend, targeting MI350X / CDNA4. That port mostly swapped each NVIDIA component for its AMD counterpart.
| NVIDIA / CUDA | AMD / ROCm |
|---|---|
| CUDA Graphs | HIP Graphs |
| cuBLASLt | hipBLASLt |
| custom CUDA | custom HIP |
| MMA | MFMA |
| Flash Attention | AITER / custom attention |
Attempting to port FlashRT to MI300X showed where that backend was actually specific to gfx950.
| Property | CDNA3 / MI300X | CDNA4 / MI350X |
|---|---|---|
| Compiler target | gfx942 | gfx950 |
| FP8 storage | E4M3 FNUZ | OCP E4M3 |
| Maximum finite E4M3 value | 240 | 448 |
| hipBLASLt datatype | HIP_R_8F_E4M3_FNUZ | HIP_R_8F_E4M3 |
| BF16 MFMA used by our packed path | 16×16×16 | 16×16×32 |
| BF16 values consumed per lane | 4 | 8 |
| K depth per MFMA step | 16 | 32 |
| Fragment size per lane | 8 bytes | 16 bytes |
| MXFP4 | unavailable | available |
Different FP8 representation
CDNA3 and CDNA4 use different FP8 formats: CDNA3 uses E4M3 FNUZ, while CDNA4 uses OCP E4M3. Because the formats interpret the same bit patterns differently, FP8 weights cannot be reused directly between the two architectures. For the CDNA3 port, we updated the FP8 quantization path and the relevant kernels so weights are produced and consumed consistently in FNUZ format.
FP8 helps the encoder and hurts the decoder
π0.5 has two very different compute patterns. The encoder processes hundreds of rows at once, while the action decoder operates on only 10 rows and repeats this small computation throughout the denoising process. This makes FP8 much more effective in the encoder than in the decoder.
Encoder
For the large encoder projections, FP8 consistently reduced latency.
| Encoder projection | Shape (M,N,K) | BF16 | FP8 FNUZ |
|---|---|---|---|
| QKV | (571,2560,2048) | 25.06 μs | 12.51 μs |
| Attention output | (571,2048,2048) | 21.12 μs | 11.80 μs |
| Merged gate/up | (571,32768,2048) | 130.38 μs | 100.37 μs |
| Down | (571,2048,16384) | 87.54 μs | 57.48 μs |
With 571 rows, there is enough parallel work for the GPU to use its compute resources efficiently. FP8 also reduces the amount of data that needs to be read for the weights, so these larger projections see a clear latency improvement.
Decoder
FP8 was slower than BF16 on the 10-row decoder GEMMs:
| Decoder projection | Shape (M,N,K) | hipBLASLt BF16 | hipBLASLt FP8 |
|---|---|---|---|
| QKV | (10,2560,1024) | 6.88 μs | 7.99 μs |
| Attention output | (10,1024,2048) | 7.25 μs | 7.74 μs |
| Merged gate/up | (10,8192,1024) | 7.58 μs | 7.62 μs |
| Down | (10,1024,4096) | 8.02 μs | 11.23 μs |
The decoder only processes 10 rows at a time, so there is not enough work to make up for the extra overhead of the FP8 path. Since these GEMMs are repeated across all 18 layers and 10 denoising steps, even a few extra microseconds per GEMM add up quickly.
Writing a CDNA3 BF16 MFMA kernel
CDNA4 uses a different BF16 MFMA layout, so for CDNA3 we wrote a new kernel around V_MFMA_F32_16X16X16BF16_1K.
Packing around the instruction
Since the weights do not change during inference, we pack them once at model load time into the layout expected by the CDNA3 MFMA instruction.
packed = ( W.view(K // 16, 4, 4, N // 16, 16) .permute(3, 0, 1, 4, 2) .contiguous())output tile ↓MFMA K step ↓lane ↓four contiguous BF16 valuesThe packed weights are created once when the model is loaded and reused for every inference. Their addresses stay fixed, so the HIP Graph can reference them directly on every replay.
Keeping enough loads in flight
CDNA3’s smaller BF16 MFMA fragment (i.e. four BF16 values per lane per instruction) means each instruction consumes less weight data, so loading one fragment at a time can leave the compute units waiting on memory. To avoid that, we issue the next round of loads while the current MFMA work is still running.
We issue eight independent weight loads per lane for each round. While the current MFMA operations are running, the weights for the next round are already being loaded, which avoids stalls between rounds.
Merging gate and up
A standard feed-forward block computes:
gate = A × W_gateup = A × W_upBoth multiplications consume exactly the same activation matrix. We concatenate their static weights, W_merged = [W_gate, W_up], and compute [gate, up] = A × W_merged. For π0.5 this gives the decoder GEMM (10, 8192, 1024). The combined matrix is packed once and one MFMA call replaces the two projection paths. The kernel also performs the bias operation in its epilogue.
Beating the hipBLASLt baseline
We benchmarked the packed BF16 implementation against the original hipBLASLt BF16 route on every decoder projection.
| Decoder projection | Shape (M,N,K) | hipBLASLt baseline | Packed MFMA | Result |
|---|---|---|---|---|
| QKV | (10,2560,1024) | 6.88 μs | 3.93 μs | 1.75× faster |
| Attention output | (10,1024,2048) | 7.25 μs | 5.68 μs | 1.28× faster |
| Merged gate/up | (10,8192,1024) | 7.58 μs | 6.13 μs | 1.24× faster |
| Down | (10,1024,4096) | 8.02 μs | 9.37 μs | keep hipBLASLt |
The first three shapes benefit from the packed MFMA kernel, but the down projection does not. With only 1024 output columns, it launches just 64 workgroups, while each workgroup still has to process a relatively long K dimension of 4096. hipBLASLt performs better for this shape, so we keep it for the down projection and use the custom kernel for the others.
QKV → packed CDNA3 MFMAAttention → optimized attentionAttention output → packed CDNA3 MFMAGate + Up → packed CDNA3 MFMADown → hipBLASLtFusing the BF16 decoder boundary
After the matrix multiplications were faster, profiling exposed another repeated overhead between decoder blocks. The graph originally performed residual = BF16(residual + output × gate) followed by normalized, next_gate = adaptive_rms_norm(residual, style). These are both small operations. But they appear repeatedly through the 18 decoder layers and 10 denoising steps. We combined them into a single gate_residual_ada_norm_bf16 kernel. This eliminates roughly 350 kernel launches per inference.
Results
With the packed MFMA routing plus the fused BF16 decoder boundary, the final π0.5 implementation reaches:
| Metric | Optimized CDNA3 BF16 |
|---|---|
| Minimum | 34.97 ms |
| Median | 35.14 ms |
| p95 | 35.43 ms |
| Maximum | 50.15 ms |
The 35.14 ms median end-to-end latency is above the 20 ms budget from the introduction. The MI350X backend, at roughly 16.3 ms, fits it. The benchmark uses real LIBERO Spatial camera observations, pinned denoising noise, GPU-local CPU affinity, 50 untimed warmup calls, and 100 measured full inference calls. The same CDNA3 port also brought GR00T N1.7 onto MI300X. Its backbone plus four-step action chain measured 19.58 ms median, with 0.999934 cosine similarity against the official policy across 680 action values.
Acknowledgements
The original CDNA4 FlashRT backend and MI350X optimization work was implemented by Liang Su in PR #189. The CDNA3 work described here is implemented in PR #214, extending that backend to MI300X while preserving the existing CDNA4 path.