blockwise fp8 gemm integration for gfx942 and gfx950#658
Conversation
…dant HIP guards, revert unnecessary common.h change
# Conflicts: # tests/cpp/operator/CMakeLists.txt
…hipkittens_optim # Conflicts: # transformer_engine/common/CMakeLists.txt # transformer_engine/common/gemm/kittens/CMakeLists.txt # transformer_engine/common/gemm/rocm_gemm.cu # transformer_engine/pytorch/quantization.py
…gemm_hipkittens_optim # Conflicts: # transformer_engine/common/gemm/kittens/cdna4/mxfp8_gemm.cpp # transformer_engine/pytorch/quantization.py
Claude review — PR 658Reviewed the ROCm-only 3-dot diff ( Verdict: solid direction, but there are a few things worth fixing before merge. Findings posted inline:
Copyright headers: OK — every touched file carries an up-to-date AMD 2026 header in the correct format; NVIDIA lines preserved on modified upstream files. |
|
|
Is CDNA3 support expected to be ported to HK? |
Currently HK supports CDNA3, CDNA4, and UDNA1. There is currently effort with the external HK team to move away from having separate git branches for each arch to having directory level separation. |
In this case can we avoid having 2 copies of HK submodule? |
Yes, looks like they merged it in late last week. I think that will require some integration/cmake changes so that the right headers are imported for each kernel. We'll probably want to move our kittens kernels into cdna3/cdna4/udna1 folders as well, so might be better as a follow up ticket? |
| emulated = get_device_compute_capability() >= (10, 0) | ||
| return supported and not emulated | ||
|
|
||
| def rocm_blockwise_unsupported_reason( |
There was a problem hiding this comment.
Here generally I think we should return a tuple -- True/False, followed by an optional reason. Maybe rename to rocm_blockwise_is_supported too.
| ): | ||
| if not fp8_blockwise_gemm_supported(): | ||
| pytest.skip("CUDA version does not support blockwise FP8 gemm.") | ||
| if IS_HIP_EXTENSION: |
There was a problem hiding this comment.
We could also potentially use the rocm function above here too?
There was a problem hiding this comment.
yes, now we are using the function
| is_w_1d_scaled, | ||
| ) -> None: | ||
| # NOTE: BGRAD epilogue is not supported for fp8. | ||
| if IS_HIP_EXTENSION: |
There was a problem hiding this comment.
One of the great things about hipKittens is how easy it is to fuse epilogues into kernels. If you have time it might be worth looking into bgrad/accumulator support.
There was a problem hiding this comment.
I will leave this as a future support. The current implementation supports the minimal cases that the test is expecting
| ) -> None: | ||
| # e5m2 by e5m2 not supported. | ||
| if IS_HIP_EXTENSION: | ||
| expected_err_msg = "does not support e5m2 by e5m2 inputs" |
There was a problem hiding this comment.
e5m2 should be supported by hipkittens, have a look at the mxfp8 kernels to see how we template for that
There was a problem hiding this comment.
Same here. The test expects e4m3 x e4m3 / e4m3 x e5m2 / e5m2 x e4m3 but not e5m2 x e5m2 and the current implementation also supports the former 3 but not the e5m2 x e5m2. While this is an easy support, I will leave it as a future PR
| use_split_accumulator, | ||
| is_x_1d_scaled, | ||
| is_w_1d_scaled, | ||
| expected_err_msg="dimension requirement", |
There was a problem hiding this comment.
This change should be hip guarded
There was a problem hiding this comment.
fixed. To keep the upstream lines untouched, I guarded early-return block, which makes the code a bit longer. Let me know if you have a cleaner approach.
|
|
||
| #ifndef KITTENS_DTYPE_ENUM_DEFINED | ||
| #define KITTENS_DTYPE_ENUM_DEFINED | ||
| enum KittensDType { |
There was a problem hiding this comment.
I'd rather not have to redefine these. I think we can have a parent folder header file that contains these enums and any other shared values.
There was a problem hiding this comment.
fixed. Added kittens_gemm_enums.h
| #endif // KITTENS_SCALING_MODE_DEFINED | ||
|
|
||
| namespace blockwise_gfx942 { | ||
| void kittens_blockwise_fp8_gemm_impl_cdna3( |
There was a problem hiding this comment.
Rather than a namespace for gfx942 and gfx950, etc, I think we would be better off with a purely virtual class at the parent folder level, with child class definitions for functions within each architectural folder. That way we can track more easily what is available and what is not for each architecture.
| @@ -0,0 +1,1459 @@ | |||
| /************************************************************************* | |||
There was a problem hiding this comment.
Most of my comments regarding the cdna3 kernel applies here as well -- otherwise the kernel looks good.
| option(NVTE_KITTENS_USE_POWER_OF_2_SCALE | ||
| "Use HW E8M0 MFMA scaling for power-of-2 blockwise FP8 scales (gfx950)" ON) | ||
|
|
||
| function(try_enable_hipkittens_gemm) |
There was a problem hiding this comment.
Not sure this needs to bea function if we are only calling it once?
| return; | ||
| } | ||
| } | ||
| #endif |
There was a problem hiding this comment.
if hipkittens is not being used, we should add an nvte_error here for blockwise usage
Description
Integrating blockwise FP8 gemm for gfx942 and gfx950 using hipkittens. Using 2 submodules for now. The hipkittens submodule uses separate branch for CDNA3 and CDNA4.
Fixes # (issue)
Type of change
Changes
Please list the changes introduced in this PR:
separate the kernel for gfx942 and gfx950. The host selects the kernel in runtime.
The kernel supports 1d2d, 1d1d GEMM with output bf16, fp32, fp16 / input fp8e5m2, fp8e4m3 / and different epilogues.
Checklist: