Skip to content

[Common] Requantize grouped MXFP8 tensors with 128x128 tiles - #3656

Merged
phu0ngng merged 3 commits into
NVIDIA:mainfrom
sraman-rgb:sraman/mxfp8-group-requantize-wide-tiles
Oct 9, 2026
Merged

phu0ngng merged 3 commits into
NVIDIA:mainfrom
sraman-rgb:sraman/mxfp8-group-requantize-wide-tiles

Conversation

@sraman-rgb

Copy link
Copy Markdown
Collaborator

nvte_group_requantize turns a row-wise MXFP8 grouped tensor, such as tokens dispatched in FP8, into its column-wise copy with GEMM-swizzled scaling factors. Its kernel runs one 32x128 tile per 128-thread CTA, binary-searches the group offsets in thread 0 before the tile can finish, and writes both scale outputs one byte at a time; on a capacity-sized buffer every tile past the live rows still costs a CTA.

Add a kernel for the fast-math case without a dequantized output, the case the grouped MLP uses. A CTA owns a 128x128 tile and keeps it in registers, loading and storing coalesced 128-byte row segments; each swizzled scale output is one contiguous 512-byte block per tile; and the CTAs of a persistent grid walk the tiles in row-block order, so they stop at the first tile past the live rows. The arithmetic is that of the existing kernel and the outputs are bit-identical. Other cases keep the existing kernel.

Description

Please include a brief summary of the changes, relevant motivation and context.

Fixes # (issue)

Type of change

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

Please list the changes introduced in this PR:

  • Change A
  • Change B

Checklist:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes

nvte_group_requantize turns a row-wise MXFP8 grouped tensor, such as
tokens dispatched in FP8, into its column-wise copy with GEMM-swizzled
scaling factors. Its kernel runs one 32x128 tile per 128-thread CTA,
binary-searches the group offsets in thread 0 before the tile can
finish, and writes both scale outputs one byte at a time; on a
capacity-sized buffer every tile past the live rows still costs a CTA.

Add a kernel for the fast-math case without a dequantized output, the
case the grouped MLP uses. A CTA owns a 128x128 tile and keeps it in
registers, loading and storing coalesced 128-byte row segments; each
swizzled scale output is one contiguous 512-byte block per tile; and
the CTAs of a persistent grid walk the tiles in row-block order, so
they stop at the first tile past the live rows. The arithmetic is that
of the existing kernel and the outputs are bit-identical. Other cases
keep the existing kernel.

Signed-off-by: Siddhartha Raman S <sraman@nvidia.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@greptile-apps

greptile-apps Bot commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 5/5

[Medium impact] The PR appears safe to merge; no actionable new issue was found.

Summary

Adds a persistent 128×128-tile kernel for grouped MXFP8 requantization when fast math is enabled and no dequantized output is requested.

  • Moves FP8 data in packed words and writes scales in 512-byte blocks.
  • Stops work at the live-row boundary instead of visiting the full capacity buffer.
  • Keeps the existing kernel for other requests.
  • No actionable new issues were found. Existing C++ tests cover both FP8 input types, empty groups, capacity tails, and the new dispatch combination; tests were not run.

Diagram

%%{init: {'theme': 'neutral'}}%%
flowchart TD
  A["Grouped MXFP8 input"] --> B{"Fast math, no dequantized output,<br/>supported group count and alignment?"}
  B -->|Yes| C["Persistent 128×128 tiles"]
  B -->|No| D["Existing 32×128 tiles"]
  C --> E["Columnwise FP8 data and both swizzled scale outputs"]
  D --> E
Loading

Reviews (3) · Last reviewed commit: "Merge branch 'main' into sraman/mxfp8-gr..." · Reviewed by Greptile

@phu0ngng

phu0ngng commented Oct 8, 2026

Copy link
Copy Markdown
Collaborator

/te-ci L1 Pytorch

@phu0ngng

phu0ngng commented Oct 9, 2026

Copy link
Copy Markdown
Collaborator

/te-ci

@phu0ngng
phu0ngng merged commit 15dcc84 into NVIDIA:main Oct 9, 2026
11 of 12 checks passed
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.

2 participants