Skip to content

perf(metax): vectorize small fused add RMSNorm rows - #998

Open
JoeZhang-0x000 wants to merge 1 commit into
InfiniTensor:masterfrom
JoeZhang-0x000:perf/metax-vectorized-fused-add-rms-norm
Open

JoeZhang-0x000 wants to merge 1 commit into
InfiniTensor:masterfrom
JoeZhang-0x000:perf/metax-vectorized-fused-add-rms-norm

Conversation

@JoeZhang-0x000

Copy link
Copy Markdown

Small MetaX fused Add+RMSNorm rows spend substantial time in the generic reduction path. Add an aligned, vectorized MetaX kernel for hidden=4096 and at most 128 rows, using one 256-thread block per row. Accumulate in FP32 after rounding the residual sum to the input dtype, matching the existing fused contract. Weight remains optional.

Keep larger prefill batches and unsupported/unaligned layouts on the existing kernel: enabling this small-batch kernel unconditionally regressed large-prefill model throughput. Add tests at 1, 128 and 129 rows and with aligned/unaligned padded row strides, covering FP32/FP16/BF16 and optional weight.

Validation on MetaX C550, MACA 3.8.0.23, PyTorch 2.10.0+metax3.8.0.7:

  • Built master 4c014ca with this change and the two separately proposed MetaX Embedding/RmsNorm changes; 301 focused operator tests passed, plus the smoke selection (2 tests).
  • All three patches also apply without fuzz to the plugin's locked InfiniOps 8c2f70a.
  • Prior BF16 Graph measurements on 8c2f70a: 1 row, hidden=4096, 40.38 -> 27.20 us; 64 rows, 96.13 -> 82.28 us. Both timings include two clones to reset mutable inputs.
  • In the existing Qwen3-8B TP=2 workload, limiting the new kernel to <=128 rows improved the 2048-in/512-out, batch=64 run from 1818.56 to 2112.39 output tokens/s (+16.16%); plugin code and benchmark settings were unchanged in that comparison.

NVIDIA hardware was not available for testing. Its shared implementation is unchanged.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant