From d185a93511cab71c48c44e08196c5edbc61f1064 Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Tue, 4 Aug 2026 22:33:00 +0000 Subject: [PATCH 1/3] gfx1250: fix MXFP8 scale_inv shape mismatch --- tests/cpp/test_common.cu | 29 +++++++++++++++++++++++------ 1 file changed, 23 insertions(+), 6 deletions(-) diff --git a/tests/cpp/test_common.cu b/tests/cpp/test_common.cu index 2c919488d..5bb12c9ba 100644 --- a/tests/cpp/test_common.cu +++ b/tests/cpp/test_common.cu @@ -160,14 +160,25 @@ std::pair get_scales(const NVTEShape& shape, scale_inv_meta ret_rowwise, ret_colwise; + size_t align_Y_rowwise = scale_tensor_alignment_Y_rowwise; + size_t align_X_rowwise = scale_tensor_alignment_X_rowwise; + size_t align_Y_colwise = scale_tensor_alignment_Y_colwise; + size_t align_X_colwise = scale_tensor_alignment_X_colwise; +#ifdef __HIP_PLATFORM_AMD__ + // gfx1250 MX pre-swizzle requires MXFP8 scales padded to a multiple of 4 in both dims + if (getDeviceComputeCapability() == 125) { + align_Y_rowwise = align_X_rowwise = align_Y_colwise = align_X_colwise = 4; + } +#endif + const size_t block_size_X_rowwise = 32; - size_t scale_dim_Y_rowwise = DIVUP_TO_MULTIPLE(first_dim, scale_tensor_alignment_Y_rowwise); - size_t scale_dim_X_rowwise = DIVUP_TO_MULTIPLE(DIVUP(last_dim, block_size_X_rowwise), scale_tensor_alignment_X_rowwise); + size_t scale_dim_Y_rowwise = DIVUP_TO_MULTIPLE(first_dim, align_Y_rowwise); + size_t scale_dim_X_rowwise = DIVUP_TO_MULTIPLE(DIVUP(last_dim, block_size_X_rowwise), align_X_rowwise); ret_rowwise.shape = {scale_dim_Y_rowwise, scale_dim_X_rowwise}; const size_t block_size_Y_colwise = 32; - size_t scale_dim_Y_colwise = DIVUP_TO_MULTIPLE(DIVUP(first_dim, block_size_Y_colwise), scale_tensor_alignment_Y_colwise); - size_t scale_dim_X_colwise = DIVUP_TO_MULTIPLE(last_dim, scale_tensor_alignment_X_colwise); + size_t scale_dim_Y_colwise = DIVUP_TO_MULTIPLE(DIVUP(first_dim, block_size_Y_colwise), align_Y_colwise); + size_t scale_dim_X_colwise = DIVUP_TO_MULTIPLE(last_dim, align_X_colwise); ret_colwise.shape = {scale_dim_Y_colwise, scale_dim_X_colwise}; ret_rowwise.type = DType::kFloat8E8M0; @@ -1322,8 +1333,14 @@ std::array get_scale_tensor_dims(const size_t rows, alignment_X = is_rowwise ? nvfp4_scale_tensor_alignment_X_rowwise : nvfp4_scale_tensor_alignment_X_colwise; } else { - alignment_Y = 1; - alignment_X = 1; + // MXFP8: gfx1250 MX pre-swizzle requires scales padded to a multiple of 4 in both dims + if (getDeviceComputeCapability() == 125) { + alignment_Y = 4; + alignment_X = 4; + } else { + alignment_Y = 1; + alignment_X = 1; + } } #else const size_t alignment_Y = is_rowwise From 7948ad8961778126cee92dbfc4d07cfbebe2f74e Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Wed, 5 Aug 2026 21:45:34 +0000 Subject: [PATCH 2/3] factor out 4 --- tests/cpp/test_common.cu | 12 ++++-------- tests/cpp/test_common.h | 2 ++ 2 files changed, 6 insertions(+), 8 deletions(-) diff --git a/tests/cpp/test_common.cu b/tests/cpp/test_common.cu index 5bb12c9ba..8888d489d 100644 --- a/tests/cpp/test_common.cu +++ b/tests/cpp/test_common.cu @@ -167,7 +167,8 @@ std::pair get_scales(const NVTEShape& shape, #ifdef __HIP_PLATFORM_AMD__ // gfx1250 MX pre-swizzle requires MXFP8 scales padded to a multiple of 4 in both dims if (getDeviceComputeCapability() == 125) { - align_Y_rowwise = align_X_rowwise = align_Y_colwise = align_X_colwise = 4; + align_Y_rowwise = align_X_rowwise = align_Y_colwise = align_X_colwise = + mxfp8_gfx1250_scale_tensor_alignment; } #endif @@ -1334,13 +1335,8 @@ std::array get_scale_tensor_dims(const size_t rows, : nvfp4_scale_tensor_alignment_X_colwise; } else { // MXFP8: gfx1250 MX pre-swizzle requires scales padded to a multiple of 4 in both dims - if (getDeviceComputeCapability() == 125) { - alignment_Y = 4; - alignment_X = 4; - } else { - alignment_Y = 1; - alignment_X = 1; - } + alignment_Y = alignment_X = + (getDeviceComputeCapability() == 125) ? mxfp8_gfx1250_scale_tensor_alignment : 1; } #else const size_t alignment_Y = is_rowwise diff --git a/tests/cpp/test_common.h b/tests/cpp/test_common.h index 663f2e764..d0f068851 100644 --- a/tests/cpp/test_common.h +++ b/tests/cpp/test_common.h @@ -422,6 +422,8 @@ constexpr size_t scale_tensor_alignment_X_rowwise = 1; constexpr size_t scale_tensor_alignment_Y_rowwise = 1; constexpr size_t scale_tensor_alignment_X_colwise = 1; constexpr size_t scale_tensor_alignment_Y_colwise = 1; +// gfx1250 MX pre-swizzle pads MXFP8 scales to a multiple of 4 in both dims +constexpr size_t mxfp8_gfx1250_scale_tensor_alignment = 4; // For nvfp4: constexpr size_t nvfp4_scale_tensor_alignment_Y_rowwise = 128; From 3a658c2ef7538ce25039373d370e27e008ff14f4 Mon Sep 17 00:00:00 2001 From: Matthias Diener Date: Wed, 5 Aug 2026 22:04:39 +0000 Subject: [PATCH 3/3] restructure --- tests/cpp/test_common.cu | 37 +++++++++++++++++++------------------ 1 file changed, 19 insertions(+), 18 deletions(-) diff --git a/tests/cpp/test_common.cu b/tests/cpp/test_common.cu index 8888d489d..7d2c584da 100644 --- a/tests/cpp/test_common.cu +++ b/tests/cpp/test_common.cu @@ -160,28 +160,27 @@ std::pair get_scales(const NVTEShape& shape, scale_inv_meta ret_rowwise, ret_colwise; - size_t align_Y_rowwise = scale_tensor_alignment_Y_rowwise; - size_t align_X_rowwise = scale_tensor_alignment_X_rowwise; - size_t align_Y_colwise = scale_tensor_alignment_Y_colwise; - size_t align_X_colwise = scale_tensor_alignment_X_colwise; -#ifdef __HIP_PLATFORM_AMD__ - // gfx1250 MX pre-swizzle requires MXFP8 scales padded to a multiple of 4 in both dims - if (getDeviceComputeCapability() == 125) { - align_Y_rowwise = align_X_rowwise = align_Y_colwise = align_X_colwise = - mxfp8_gfx1250_scale_tensor_alignment; - } -#endif - const size_t block_size_X_rowwise = 32; - size_t scale_dim_Y_rowwise = DIVUP_TO_MULTIPLE(first_dim, align_Y_rowwise); - size_t scale_dim_X_rowwise = DIVUP_TO_MULTIPLE(DIVUP(last_dim, block_size_X_rowwise), align_X_rowwise); + size_t scale_dim_Y_rowwise = DIVUP_TO_MULTIPLE(first_dim, scale_tensor_alignment_Y_rowwise); + size_t scale_dim_X_rowwise = DIVUP_TO_MULTIPLE(DIVUP(last_dim, block_size_X_rowwise), scale_tensor_alignment_X_rowwise); ret_rowwise.shape = {scale_dim_Y_rowwise, scale_dim_X_rowwise}; const size_t block_size_Y_colwise = 32; - size_t scale_dim_Y_colwise = DIVUP_TO_MULTIPLE(DIVUP(first_dim, block_size_Y_colwise), align_Y_colwise); - size_t scale_dim_X_colwise = DIVUP_TO_MULTIPLE(last_dim, align_X_colwise); + size_t scale_dim_Y_colwise = DIVUP_TO_MULTIPLE(DIVUP(first_dim, block_size_Y_colwise), scale_tensor_alignment_Y_colwise); + size_t scale_dim_X_colwise = DIVUP_TO_MULTIPLE(last_dim, scale_tensor_alignment_X_colwise); ret_colwise.shape = {scale_dim_Y_colwise, scale_dim_X_colwise}; +#ifdef __HIP_PLATFORM_AMD__ + // gfx1250 MX pre-swizzle pads MXFP8 scales to a multiple of 4 in both dims + if (getDeviceComputeCapability() == 125) { + const size_t align = mxfp8_gfx1250_scale_tensor_alignment; + ret_rowwise.shape = {DIVUP_TO_MULTIPLE(ret_rowwise.shape[0], align), + DIVUP_TO_MULTIPLE(ret_rowwise.shape[1], align)}; + ret_colwise.shape = {DIVUP_TO_MULTIPLE(ret_colwise.shape[0], align), + DIVUP_TO_MULTIPLE(ret_colwise.shape[1], align)}; + } +#endif + ret_rowwise.type = DType::kFloat8E8M0; ret_rowwise.type_size_bits = typeToNumBits(DType::kFloat8E8M0); ret_colwise.type = DType::kFloat8E8M0; @@ -1333,10 +1332,12 @@ std::array get_scale_tensor_dims(const size_t rows, : nvfp4_scale_tensor_alignment_Y_colwise; alignment_X = is_rowwise ? nvfp4_scale_tensor_alignment_X_rowwise : nvfp4_scale_tensor_alignment_X_colwise; - } else { - // MXFP8: gfx1250 MX pre-swizzle requires scales padded to a multiple of 4 in both dims + } else if (scaling_mode == NVTE_MXFP8_1D_SCALING) { + // gfx1250 MX pre-swizzle requires MXFP8 scales padded to a multiple of 4 in both dims (1 on other architectures) alignment_Y = alignment_X = (getDeviceComputeCapability() == 125) ? mxfp8_gfx1250_scale_tensor_alignment : 1; + } else { + alignment_Y = alignment_X = 1; } #else const size_t alignment_Y = is_rowwise