⌈ Blog ⌋

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

View the code on GitHub

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.

A two-row diagram. On the top row, an observation feeds a stack of static input buffers that get updated, with robot actions shown as a robot arm on the right. An arrow leads down from the buffers into a box on the bottom row labelled HIP Graph Replay, which holds four stages (hipBLASLt GEMMs, AITER or custom attention, custom MFMA GEMMs, fused elementwise kernels), and an arrow leads back up from the box to the robot actions.
Fig. 01What FlashRT does under the hood to optimize inference for small-batch, latency-sensitive models.

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.

Two panels side by side. On CDNA3 (MI300X, gfx942) one thread out of a 64-thread wavefront supplies 4 BF16 values to a 16 × 16 × 16 MFMA instruction. On CDNA4 (MI350X, gfx950) one thread out of 64 supplies 8 BF16 values to a 16 × 16 × 32 MFMA instruction.
Fig. 02Each of the 64 threads in a wavefront provides 4 BF16 values per MFMA instruction on CDNA3 and 8 on CDNA4.

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 / CUDAAMD / ROCm
CUDA GraphsHIP Graphs
cuBLASLthipBLASLt
custom CUDAcustom HIP
MMAMFMA
Flash AttentionAITER / custom attention

Attempting to port FlashRT to MI300X showed where that backend was actually specific to gfx950.

PropertyCDNA3 / MI300XCDNA4 / MI350X
Compiler targetgfx942gfx950
FP8 storageE4M3 FNUZOCP E4M3
Maximum finite E4M3 value240448
hipBLASLt datatypeHIP_R_8F_E4M3_FNUZHIP_R_8F_E4M3
BF16 MFMA used by our packed path16×16×1616×16×32
BF16 values consumed per lane48
K depth per MFMA step1632
Fragment size per lane8 bytes16 bytes
MXFP4unavailableavailable

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.

The byte 0x4C drawn as eight bits: a sign bit of 0, four exponent bits 1001 equal to 9, and three mantissa bits 100 equal to 1.5. Decoded as OCP E4M3 on CDNA4 with bias 7 it is 2 to the power 2 times 1.5, which is 6.0. Decoded as E4M3 FNUZ on CDNA3 with bias 8 it is 2 to the power 1 times 1.5, which is 3.0.
Fig. 03The byte 0x4C decoded both ways. FNUZ uses an exponent bias of 8 and OCP uses 7, so the same bits give half the value on CDNA3.

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 projectionShape (M,N,K)BF16FP8 FNUZ
QKV(571,2560,2048)25.06 μs12.51 μs
Attention output(571,2048,2048)21.12 μs11.80 μs
Merged gate/up(571,32768,2048)130.38 μs100.37 μs
Down(571,2048,16384)87.54 μs57.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 projectionShape (M,N,K)hipBLASLt BF16hipBLASLt FP8
QKV(10,2560,1024)6.88 μs7.99 μs
Attention output(10,1024,2048)7.25 μs7.74 μs
Merged gate/up(10,8192,1024)7.58 μs7.62 μs
Down(10,1024,4096)8.02 μs11.23 μs
A bar chart of FP8 time as a multiple of BF16 time, with bars growing left or right from a parity line at 1.0. For the encoder at M = 571 all four bars point left: QKV 0.50, attention output 0.56, merged gate/up 0.77, down 0.66. For the decoder at M = 10 all four point right: QKV 1.16, attention output 1.07, merged gate/up 1.01, down 1.40.
Fig. 04FP8 time relative to BF16 for each projection. FP8 is faster on the 571-row encoder GEMMs and slower on the 10-row decoder GEMMs.

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.

python
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 values

The 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.

Two timelines on the same clock, each with a weight loads lane and an MFMA chain lane. In the first, one load is followed by one MFMA round, and the MFMA lane waits during every load, so three rounds fit. In the second, eight loads for round r + 1 run during the MFMA chain for round r, the MFMA lane waits only once at the start, and five rounds fit in the same time.
Fig. 05Loading one fragment at a time leaves the matrix cores waiting between rounds. Issuing the next round's eight loads during the current MFMA chain keeps them busy. Schematic, not to scale.

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_up

Both 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 projectionShape (M,N,K)hipBLASLt baselinePacked MFMAResult
QKV(10,2560,1024)6.88 μs3.93 μs1.75× faster
Attention output(10,1024,2048)7.25 μs5.68 μs1.28× faster
Merged gate/up(10,8192,1024)7.58 μs6.13 μs1.24× faster
Down(10,1024,4096)8.02 μs9.37 μskeep 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              → hipBLASLt

Fusing 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:

MetricOptimized CDNA3 BF16
Minimum34.97 ms
Median35.14 ms
p9535.43 ms
Maximum50.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.