diff --git a/tests/cpp/operator/CMakeLists.txt b/tests/cpp/operator/CMakeLists.txt index 85bad455d2..e9b0f4f133 100644 --- a/tests/cpp/operator/CMakeLists.txt +++ b/tests/cpp/operator/CMakeLists.txt @@ -15,6 +15,7 @@ add_executable(test_operator test_cast_mxfp8_grouped.cu test_cast_mxfp8_grouped_scaled_swiglu.cu test_cast_nvfp4_transpose.cu + test_cast_nvfp4_transpose_grouped.cu test_cast_float8blockwise.cu test_cast_float8blockwise_grouped.cu test_dequantize_mxfp8.cu diff --git a/tests/cpp/operator/test_cast_nvfp4_transpose_grouped.cu b/tests/cpp/operator/test_cast_nvfp4_transpose_grouped.cu new file mode 100644 index 0000000000..b6e77de77c --- /dev/null +++ b/tests/cpp/operator/test_cast_nvfp4_transpose_grouped.cu @@ -0,0 +1,1055 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +#include +#include +#include +#include +#include + +#include +#include "../test_common.h" +#include "transformer_engine/transformer_engine.h" +#include +#include +#include +#include + +using namespace transformer_engine; +using namespace test; + +namespace { + +enum ShapeRepresentation { + SAME_BOTH_DIMS = 0, + VARYING_FIRST_DIM = 1, + VARYING_LAST_DIM = 2, + VARYING_BOTH_DIMS = 3 +}; + +using nvfp4_scale_t = fp8e4m3; + +double2 cvt_fp4x2_to_double2(fp4e2m1x2 fp4_pair) { + const __half2_raw raw_truncated_to_fp4e2m1_pair = + __nv_cvt_fp4x2_to_halfraw2(*reinterpret_cast<__nv_fp4x2_storage_t*>(&fp4_pair), __NV_E2M1); + + const __half2 truncated_to_fp4e2m1_pair(raw_truncated_to_fp4e2m1_pair); + const double truncated_to_fp4e2m1_x = static_cast(truncated_to_fp4e2m1_pair.x); + const double truncated_to_fp4e2m1_y = static_cast(truncated_to_fp4e2m1_pair.y); + return {truncated_to_fp4e2m1_x, truncated_to_fp4e2m1_y}; +} + +template +std::vector create_transpose(const InputType* const input, const size_t rows, size_t cols) { + std::vector input_t(cols * rows); + for (size_t i = 0; i < rows; ++i) { + for (size_t j = 0; j < cols; ++j) { + const size_t idx = i * cols + j; + const size_t idx_t = j * rows + i; + input_t[idx_t] = input[idx]; + } + } + return input_t; +} + +template +struct TypeExtrema; + +template <> +struct TypeExtrema { + // Hex float format of 1.(7 bits of 1) * 2 ^ 127 + static constexpr float max = 0x1.FEp127; +}; + +template +struct TypeExtrema { + static constexpr float max = std::numeric_limits::max(); +}; + +// Compute "correct" per-block encoding scaling factor +float compute_scaling_coefficient(const nvfp4_scale_t S_dec_block, const float S_enc, + const bool use_fast_math) { + const float S_dec_block_as_fp32 = static_cast(S_dec_block); + float scale_rcp = 0.0f; + if (use_fast_math) { + scale_rcp = fminf(S_enc / S_dec_block_as_fp32, TypeExtrema::max); + scale_rcp = static_cast(static_cast(scale_rcp)); + } else { + const float S_dec = 1.0f / S_enc; + scale_rcp = fminf(1.0f / (S_dec_block_as_fp32 * S_dec), TypeExtrema::max); + } + return scale_rcp; +} + +nvfp4_scale_t compute_decoding_scaling_factor(const float block_amax, const float S_enc) { + constexpr float fp4_max = 6.0f; + const float S_dec_b = block_amax / fp4_max * S_enc; + return static_cast(fminf(S_dec_b, TypeExtrema::max)); +} + +// Compute the global encode scale factor for a given global amax +float compute_global_encode_scaling_factor(const float global_amax) { + constexpr float fp8_max = 448.0f; // 448.0f; + constexpr float fp4_max = 6.0f; // 6.0f; + float global_encode_scale = fp8_max * fp4_max / global_amax; + + global_encode_scale = fminf(global_encode_scale, Numeric_Traits::maxNorm); + // If global amax is 0 or infinity, return 1 + if (global_amax == 0.0f || global_encode_scale == 0.0f) { + return 1.0f; + } + return global_encode_scale; +} + +// 1D Scaling: Original implementation with 1x16 blocks +template +void quantize_nvfp4(const InputType* const input, + fp4e2m1x2* const output, + fp8e4m3* const scales, + const size_t rows, + const size_t cols, + const size_t scales_stride, + const float global_amax, + const bool use_fast_math) { + + // Compute a global encoding/decoding scaling factor for all S_dec_b + const float S_enc = compute_global_encode_scaling_factor(global_amax); + + constexpr size_t block_size_X = 16; + const size_t blocks_X = divide_round_up(cols, block_size_X); + + std::array cache_buffer; + for (size_t i = 0; i < block_size_X; ++i) { + cache_buffer[i] = 0.0f; + } + + for (size_t i = 0; i < rows; ++i) { + for (size_t block_X = 0; block_X < blocks_X; ++block_X) { + const size_t j_min = block_X * block_size_X; + const size_t j_max = j_min + block_size_X; + + // Find block amax + float block_amax = 0.0f; + for (size_t j = j_min; j < j_max; ++j) { + const size_t idx = i * cols + j; + const size_t cache_idx = j - j_min; + + const float input_elt = static_cast(input[idx]); + const float act_elt = input_elt; + + // Numerical truncation: after downcast to InputType (BF16/FP16), upcast it back to FP32 + const float elt = static_cast(static_cast(act_elt)); + cache_buffer[cache_idx] = elt; + block_amax = std::max(block_amax, std::abs(elt)); + } + + // Compute E4M3 scaling factor + const nvfp4_scale_t S_dec_b_fp8 = compute_decoding_scaling_factor(block_amax, S_enc); + const float SFcoefficient = compute_scaling_coefficient(S_dec_b_fp8, S_enc, use_fast_math); + + const size_t scale_idx = i * scales_stride + block_X; + scales[scale_idx] = S_dec_b_fp8; + + for (size_t j = j_min; j < j_max; j += 2) { + const int idx_pair = (i * cols + j) / 2; + const int cache_idx_x = j - j_min; + const int cache_idx_y = cache_idx_x + 1; + const float cached_x = cache_buffer[cache_idx_x]; + const float cached_y = cache_buffer[cache_idx_y]; + const float scaled_elt_x = cached_x * SFcoefficient; + const float scaled_elt_y = cached_y * SFcoefficient; + const float2 scaled_elt_pair = {scaled_elt_x, scaled_elt_y}; + + fp4e2m1x2 casted_to_e2m1_pair(scaled_elt_pair); + output[idx_pair] = casted_to_e2m1_pair; + + // const double2 truncated_pair = cvt_fp4x2_to_double2(casted_to_e2m1_pair); + } + } + } +} + +template +void compute_ref(const InputType* input, + fp4e2m1x2* output, + fp4e2m1x2* output_t, + fp8e4m3* scales, + fp8e4m3* scales_t, + const float global_amax, + const size_t rows, + const size_t cols, + const size_t scales_stride, + const size_t scales_stride_t, + const bool use_fast_math) +{ + std::vector input_t = create_transpose(input, rows, cols); + + quantize_nvfp4(input, output, scales, rows, cols, scales_stride, global_amax, use_fast_math); + quantize_nvfp4(input_t.data(), output_t, scales_t, cols, rows, scales_stride_t, global_amax, use_fast_math); +} + +void compare_nvfp4_tensors(const std::string& name, + const fp4e2m1 *test_data, const fp4e2m1 *ref_data, + const int rows, const int cols, + double atol = 1e-5, double rtol = 1e-8) { + constexpr int max_mismatches_to_print = 3; + + std::vector mismatch_messages; + size_t total_mismatches = 0; + + for (int i = 0; i < rows; ++i) { + for (int j = 0; j < cols; j += 2) { + const int idx = i * cols + j; + double2 test_data_pair = cvt_fp4x2_to_double2(*reinterpret_cast(&test_data[idx/2])); + double2 ref_data_pair = cvt_fp4x2_to_double2(*reinterpret_cast(&ref_data[idx/2])); + + for (int k = 0; k < 2; ++k) { + const double t = (k == 0 ? test_data_pair.x : test_data_pair.y); + const double r = (k == 0 ? ref_data_pair.x : ref_data_pair.y); + + const bool mismatch = fabs(t - r) > (atol + fabs(r) * rtol); + if (mismatch) { + total_mismatches++; + // Optional: limit number of detailed messages to avoid overwhelming output + if (total_mismatches <= max_mismatches_to_print) { + std::string msg = "Mismatch at place (" + std::to_string(idx + k) + "): " + + std::to_string(t) + " vs " + std::to_string(r) + + " (abs_diff: " + std::to_string(fabs(t - r)) + + ", rel_diff: " + std::to_string(r == 0 ? 0.0 : fabs((t - r) / r)) + ")"; + mismatch_messages.push_back(msg); + std::cout << "Error in tensor " << name << ": " << msg << std::endl; + } + } + } + } + } + + bool print_detailed_summary = false; + if (print_detailed_summary) { + // Always report summary - either success or failure + std::cout << "=== SUMMARY for tensor " << name << " ===" << std::endl; + std::cout << "Total elements checked: " << (rows * cols) << std::endl; + + if (total_mismatches > 0) { + std::cout << "STATUS: FAILED for output" << std::endl; + std::cout << "Total mismatches found: " << total_mismatches << std::endl; + std::cout << "Mismatch rate: " << (100.0 * total_mismatches) / (rows * cols) << "%" << std::endl; + if (total_mismatches > max_mismatches_to_print) { + std::cout << "... and " << (total_mismatches - max_mismatches_to_print) + << " more mismatches (showing first " << max_mismatches_to_print << ")" << std::endl; + } + std::cout << "============================" << std::endl; + + GTEST_FAIL() << "Found " << total_mismatches << " mismatches in tensor " << name; + } else { + std::cout << "STATUS: PASSED for output" << std::endl; + std::cout << "All elements match within tolerance!" << std::endl; + std::cout << "Tensor " << name << " is IDENTICAL to reference" << std::endl; + std::cout << "============================" << std::endl; + } + } else { + if (total_mismatches > 0) { + GTEST_FAIL() << "Found " << total_mismatches << " mismatches in tensor " << name; + } + } +} + +template +void performTest(const ShapeRepresentation shape_rep, + const size_t num_tensors, + const std::vector& logical_shape, + const std::vector& first_dims, + const std::vector& last_dims, + const std::vector& offsets, + const bool use_fast_math) { + using namespace test; + + DType itype = TypeInfo::dtype; + DType otype = DType::kFloat4E2M1; + const size_t total_elts = offsets.back(); + std::vector grouped_input(total_elts); + + // Validate logical shape against the offsets-based flattened size. + size_t expected_total_elts = logical_shape[0] * logical_shape[1]; + if (shape_rep == VARYING_LAST_DIM) { + expected_total_elts = logical_shape[0] + * std::accumulate(last_dims.begin(), last_dims.end(), static_cast(0)); + } + ASSERT_GE(expected_total_elts, total_elts); + + Tensor grouped_input_tensor("grouped_input", std::vector{total_elts}, itype); + fillCase(&grouped_input_tensor, InputsFillCase::uniform); + std::copy(grouped_input_tensor.rowwise_cpu_dptr(), + grouped_input_tensor.rowwise_cpu_dptr() + total_elts, + grouped_input.begin()); + + const double atol = 1.0E-6; + const double rtol = 1.0E-6; + + std::vector rowwise_scales_stride(num_tensors, 0); + std::vector colwise_scales_stride(num_tensors, 0); + std::vector rowwise_unpadded_blocks_X(num_tensors, 0); + std::vector colwise_unpadded_blocks_X(num_tensors, 0); + std::vector rowwise_scale_offsets(num_tensors, 0); + std::vector colwise_scale_offsets(num_tensors, 0); + + size_t rowwise_scales_num = 0; + size_t colwise_scales_num = 0; + + for (size_t t = 0; t < num_tensors; ++t) { + const size_t rows = first_dims[t]; + const size_t cols = last_dims[t]; + rowwise_unpadded_blocks_X[t] = divide_round_up(cols, static_cast(16)); + colwise_unpadded_blocks_X[t] = divide_round_up(rows, static_cast(16)); + + rowwise_scales_stride[t] = round_up_to_nearest_multiple(rowwise_unpadded_blocks_X[t], static_cast(4)); + colwise_scales_stride[t] = round_up_to_nearest_multiple(colwise_unpadded_blocks_X[t], static_cast(4)); + + rowwise_scale_offsets[t] = rowwise_scales_num; + colwise_scale_offsets[t] = colwise_scales_num; + + rowwise_scales_num += rows * rowwise_scales_stride[t]; + colwise_scales_num += cols * colwise_scales_stride[t]; + } + + std::vector out_data_rowwise_h(total_elts / 2); + std::vector out_data_colwise_h(total_elts / 2); + std::vector out_scales_rowwise_h(rowwise_scales_num); + std::vector out_scales_colwise_h(colwise_scales_num); + + std::vector out_data_rowwise_ref(total_elts / 2); + std::vector out_data_colwise_ref(total_elts / 2); + std::vector> out_scales_rowwise_ref(num_tensors); + std::vector> out_scales_colwise_ref(num_tensors); + std::vector amax_per_tensor(num_tensors, 0.0f); + + for (size_t t = 0; t < num_tensors; ++t) { + const size_t rows = first_dims[t]; + const size_t cols = last_dims[t]; + const size_t tensor_offset = offsets[t]; + const size_t tensor_numel = rows * cols; + ASSERT_EQ(offsets[t + 1] - offsets[t], tensor_numel); + ASSERT_LE(tensor_offset + tensor_numel, total_elts); + ASSERT_EQ(tensor_numel % 2, 0U); + + float amax = 0.0f; + for (size_t i = 0; i < tensor_numel; ++i) { + amax = fmaxf(amax, fabs(static_cast(grouped_input[tensor_offset + i]))); + } + amax_per_tensor[t] = amax; + + std::unique_ptr ref_output = std::make_unique(tensor_numel / 2); + std::unique_ptr ref_output_t = std::make_unique(tensor_numel / 2); + std::unique_ptr ref_scales = std::make_unique(rows * rowwise_scales_stride[t]); + std::unique_ptr ref_scales_t = std::make_unique(cols * colwise_scales_stride[t]); + + compute_ref(grouped_input.data() + tensor_offset, + ref_output.get(), + ref_output_t.get(), + ref_scales.get(), + ref_scales_t.get(), + amax_per_tensor[t], + rows, + cols, + rowwise_scales_stride[t], + colwise_scales_stride[t], + use_fast_math); + + std::memcpy(out_data_rowwise_ref.data() + tensor_offset / 2, ref_output.get(), + (tensor_numel / 2) * sizeof(fp4e2m1x2)); + std::memcpy(out_data_colwise_ref.data() + tensor_offset / 2, ref_output_t.get(), + (tensor_numel / 2) * sizeof(fp4e2m1x2)); + + out_scales_rowwise_ref[t].assign(ref_scales.get(), ref_scales.get() + rows * rowwise_scales_stride[t]); + out_scales_colwise_ref[t].assign(ref_scales_t.get(), ref_scales_t.get() + cols * colwise_scales_stride[t]); + } + + const size_t in_data_size = total_elts * sizeof(InputType); + const size_t out_data_size = (total_elts * typeToNumBits(otype)) / 8; + const size_t rowwise_scales_size = rowwise_scales_num * sizeof(fp8e4m3); + const size_t colwise_scales_size = colwise_scales_num * sizeof(fp8e4m3); + const size_t amax_size = num_tensors * sizeof(float); + + std::vector first_dims_h(num_tensors, 0); + std::vector last_dims_h(num_tensors, 0); + std::vector offsets_h(num_tensors + 1, 0); + for (size_t t = 0; t < num_tensors; ++t) { + first_dims_h[t] = static_cast(first_dims[t]); + last_dims_h[t] = static_cast(last_dims[t]); + } + for (size_t t = 0; t < num_tensors + 1; ++t) { + offsets_h[t] = static_cast(offsets[t]); + } + + InputType* in_data_d = nullptr; + fp4e2m1* out_data_rowwise_d = nullptr; + fp4e2m1* out_data_colwise_d = nullptr; + fp8e4m3* out_scales_rowwise_d = nullptr; + fp8e4m3* out_scales_colwise_d = nullptr; + float* out_amax_rowwise_d = nullptr; + float* out_amax_colwise_d = nullptr; + int64_t* first_dims_d = nullptr; + int64_t* last_dims_d = nullptr; + int64_t* offsets_d = nullptr; + + cudaMalloc((void**)&in_data_d, in_data_size); + cudaMalloc((void**)&out_data_rowwise_d, out_data_size); + cudaMalloc((void**)&out_data_colwise_d, out_data_size); + cudaMalloc((void**)&out_scales_rowwise_d, rowwise_scales_size); + cudaMalloc((void**)&out_scales_colwise_d, colwise_scales_size); + cudaMalloc((void**)&out_amax_rowwise_d, amax_size); + cudaMalloc((void**)&out_amax_colwise_d, amax_size); + + cudaMalloc((void**)&first_dims_d, num_tensors * sizeof(int64_t)); + cudaMalloc((void**)&last_dims_d, num_tensors * sizeof(int64_t)); + cudaMalloc((void**)&offsets_d, (num_tensors + 1) * sizeof(int64_t)); + + cudaMemcpy(in_data_d, grouped_input.data(), in_data_size, cudaMemcpyHostToDevice); + cudaMemcpy(out_amax_rowwise_d, amax_per_tensor.data(), amax_size, cudaMemcpyHostToDevice); + cudaMemcpy(out_amax_colwise_d, amax_per_tensor.data(), amax_size, cudaMemcpyHostToDevice); + cudaMemcpy(first_dims_d, first_dims_h.data(), num_tensors * sizeof(int64_t), cudaMemcpyHostToDevice); + cudaMemcpy(last_dims_d, last_dims_h.data(), num_tensors * sizeof(int64_t), cudaMemcpyHostToDevice); + cudaMemcpy(offsets_d, offsets_h.data(), (num_tensors + 1) * sizeof(int64_t), cudaMemcpyHostToDevice); + + cudaMemset(out_data_rowwise_d, 0, out_data_size); + cudaMemset(out_data_colwise_d, 0, out_data_size); + cudaMemset(out_scales_rowwise_d, 0, rowwise_scales_size); + cudaMemset(out_scales_colwise_d, 0, colwise_scales_size); + + NVTEShape logical_shape_ = nvte_make_shape(logical_shape.data(), logical_shape.size()); + + NVTEShape first_dims_shape_; + NVTEShape last_dims_shape_; + NVTEShape offsets_shape_; + first_dims_shape_.ndim = 1; + last_dims_shape_.ndim = 1; + offsets_shape_.ndim = 1; + first_dims_shape_.data[0] = num_tensors; + last_dims_shape_.data[0] = num_tensors; + offsets_shape_.data[0] = num_tensors + 1; + + NVTEGroupedTensor in_group_tensor = nvte_create_grouped_tensor(NVTE_DELAYED_TENSOR_SCALING, num_tensors, logical_shape_); + NVTEGroupedTensor out_group_tensor = nvte_create_grouped_tensor(NVTE_NVFP4_1D_SCALING, num_tensors, logical_shape_); + + NVTEBasicTensor in_data_tensor = {in_data_d, static_cast(itype), logical_shape_}; + nvte_set_grouped_tensor_param(in_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedRowwiseData, + &in_data_tensor, sizeof(in_data_tensor)); + + NVTEBasicTensor out_data_rowwise_tensor = {out_data_rowwise_d, NVTEDType::kNVTEFloat4E2M1, logical_shape_}; + NVTEBasicTensor out_data_colwise_tensor = {out_data_colwise_d, NVTEDType::kNVTEFloat4E2M1, logical_shape_}; + nvte_set_grouped_tensor_param(out_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedRowwiseData, + &out_data_rowwise_tensor, sizeof(out_data_rowwise_tensor)); + nvte_set_grouped_tensor_param(out_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedColumnwiseData, + &out_data_colwise_tensor, sizeof(out_data_colwise_tensor)); + + std::vector rowwise_scales_shape = {rowwise_scales_num}; + std::vector colwise_scales_shape = {colwise_scales_num}; + std::vector amax_shape = {num_tensors}; + NVTEShape rowwise_scales_shape_ = nvte_make_shape(rowwise_scales_shape.data(), rowwise_scales_shape.size()); + NVTEShape colwise_scales_shape_ = nvte_make_shape(colwise_scales_shape.data(), colwise_scales_shape.size()); + NVTEShape amax_shape_ = nvte_make_shape(amax_shape.data(), amax_shape.size()); + NVTEBasicTensor out_scales_rowwise_tensor = { out_scales_rowwise_d, NVTEDType::kNVTEFloat8E4M3, rowwise_scales_shape_}; + NVTEBasicTensor out_scales_colwise_tensor = { out_scales_colwise_d, NVTEDType::kNVTEFloat8E4M3, colwise_scales_shape_}; + NVTEBasicTensor out_amax_rowwise_tensor = {out_amax_rowwise_d, NVTEDType::kNVTEFloat32, amax_shape_}; + NVTEBasicTensor out_amax_colwise_tensor = {out_amax_colwise_d, NVTEDType::kNVTEFloat32, amax_shape_}; + nvte_set_grouped_tensor_param(out_group_tensor, + NVTEGroupedTensorParam::kNVTEGroupedRowwiseScaleInv, + &out_scales_rowwise_tensor, sizeof(out_scales_rowwise_tensor)); + nvte_set_grouped_tensor_param(out_group_tensor, + NVTEGroupedTensorParam::kNVTEGroupedColumnwiseScaleInv, + &out_scales_colwise_tensor, sizeof(out_scales_colwise_tensor)); + nvte_set_grouped_tensor_param(in_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedAmax, + &out_amax_rowwise_tensor, sizeof(out_amax_rowwise_tensor)); + nvte_set_grouped_tensor_param(in_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedColumnwiseAmax, + &out_amax_colwise_tensor, sizeof(out_amax_colwise_tensor)); + + if ((shape_rep == VARYING_FIRST_DIM) || (shape_rep == VARYING_BOTH_DIMS)) { + NVTEBasicTensor first_dims_tensor = {first_dims_d, kNVTEInt64, first_dims_shape_}; + nvte_set_grouped_tensor_param(in_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedFirstDims, + &first_dims_tensor, sizeof(first_dims_tensor)); + nvte_set_grouped_tensor_param(out_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedFirstDims, + &first_dims_tensor, sizeof(first_dims_tensor)); + } + + if ((shape_rep == VARYING_LAST_DIM) || (shape_rep == VARYING_BOTH_DIMS)) { + NVTEBasicTensor last_dims_tensor = {last_dims_d, kNVTEInt64, last_dims_shape_}; + nvte_set_grouped_tensor_param(in_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedLastDims, + &last_dims_tensor, sizeof(last_dims_tensor)); + nvte_set_grouped_tensor_param(out_group_tensor, NVTEGroupedTensorParam::kNVTEGroupedLastDims, + &last_dims_tensor, sizeof(last_dims_tensor)); + } + + if (shape_rep != SAME_BOTH_DIMS) { + NVTEBasicTensor offsets_tensor = {offsets_d, kNVTEInt64, offsets_shape_}; + nvte_set_grouped_tensor_param(in_group_tensor, + NVTEGroupedTensorParam::kNVTEGroupedTensorOffsets, + &offsets_tensor, sizeof(offsets_tensor)); + nvte_set_grouped_tensor_param(out_group_tensor, + NVTEGroupedTensorParam::kNVTEGroupedTensorOffsets, + &offsets_tensor, sizeof(offsets_tensor)); + } + + QuantizationConfigWrapper quant_config; + quant_config.set_use_fast_math(use_fast_math); + quant_config.set_stochastic_rounding(false); + + nvte_group_quantize(in_group_tensor, out_group_tensor, quant_config, 0); + cudaDeviceSynchronize(); + auto err = cudaGetLastError(); + ASSERT_EQ(err, cudaSuccess) << cudaGetErrorString(err); + + cudaMemcpy(out_data_rowwise_h.data(), out_data_rowwise_d, out_data_size, cudaMemcpyDeviceToHost); + cudaMemcpy(out_data_colwise_h.data(), out_data_colwise_d, out_data_size, cudaMemcpyDeviceToHost); + cudaMemcpy(out_scales_rowwise_h.data(), out_scales_rowwise_d, rowwise_scales_size, cudaMemcpyDeviceToHost); + cudaMemcpy(out_scales_colwise_h.data(), out_scales_colwise_d, colwise_scales_size, cudaMemcpyDeviceToHost); + + for (size_t t = 0; t < num_tensors; ++t) { + const size_t rows = first_dims[t]; + const size_t cols = last_dims[t]; + const size_t tensor_offset = offsets[t]; + + const fp4e2m1* test_output = out_data_rowwise_h.data() + tensor_offset / 2; + const fp4e2m1* ref_output = out_data_rowwise_ref.data() + tensor_offset / 2; + const fp4e2m1* test_output_t = out_data_colwise_h.data() + tensor_offset / 2; + const fp4e2m1* ref_output_t = out_data_colwise_ref.data() + tensor_offset / 2; + + compare_nvfp4_tensors("output_" + std::to_string(t), test_output, ref_output, + static_cast(rows), static_cast(cols), atol, rtol); + compare_nvfp4_tensors("output_t_" + std::to_string(t), test_output_t, ref_output_t, + static_cast(cols), static_cast(rows), atol, rtol); + + size_t scale_mismatches_num = 0; + compare_scaling_factors( + "scales_" + std::to_string(t), + out_scales_rowwise_h.data() + rowwise_scale_offsets[t], + out_scales_rowwise_ref[t].data(), + rows, rowwise_unpadded_blocks_X[t], rowwise_scales_stride[t], scale_mismatches_num); + + compare_scaling_factors( + "scales_t_" + std::to_string(t), + out_scales_colwise_h.data() + colwise_scale_offsets[t], + out_scales_colwise_ref[t].data(), + cols, colwise_unpadded_blocks_X[t], colwise_scales_stride[t], scale_mismatches_num); + } + + nvte_destroy_grouped_tensor(in_group_tensor); + nvte_destroy_grouped_tensor(out_group_tensor); + + cudaFree(in_data_d); + cudaFree(out_data_rowwise_d); + cudaFree(out_data_colwise_d); + cudaFree(out_scales_rowwise_d); + cudaFree(out_scales_colwise_d); + cudaFree(out_amax_rowwise_d); + cudaFree(out_amax_colwise_d); + cudaFree(first_dims_d); + cudaFree(last_dims_d); + cudaFree(offsets_d); +} + +// {shape_representation, num_tensors, [logical_shape_M, logical_shape_K], [M_i], [K_i]} +std::vector> grouped_input_config = { + {SAME_BOTH_DIMS, 1, 1024,1024}, + {SAME_BOTH_DIMS, 1, 2048,2048}, + {SAME_BOTH_DIMS, 1, 3072,3072}, + {SAME_BOTH_DIMS, 1, 4096,4096}, + {SAME_BOTH_DIMS, 1, 5120,5120}, + {SAME_BOTH_DIMS, 1, 6144,6144}, + {SAME_BOTH_DIMS, 1, 7168,7168}, + {SAME_BOTH_DIMS, 1, 8192,8192}, + {SAME_BOTH_DIMS, 1, 10240,10240}, + {SAME_BOTH_DIMS, 1, 12288,12288}, + {SAME_BOTH_DIMS, 1, 14336,14336}, + {SAME_BOTH_DIMS, 1, 16384,16384}, + + {SAME_BOTH_DIMS, 2, 1024,1024}, + {SAME_BOTH_DIMS, 2, 2048,2048}, + {SAME_BOTH_DIMS, 2, 3072,3072}, + {SAME_BOTH_DIMS, 2, 4096,4096}, + {SAME_BOTH_DIMS, 2, 5120,5120}, + {SAME_BOTH_DIMS, 2, 6144,6144}, + {SAME_BOTH_DIMS, 2, 7168,7168}, + {SAME_BOTH_DIMS, 2, 8192,8192}, + {SAME_BOTH_DIMS, 2, 10240,10240}, + {SAME_BOTH_DIMS, 2, 12288,12288}, + {SAME_BOTH_DIMS, 2, 14336,14336}, + {SAME_BOTH_DIMS, 2, 16384,16384}, + + {SAME_BOTH_DIMS, 4, 1024,1024}, + {SAME_BOTH_DIMS, 4, 2048,2048}, + {SAME_BOTH_DIMS, 4, 3072,3072}, + {SAME_BOTH_DIMS, 4, 4096,4096}, + {SAME_BOTH_DIMS, 4, 5120,5120}, + {SAME_BOTH_DIMS, 4, 6144,6144}, + {SAME_BOTH_DIMS, 4, 7168,7168}, + {SAME_BOTH_DIMS, 4, 8192,8192}, + {SAME_BOTH_DIMS, 4, 10240,10240}, + {SAME_BOTH_DIMS, 4, 12288,12288}, + {SAME_BOTH_DIMS, 4, 14336,14336}, + {SAME_BOTH_DIMS, 4, 16384,16384}, + + {SAME_BOTH_DIMS, 8, 1024,1024}, + {SAME_BOTH_DIMS, 8, 2048,2048}, + {SAME_BOTH_DIMS, 8, 3072,3072}, + {SAME_BOTH_DIMS, 8, 4096,4096}, + {SAME_BOTH_DIMS, 8, 5120,5120}, + {SAME_BOTH_DIMS, 8, 6144,6144}, + {SAME_BOTH_DIMS, 8, 7168,7168}, + {SAME_BOTH_DIMS, 8, 8192,8192}, + {SAME_BOTH_DIMS, 8, 10240,10240}, + {SAME_BOTH_DIMS, 8, 12288,12288}, + {SAME_BOTH_DIMS, 8, 14336,14336}, + {SAME_BOTH_DIMS, 8, 16384,16384}, + + {SAME_BOTH_DIMS, 16, 1024,1024}, + {SAME_BOTH_DIMS, 16, 2048,2048}, + {SAME_BOTH_DIMS, 16, 3072,3072}, + {SAME_BOTH_DIMS, 16, 4096,4096}, + {SAME_BOTH_DIMS, 16, 5120,5120}, + {SAME_BOTH_DIMS, 16, 6144,6144}, + {SAME_BOTH_DIMS, 16, 7168,7168}, + {SAME_BOTH_DIMS, 16, 8192,8192}, + {SAME_BOTH_DIMS, 16, 10240,10240}, + {SAME_BOTH_DIMS, 16, 12288,12288}, + {SAME_BOTH_DIMS, 16, 14336,14336}, + {SAME_BOTH_DIMS, 16, 16384,16384}, + + {SAME_BOTH_DIMS, 32, 1024,1024}, + {SAME_BOTH_DIMS, 32, 2048,2048}, + {SAME_BOTH_DIMS, 32, 3072,3072}, + {SAME_BOTH_DIMS, 32, 4096,4096}, + {SAME_BOTH_DIMS, 32, 5120,5120}, + {SAME_BOTH_DIMS, 32, 6144,6144}, + {SAME_BOTH_DIMS, 32, 7168,7168}, + {SAME_BOTH_DIMS, 32, 8192,8192}, + {SAME_BOTH_DIMS, 32, 10240,10240}, + {SAME_BOTH_DIMS, 32, 12288,12288}, + {SAME_BOTH_DIMS, 32, 14336,14336}, + {SAME_BOTH_DIMS, 32, 16384,16384}, + + {SAME_BOTH_DIMS, 64, 1024,1024}, + {SAME_BOTH_DIMS, 64, 2048,2048}, + {SAME_BOTH_DIMS, 64, 3072,3072}, + {SAME_BOTH_DIMS, 64, 4096,4096}, + {SAME_BOTH_DIMS, 64, 5120,5120}, + {SAME_BOTH_DIMS, 64, 6144,6144}, + {SAME_BOTH_DIMS, 64, 7168,7168}, + {SAME_BOTH_DIMS, 64, 8192,8192}, + {SAME_BOTH_DIMS, 64, 10240,10240}, + {SAME_BOTH_DIMS, 64, 12288,12288}, + {SAME_BOTH_DIMS, 64, 14336,14336}, + {SAME_BOTH_DIMS, 64, 16384,16384}, + + {VARYING_FIRST_DIM, 1, 1024,1024, 1024}, + {VARYING_FIRST_DIM, 1, 2048,2048, 2048}, + {VARYING_FIRST_DIM, 1, 3072,3072, 3072}, + {VARYING_FIRST_DIM, 1, 4096,4096, 4096}, + {VARYING_FIRST_DIM, 1, 5120,5120, 5120}, + {VARYING_FIRST_DIM, 1, 6144,6144, 6144}, + {VARYING_FIRST_DIM, 1, 7168,7168, 7168}, + {VARYING_FIRST_DIM, 1, 8192,8192, 8192}, + {VARYING_FIRST_DIM, 1, 10240,10240, 10240}, + {VARYING_FIRST_DIM, 1, 12288,12288, 12288}, + {VARYING_FIRST_DIM, 1, 14336,14336, 14336}, + {VARYING_FIRST_DIM, 1, 16384,16384, 16384}, + + {VARYING_FIRST_DIM, 2, 1024,1024, 512,512}, + {VARYING_FIRST_DIM, 2, 2048,2048, 1024,1024}, + {VARYING_FIRST_DIM, 2, 3072,3072, 1536,1536}, + {VARYING_FIRST_DIM, 2, 4096,4096, 2048,2048}, + {VARYING_FIRST_DIM, 2, 5120,5120, 2560,2560}, + {VARYING_FIRST_DIM, 2, 6144,6144, 3072,3072}, + {VARYING_FIRST_DIM, 2, 7168,7168, 3584,3584}, + {VARYING_FIRST_DIM, 2, 8192,8192, 4096,4096}, + {VARYING_FIRST_DIM, 2, 10240,10240, 5120,5120}, + {VARYING_FIRST_DIM, 2, 12288,12288, 6144,6144}, + {VARYING_FIRST_DIM, 2, 14336,14336, 7168,7168}, + {VARYING_FIRST_DIM, 2, 16384,16384, 8192,8192}, + + {VARYING_FIRST_DIM, 4, 1024,1024, 256,256,256,256}, + {VARYING_FIRST_DIM, 4, 2048,2048, 512,512,512,512}, + {VARYING_FIRST_DIM, 4, 3072,3072, 768,768,768,768}, + {VARYING_FIRST_DIM, 4, 4096,4096, 1024,1024,1024,1024}, + {VARYING_FIRST_DIM, 4, 5120,5120, 1280,1280,1280,1280}, + {VARYING_FIRST_DIM, 4, 6144,6144, 1536,1536,1536,1536}, + {VARYING_FIRST_DIM, 4, 7168,7168, 1792,1792,1792,1792}, + {VARYING_FIRST_DIM, 4, 8192,8192, 2048,2048,2048,2048}, + {VARYING_FIRST_DIM, 4, 10240,10240, 2560,2560,2560,2560}, + {VARYING_FIRST_DIM, 4, 12288,12288, 3072,3072,3072,3072}, + {VARYING_FIRST_DIM, 4, 14336,14336, 3584,3584,3584,3584}, + {VARYING_FIRST_DIM, 4, 16384,16384, 4096,4096,4096,4096}, + + {VARYING_FIRST_DIM, 8, 1024,1024, 128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 8, 2048,2048, 256,256,256,256,256,256,256,256}, + {VARYING_FIRST_DIM, 8, 3072,3072, 384,384,384,384,384,384,384,384}, + {VARYING_FIRST_DIM, 8, 4096,4096, 512,512,512,512,512,512,512,512}, + {VARYING_FIRST_DIM, 8, 5120,5120, 640,640,640,640,640,640,640,640}, + {VARYING_FIRST_DIM, 8, 6144,6144, 768,768,768,768,768,768,768,768}, + {VARYING_FIRST_DIM, 8, 7168,7168, 896,896,896,896,896,896,896,896}, + {VARYING_FIRST_DIM, 8, 8192,8192, 1024,1024,1024,1024,1024,1024,1024,1024}, + {VARYING_FIRST_DIM, 8, 10240,10240, 1280,1280,1280,1280,1280,1280,1280,1280}, + {VARYING_FIRST_DIM, 8, 12288,12288, 1536,1536,1536,1536,1536,1536,1536,1536}, + {VARYING_FIRST_DIM, 8, 14336,14336, 1792,1792,1792,1792,1792,1792,1792,1792}, + {VARYING_FIRST_DIM, 8, 16384,16384, 2048,2048,2048,2048,2048,2048,2048,2048}, + + {VARYING_FIRST_DIM, 16, 1024,1024, 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 16, 2048,2048, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 16, 3072,3072, 256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 16, 4096,4096, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256}, + {VARYING_FIRST_DIM, 16, 5120,5120, 512,512,512,512,256,256,256,256,256,256,256,256,256,256,256,256}, + {VARYING_FIRST_DIM, 16, 6144,6144, 384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384}, + {VARYING_FIRST_DIM, 16, 7168,7168, 512,512,512,512,512,512,512,512,512,512,512,512,256,256,256,256}, + {VARYING_FIRST_DIM, 16, 8192,8192, 512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512}, + {VARYING_FIRST_DIM, 16, 10240,10240, 640,640,640,640,640,640,640,640,640,640,640,640,640,640,640,640}, + {VARYING_FIRST_DIM, 16, 12288,12288, 768,768,768,768,768,768,768,768,768,768,768,768,768,768,768,768}, + {VARYING_FIRST_DIM, 16, 14336,14336, 896,896,896,896,896,896,896,896,896,896,896,896,896,896,896,896}, + {VARYING_FIRST_DIM, 16, 16384,16384, 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024}, + + {VARYING_FIRST_DIM, 32, 1024,1024, 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 32, 2048,2048, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 32, 3072,3072, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 32, 4096,4096, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 32, 5120,5120, 256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 32, 6144,6144, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 32, 7168,7168, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 32, 8192,8192, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256}, + {VARYING_FIRST_DIM, 32, 10240,10240, 384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384, + 384,384,384,384,256,256,256,256,256,256,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 32, 12288,12288, 384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384, + 384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384}, + {VARYING_FIRST_DIM, 32, 14336,14336, 512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512, + 512,512,512,512,512,512,512,512,256,256,256,256,256,256,256,256,}, + {VARYING_FIRST_DIM, 32, 16384,16384, 512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512, + 512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512}, + + {VARYING_FIRST_DIM, 64, 1024,1024, 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 64, 2048,2048, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 64, 3072,3072, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 64, 4096,4096, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 64, 5120,5120, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 64, 6144,6144, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 64, 7168,7168, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0}, + {VARYING_FIRST_DIM, 64, 8192,8192, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 64, 10240,10240, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 64, 12288,12288, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 64, 14336,14336, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128}, + {VARYING_FIRST_DIM, 64, 16384,16384, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256}, + + // {VARYING_BOTH_DIMS, 1, 1,1024*1024, 1024, 1024}, + // {VARYING_BOTH_DIMS, 1, 1,2048*2048, 2048, 2048}, + // {VARYING_BOTH_DIMS, 1, 1,3072*3072, 3072, 3072}, + // {VARYING_BOTH_DIMS, 1, 1,4096*4096, 4096, 4096}, + // {VARYING_BOTH_DIMS, 1, 1,5120*5120, 5120, 5120}, + // {VARYING_BOTH_DIMS, 1, 1,6144*6144, 6144, 6144}, + // {VARYING_BOTH_DIMS, 1, 1,7168*7168, 7168, 7168}, + // {VARYING_BOTH_DIMS, 1, 1,8192*8192, 8192, 8192}, + // {VARYING_BOTH_DIMS, 1, 1,10240*10240, 10240, 10240}, + // {VARYING_BOTH_DIMS, 1, 1,12288*12288, 12288, 12288}, + // {VARYING_BOTH_DIMS, 1, 1,14336*14336, 14336, 14336}, + // {VARYING_BOTH_DIMS, 1, 1,16384*16384, 16384, 16384}, + + // {VARYING_BOTH_DIMS, 2, 1,1024*1024, 512,512, 1024,1024}, + // {VARYING_BOTH_DIMS, 2, 1,2048*2048, 1024,1024, 2048,2048}, + // {VARYING_BOTH_DIMS, 2, 1,3072*3072, 1536,1536, 3072,3072}, + // {VARYING_BOTH_DIMS, 2, 1,4096*4096, 2048,2048, 4096,4096}, + // {VARYING_BOTH_DIMS, 2, 1,5120*5120, 2560,2560, 5120,5120}, + // {VARYING_BOTH_DIMS, 2, 1,6144*6144, 3072,3072, 6144,6144}, + // {VARYING_BOTH_DIMS, 2, 1,7168*7168, 3584,3584, 7168,7168}, + // {VARYING_BOTH_DIMS, 2, 1,8192*8192, 4096,4096, 8192,8192}, + // {VARYING_BOTH_DIMS, 2, 1,10240*10240, 5120,5120, 10240,10240}, + // {VARYING_BOTH_DIMS, 2, 1,12288*12288, 6144,6144, 12288,12288}, + // {VARYING_BOTH_DIMS, 2, 1,14336*14336, 7168,7168, 14336,14336}, + // {VARYING_BOTH_DIMS, 2, 1,16384*16384, 8192,8192, 16384,16384}, + + // {VARYING_BOTH_DIMS, 4, 1,1024*1024, 256,256,256,256, 1024,1024,1024,1024}, + // {VARYING_BOTH_DIMS, 4, 1,2048*2048, 512,512,512,512, 2048,2048,2048,2048}, + // {VARYING_BOTH_DIMS, 4, 1,3072*3072, 768,768,768,768, 3072,3072,3072,3072}, + // {VARYING_BOTH_DIMS, 4, 1,4096*4096, 1024,1024,1024,1024, 4096,4096,4096,4096}, + // {VARYING_BOTH_DIMS, 4, 1,5120*5120, 1280,1280,1280,1280, 5120,5120,5120,5120}, + // {VARYING_BOTH_DIMS, 4, 1,6144*6144, 1536,1536,1536,1536, 6144,6144,6144,6144}, + // {VARYING_BOTH_DIMS, 4, 1,7168*7168, 1792,1792,1792,1792, 7168,7168,7168,7168}, + // {VARYING_BOTH_DIMS, 4, 1,8192*8192, 2048,2048,2048,2048, 8192,8192,8192,8192}, + // {VARYING_BOTH_DIMS, 4, 1,10240*10240, 2560,2560,2560,2560, 10240,10240,10240,10240}, + // {VARYING_BOTH_DIMS, 4, 1,12288*12288, 3072,3072,3072,3072, 12288,12288,12288,12288}, + // {VARYING_BOTH_DIMS, 4, 1,14336*14336, 3584,3584,3584,3584, 14336,14336,14336,14336}, + // {VARYING_BOTH_DIMS, 4, 1,16384*16384, 4096,4096,4096,4096, 16384,16384,16384,16384}, + + // {VARYING_BOTH_DIMS, 8, 1,1024*1024, 128,128,128,128,128,128,128,128, 1024,1024,1024,1024,1024,1024,1024,1024}, + // {VARYING_BOTH_DIMS, 8, 1,2048*2048, 256,256,256,256,256,256,256,256, 2048,2048,2048,2048,2048,2048,2048,2048}, + // {VARYING_BOTH_DIMS, 8, 1,3072*3072, 384,384,384,384,384,384,384,384, 3072,3072,3072,3072,3072,3072,3072,3072}, + // {VARYING_BOTH_DIMS, 8, 1,4096*4096, 512,512,512,512,512,512,512,512, 4096,4096,4096,4096,4096,4096,4096,4096}, + // {VARYING_BOTH_DIMS, 8, 1,5120*5120, 640,640,640,640,640,640,640,640, 5120,5120,5120,5120,5120,5120,5120,5120}, + // {VARYING_BOTH_DIMS, 8, 1,6144*6144, 768,768,768,768,768,768,768,768, 6144,6144,6144,6144,6144,6144,6144,6144}, + // {VARYING_BOTH_DIMS, 8, 1,7168*7168, 896,896,896,896,896,896,896,896, 7168,7168,7168,7168,7168,7168,7168,7168}, + // {VARYING_BOTH_DIMS, 8, 1,8192*8192, 1024,1024,1024,1024,1024,1024,1024,1024, 8192,8192,8192,8192,8192,8192,8192,8192}, + // {VARYING_BOTH_DIMS, 8, 1,10240*10240, 1280,1280,1280,1280,1280,1280,1280,1280, 10240,10240,10240,10240,10240,10240,10240,10240}, + // {VARYING_BOTH_DIMS, 8, 1,12288*12288, 1536,1536,1536,1536,1536,1536,1536,1536, 12288,12288,12288,12288,12288,12288,12288,12288}, + // {VARYING_BOTH_DIMS, 8, 1,14336*14336, 1792,1792,1792,1792,1792,1792,1792,1792, 14336,14336,14336,14336,14336,14336,14336,14336}, + // {VARYING_BOTH_DIMS, 8, 1,16384*16384, 2048,2048,2048,2048,2048,2048,2048,2048, 16384,16384,16384,16384,16384,16384,16384,16384}, + + // {VARYING_BOTH_DIMS, 16, 1,1024*1024, 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0, + // 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024}, + // {VARYING_BOTH_DIMS, 16, 1,2048*2048, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + // 2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048}, + // {VARYING_BOTH_DIMS, 16, 1,3072*3072, 256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128, + // 3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072}, + // {VARYING_BOTH_DIMS, 16, 1,4096*4096, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + // 4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096}, + // {VARYING_BOTH_DIMS, 16, 1,5120*5120, 512,512,512,512,256,256,256,256,256,256,256,256,256,256,256,256, + // 5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120}, + // {VARYING_BOTH_DIMS, 16, 1,6144*6144, 384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384, + // 6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144}, + // {VARYING_BOTH_DIMS, 16, 1,7168*7168, 512,512,512,512,512,512,512,512,512,512,512,512,256,256,256,256, + // 7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168}, + // {VARYING_BOTH_DIMS, 16, 1,8192*8192, 512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512, + // 8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192}, + // {VARYING_BOTH_DIMS, 16, 1,10240*10240, 640,640,640,640,640,640,640,640,640,640,640,640,640,640,640,640, + // 10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240}, + // {VARYING_BOTH_DIMS, 16, 1,12288*12288, 768,768,768,768,768,768,768,768,768,768,768,768,768,768,768,768, + // 12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288}, + // {VARYING_BOTH_DIMS, 16, 1,14336*14336, 896,896,896,896,896,896,896,896,896,896,896,896,896,896,896,896, + // 14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336}, + // {VARYING_BOTH_DIMS, 16, 1,16384*16384, 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024, + // 16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384}, + + // {VARYING_BOTH_DIMS, 32, 1,1024*1024, 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + // 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024, + // 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024}, + // {VARYING_BOTH_DIMS, 32, 1,2048*2048, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + // 2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048, + // 2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048}, + // {VARYING_BOTH_DIMS, 32, 1,3072*3072, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0, + // 3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072, + // 3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072}, + // {VARYING_BOTH_DIMS, 32, 1,4096*4096, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + // 4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096, + // 4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096}, + // {VARYING_BOTH_DIMS, 32, 1,5120*5120, 256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + // 5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120, + // 5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120}, + // {VARYING_BOTH_DIMS, 32, 1,6144*6144, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + // 6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144, + // 6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144}, + // {VARYING_BOTH_DIMS, 32, 1,7168*7168, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128, + // 7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168, + // 7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168}, + // {VARYING_BOTH_DIMS, 32, 1,8192*8192, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + // 8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192, + // 8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192}, + // {VARYING_BOTH_DIMS, 32, 1,10240*10240, 384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,256,256,256,256,256,256,128,128,128,128,128,128, + // 10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240, + // 10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240}, + // {VARYING_BOTH_DIMS, 32, 1,12288*12288, 384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384,384, + // 12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288, + // 12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288}, + // {VARYING_BOTH_DIMS, 32, 1,14336*14336, 512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,256,256,256,256,256,256,256,256, + // 14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336, + // 14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336}, + // {VARYING_BOTH_DIMS, 32, 1,16384*16384, 512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512,512, + // 16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384, + // 16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384}, + + // {VARYING_BOTH_DIMS, 64, 1,1024*1024, 128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + // 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024, + // 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024, + // 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024, + // 1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024,1024}, + // {VARYING_BOTH_DIMS, 64, 1,2048*2048, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + // 2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048, + // 2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048, + // 2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048, + // 2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048,2048}, + // {VARYING_BOTH_DIMS, 64, 1,3072*3072, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + // 3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072, + // 3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072, + // 3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072, + // 3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072,3072}, + // {VARYING_BOTH_DIMS, 64, 1,4096*4096, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + // 4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096, + // 4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096, + // 4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096, + // 4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096,4096}, + // {VARYING_BOTH_DIMS, 64, 1,5120*5120, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + // 5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120, + // 5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120, + // 5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120, + // 5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120,5120}, + // {VARYING_BOTH_DIMS, 64, 1,6144*6144, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0, + // 6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144, + // 6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144, + // 6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144, + // 6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144,6144}, + // {VARYING_BOTH_DIMS, 64, 1,7168*7168, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,0,0,0,0,0,0,0,0, + // 7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168, + // 7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168, + // 7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168, + // 7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168,7168}, + // {VARYING_BOTH_DIMS, 64, 1,8192*8192, 128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + // 8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192, + // 8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192, + // 8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192, + // 8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192,8192}, + // {VARYING_BOTH_DIMS, 64, 1,10240*10240, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + // 10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240, + // 10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240, + // 10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240, + // 10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240,10240}, + // {VARYING_BOTH_DIMS, 64, 1,12288*12288, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + // 12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288, + // 12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288, + // 12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288, + // 12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288,12288}, + // {VARYING_BOTH_DIMS, 64, 1,14336*14336, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128,128, + // 14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336, + // 14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336, + // 14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336, + // 14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336,14336}, + // {VARYING_BOTH_DIMS, 64, 1,16384*16384, 256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256,256, + // 16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384, + // 16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384, + // 16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384, + // 16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384,16384}, +}; + +} // namespace + +class GroupedFusedCastTransposeNVFP4TestSuite : public ::testing::TestWithParam + , // Config + transformer_engine::DType, + bool>> {}; + +TEST_P(GroupedFusedCastTransposeNVFP4TestSuite, TestFusedCastTransposeNVFP4) { + // Skip tests for pre-Blackwell architectures + if (getDeviceComputeCapability() < blackwellComputeCapability) { + GTEST_SKIP(); + } + + using namespace transformer_engine; + using namespace test; + + const std::vector input_config = std::get<0>(GetParam()); + const DType input_type = std::get<1>(GetParam()); + const bool use_fast_math = std::get<2>(GetParam()); + + const ShapeRepresentation shape_rep = static_cast(input_config[0]); + const size_t num_tensors = input_config[1]; + const std::vector logical_shape = {input_config[2], input_config[3]}; + + std::vector first_dims(num_tensors); + std::vector last_dims(num_tensors); + std::vector offsets(num_tensors + 1, 0); + for (size_t t = 0; t < num_tensors; ++t) { + switch (shape_rep) { + case SAME_BOTH_DIMS: { + first_dims[t] = logical_shape[0] / num_tensors; + last_dims[t] = logical_shape[1]; + break; + } + case VARYING_FIRST_DIM: { + first_dims[t] = input_config[t + 4]; + last_dims[t] = logical_shape[1]; + break; + } + case VARYING_LAST_DIM: { + first_dims[t] = logical_shape[0]; + last_dims[t] = input_config[t + 4]; + break; + } + case VARYING_BOTH_DIMS: { + first_dims[t] = input_config[t + 4]; + last_dims[t] = input_config[t + (4 + num_tensors)]; + break; + } + } + offsets[t + 1] = offsets[t] + first_dims[t] * last_dims[t]; + + if (first_dims[t] % 128 != 0) { + GTEST_SKIP(); + } + + if (shape_rep == VARYING_LAST_DIM || shape_rep == VARYING_BOTH_DIMS) { + if (last_dims[t] % 128 != 0) { + GTEST_SKIP(); + } + } + } + + TRANSFORMER_ENGINE_TYPE_SWITCH_FP16_FP32_ONLY(input_type, InputType, + performTest(shape_rep, num_tensors, logical_shape, + first_dims, last_dims, offsets, use_fast_math); + ); +} + +INSTANTIATE_TEST_SUITE_P( + OperatorTest, + GroupedFusedCastTransposeNVFP4TestSuite, + ::testing::Combine( + ::testing::ValuesIn(grouped_input_config), + ::testing::Values(DType::kBFloat16), + ::testing::Values(true)), + [](const testing::TestParamInfo& info) { + std::string name = "CAST_ONLY"; + const std::vector input = std::get<0>(info.param); + + switch (static_cast(input[0])) { + case ShapeRepresentation::SAME_BOTH_DIMS: name += "_SAME_BOTH_DIMS"; break; + case ShapeRepresentation::VARYING_FIRST_DIM: name += "_VARYING_FIRST_DIM"; break; + case ShapeRepresentation::VARYING_LAST_DIM: name += "_VARYING_LAST_DIM"; break; + case ShapeRepresentation::VARYING_BOTH_DIMS: name += "_VARYING_BOTH_DIMS"; break; + }; + + name += "_N_" + std::to_string(input[1]); + name += "_SHAPE_" + std::to_string(input[2]) + "X" + std::to_string(input[3]); + name += "_" + test::typeName(std::get<1>(info.param)); + if (std::get<2>(info.param)) { + name += "_FAST_SCALING"; + } + return name; + }); diff --git a/transformer_engine/common/cast/core/grouped_layout.cuh b/transformer_engine/common/cast/core/grouped_layout.cuh index 8337e44815..7eff1057cc 100644 --- a/transformer_engine/common/cast/core/grouped_layout.cuh +++ b/transformer_engine/common/cast/core/grouped_layout.cuh @@ -236,6 +236,17 @@ __device__ __forceinline__ size_t get_tensor_rows_num( return 0; } +__device__ __forceinline__ size_t +get_tensor_base_offset(const size_t tensor_id, const ShapeRepresentation shape_rep, + const size_t first_logical_dim, const size_t last_logical_dim, + const size_t num_tensors, const int64_t *const __restrict__ offsets_ptr) { + if (shape_rep == ShapeRepresentation::SAME_BOTH_DIMS) { + const size_t rows_per_tensor = first_logical_dim / num_tensors; + return tensor_id * rows_per_tensor * last_logical_dim; + } + return static_cast(offsets_ptr[tensor_id]); +} + template __device__ __forceinline__ size_t get_tensor_cols_num(const size_t tensor_id, const size_t last_logical_dim, @@ -319,6 +330,28 @@ struct BlockDescriptor { block_offset_X(block_offset_X_) {} }; +// Per-tensor metadata prepared together with grouped TMA descriptors. +struct TensorMetadata { + size_t rows = 0; + size_t cols = 0; + size_t tensor_base = 0; + size_t rowwise_scale_base = 0; + size_t colwise_scale_base = 0; + + __host__ __device__ __forceinline__ constexpr TensorMetadata() = default; + + __host__ __device__ __forceinline__ constexpr TensorMetadata(const size_t rows_, + const size_t cols_, + const size_t tensor_base_, + const size_t rowwise_scale_base_, + const size_t colwise_scale_base_) + : rows(rows_), + cols(cols_), + tensor_base(tensor_base_), + rowwise_scale_base(rowwise_scale_base_), + colwise_scale_base(colwise_scale_base_) {} +}; + template __device__ __forceinline__ JobDescriptor decode_job( const size_t num_tensors, const size_t first_logical_dim, const size_t last_logical_dim, @@ -340,6 +373,28 @@ __device__ __forceinline__ JobDescriptor decode_job( return JobDescriptor(block_id, block_global_offset, tensor_id, rows, cols); } +template +__device__ __forceinline__ JobDescriptor +decode_job(const size_t num_tensors, const size_t first_logical_dim, const size_t last_logical_dim, + const size_t work_blocks_X, const int32_t ctaid_X, const int32_t ctaid_Y, + const int64_t *const __restrict__ offsets_ptr, + const TensorMetadata *const __restrict__ metadata_ptr) { + constexpr size_t ELTS_PER_CHUNK = CHUNK_DIM_Y * CHUNK_DIM_X; + constexpr bool is_single_tensor = (SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS || + SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM); + const size_t ctaid_X_u = static_cast(ctaid_X); + const size_t ctaid_Y_u = static_cast(ctaid_Y); + const size_t block_id = ctaid_Y_u * work_blocks_X + ctaid_X_u; + const size_t block_global_offset = + is_single_tensor ? (ctaid_Y_u * CHUNK_DIM_Y * last_logical_dim + ctaid_X_u * CHUNK_DIM_X) + : (block_id * ELTS_PER_CHUNK); + const size_t tensor_id = get_current_tensor_id( + num_tensors, block_global_offset, ctaid_Y_u, first_logical_dim, last_logical_dim, + offsets_ptr); + const TensorMetadata metadata = metadata_ptr[tensor_id]; + return JobDescriptor(block_id, block_global_offset, tensor_id, metadata.rows, metadata.cols); +} + template __device__ __forceinline__ bool is_job_valid(const JobDescriptor &job, const size_t total_work_blocks, @@ -374,6 +429,13 @@ __device__ __forceinline__ bool job_has_work(const JobDescriptor &job) { return job.rows != 0 && job.cols != 0; } +template +__device__ __forceinline__ bool is_job_valid_with_work( + const JobDescriptor &job, const size_t total_work_blocks, + const int64_t *const __restrict__ offsets_ptr) { + return is_job_valid(job, total_work_blocks, offsets_ptr) && job_has_work(job); +} + __device__ __forceinline__ void advance_to_next_job(bool &job_finished, int32_t &ctaid_X, int32_t &ctaid_Y, size_t &static_next_block_id, const size_t static_block_stride, @@ -388,6 +450,13 @@ __device__ __forceinline__ void advance_to_next_job(bool &job_finished, int32_t } } +__device__ __forceinline__ void set_cta_coords_from_block_id(const size_t block_id, + const size_t work_blocks_X, + int32_t &ctaid_X, int32_t &ctaid_Y) { + ctaid_X = static_cast(block_id % work_blocks_X); + ctaid_Y = static_cast(block_id / work_blocks_X); +} + template __device__ __forceinline__ BlockDescriptor decode_block(const JobDescriptor &job, const int64_t *const __restrict__ offsets_ptr) { @@ -406,6 +475,37 @@ decode_block(const JobDescriptor &job, const int64_t *const __restrict__ offsets block_offset_Y, block_offset_X); } +template +__device__ __forceinline__ BlockDescriptor decode_block(const JobDescriptor &job, + const size_t tensor_base, + const int32_t ctaid_X, + const int32_t ctaid_Y) { + constexpr bool is_single_tensor = (SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS || + SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM); + constexpr size_t ELTS_PER_CHUNK = CHUNK_DIM_Y * CHUNK_DIM_X; + const size_t blocks_X_num_in_current_tensor = DIVUP(job.cols, CHUNK_DIM_X); + size_t block_id_in_current_tensor = 0; + size_t block_id_Y = 0; + size_t block_id_X = 0; + if constexpr (is_single_tensor) { + block_id_X = static_cast(ctaid_X); + if constexpr (SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS) { + const size_t blocks_Y_per_tensor = DIVUP(job.rows, static_cast(CHUNK_DIM_Y)); + block_id_Y = static_cast(ctaid_Y) - job.tensor_id * blocks_Y_per_tensor; + } else { + const size_t tensor_base_row = tensor_base / job.cols; + block_id_Y = static_cast(ctaid_Y) - tensor_base_row / CHUNK_DIM_Y; + } + block_id_in_current_tensor = block_id_Y * blocks_X_num_in_current_tensor + block_id_X; + } else { + block_id_in_current_tensor = job.block_id - tensor_base / ELTS_PER_CHUNK; + block_id_Y = block_id_in_current_tensor / blocks_X_num_in_current_tensor; + block_id_X = block_id_in_current_tensor % blocks_X_num_in_current_tensor; + } + return BlockDescriptor(tensor_base, block_id_in_current_tensor, block_id_Y, block_id_X, + block_id_Y * CHUNK_DIM_Y, block_id_X * CHUNK_DIM_X); +} + } // namespace common } // namespace dispatch } // namespace transformer_engine diff --git a/transformer_engine/common/cast/core/grouped_tma.cuh b/transformer_engine/common/cast/core/grouped_tma.cuh index 61218d654a..aa402c2b82 100644 --- a/transformer_engine/common/cast/core/grouped_tma.cuh +++ b/transformer_engine/common/cast/core/grouped_tma.cuh @@ -25,6 +25,7 @@ #include "../../common.h" #include "../../util/ptx.cuh" #include "../../utils.cuh" +#include "../nvfp4/core_nvfp4.cuh" #include "grouped_layout.cuh" namespace transformer_engine { @@ -38,13 +39,11 @@ struct alignas(128) TensorMapStorage { alignas(128) CUtensorMap act_input[MAX_SUPPORTED_TENSOR_DESCRIPTORS]; alignas(128) CUtensorMap output_rowwise[MAX_SUPPORTED_TENSOR_DESCRIPTORS]; alignas(128) CUtensorMap output_colwise[MAX_SUPPORTED_TENSOR_DESCRIPTORS]; - size_t rows[MAX_SUPPORTED_TENSOR_DESCRIPTORS]; - size_t cols[MAX_SUPPORTED_TENSOR_DESCRIPTORS]; - size_t offsets[MAX_SUPPORTED_TENSOR_DESCRIPTORS]; }; // Internal linkage avoids device-link ODR issues when this header is included by multiple .cu TUs. static __device__ TensorMapStorage g_tensor_maps; +alignas(128) static __device__ TensorMetadata g_tensor_metadata[MAX_SUPPORTED_TENSOR_DESCRIPTORS]; inline bool dimensions_supported_by_TMA(const Tensor *const t) { const size_t cols = t->flat_last_dim(); @@ -59,18 +58,32 @@ __device__ __forceinline__ unsigned char *align_smem_ptr_per_TMA_requirements(un return reinterpret_cast(addr); } +__device__ __forceinline__ uintptr_t get_pointer_with_offset_bits(const uintptr_t base_ptr, + const size_t offset_elts, + const size_t data_type_bits) { + const size_t offset_bits = offset_elts * data_type_bits; + if (offset_bits % 8 != 0) { + NVTE_DEVICE_ERROR("Data offset is not byte-aligned."); + } + return base_ptr + offset_bits / 8; +} + // Copies the base tensor map to shmem, modifies the copy, stores the modified tensor map at index -__device__ __forceinline__ void modify_base_tensor_map(const CUtensorMap base_tensor_map, - CUtensorMap *global_tensor_map, - const uintptr_t global_data_ptr, - const size_t global_dim_Y, - const size_t global_dim_X, - const size_t data_type_size_bytes) { +__device__ __forceinline__ void modify_base_tensor_map_bits(const CUtensorMap base_tensor_map, + CUtensorMap *global_tensor_map, + const uintptr_t global_data_ptr, + const size_t global_dim_Y, + const size_t global_dim_X, + const size_t data_type_bits) { __shared__ CUtensorMap shared_tensor_map; shared_tensor_map = base_tensor_map; // Copy the base tensor map into shmem constexpr bool is_blackwell = ARCH_BLACKWELL_FAMILY; if constexpr (is_blackwell) { - const size_t global_stride_bytes = global_dim_X * data_type_size_bytes; + const size_t global_stride_bits = global_dim_X * data_type_bits; + if (global_stride_bits % 8 != 0) { + NVTE_DEVICE_ERROR("Shape not supported. Data stride must be byte-aligned."); + } + const size_t global_stride_bytes = global_stride_bits / 8; if (global_stride_bytes % TMA_GMEM_ALIGNMENT != 0) { NVTE_DEVICE_ERROR("Shape not supported. Data stride must be 16B aligned."); } @@ -96,7 +109,17 @@ __device__ __forceinline__ void modify_base_tensor_map(const CUtensorMap base_te } } -template +__device__ __forceinline__ void modify_base_tensor_map(const CUtensorMap base_tensor_map, + CUtensorMap *global_tensor_map, + const uintptr_t global_data_ptr, + const size_t global_dim_Y, + const size_t global_dim_X, + const size_t data_type_size_bytes) { + modify_base_tensor_map_bits(base_tensor_map, global_tensor_map, global_data_ptr, global_dim_Y, + global_dim_X, data_type_size_bytes * 8); +} + +template __global__ void __launch_bounds__(1) update_tma_descriptors(const __grid_constant__ CUtensorMap base_tensor_map_input, const __grid_constant__ CUtensorMap base_tensor_map_act_input, @@ -113,50 +136,80 @@ __global__ void __launch_bounds__(1) const int64_t *const __restrict__ last_dims_ptr, const bool rowwise, const bool colwise, const bool compute_dactivations) { const size_t tensor_id = blockIdx.x; - const size_t rows = - get_tensor_rows_num(tensor_id, shape_rep, first_logical_dim, first_dims_ptr, num_tensors); - const size_t cols = get_tensor_cols_num(tensor_id, shape_rep, last_logical_dim, last_dims_ptr); + if (tensor_id >= num_tensors) { + return; + } - const size_t offset_elts = offsets_ptr[tensor_id]; - g_tensor_maps.rows[tensor_id] = rows; - g_tensor_maps.cols[tensor_id] = cols; - g_tensor_maps.offsets[tensor_id] = offset_elts; + const bool same_both_dims = shape_rep == ShapeRepresentation::SAME_BOTH_DIMS; + const size_t descriptor_first_logical_dim = + (NVFP4_CAST && same_both_dims) ? (first_logical_dim / num_tensors) : first_logical_dim; + const size_t descriptor_num_tensors = (NVFP4_CAST && same_both_dims) ? 1 : num_tensors; + const size_t rows = get_tensor_rows_num(tensor_id, shape_rep, descriptor_first_logical_dim, + first_dims_ptr, descriptor_num_tensors); + const size_t cols = get_tensor_cols_num(tensor_id, shape_rep, last_logical_dim, last_dims_ptr); + const size_t offset_elts = + NVFP4_CAST ? get_tensor_base_offset(tensor_id, shape_rep, first_logical_dim, last_logical_dim, + num_tensors, offsets_ptr) + : static_cast(offsets_ptr[tensor_id]); // Zero-sized groups: skip TMA descriptor update. The main kernel already returns // early for rows==0 or cols==0, but creating a TMA descriptor with a zero dimension // is invalid and causes CUDA_ERROR_ILLEGAL_ADDRESS. if (rows == 0 || cols == 0) { + if constexpr (NVFP4_CAST) { + g_tensor_metadata[tensor_id] = TensorMetadata(rows, cols, offset_elts, 0, 0); + } return; } - if (tensor_id < num_tensors) { - { - CUtensorMap *modified_tensor_map_input = &g_tensor_maps.input[tensor_id]; - const uintptr_t global_data_ptr = reinterpret_cast(input_data_ptr + offset_elts); - modify_base_tensor_map(base_tensor_map_input, modified_tensor_map_input, global_data_ptr, - rows, cols, sizeof(IType)); - } - if (compute_dactivations) { - CUtensorMap *modified_tensor_map_act_input = &g_tensor_maps.act_input[tensor_id]; - const uintptr_t global_data_ptr = - reinterpret_cast(act_input_data_ptr + offset_elts); - modify_base_tensor_map(base_tensor_map_act_input, modified_tensor_map_act_input, - global_data_ptr, rows, cols, sizeof(IType)); - } - if (rowwise) { - CUtensorMap *modified_tensor_map_output_rowwise = &g_tensor_maps.output_rowwise[tensor_id]; - const uintptr_t global_data_ptr = - reinterpret_cast(output_rowwise_data_ptr + offset_elts); - modify_base_tensor_map(base_tensor_map_output_rowwise, modified_tensor_map_output_rowwise, - global_data_ptr, rows, cols, sizeof(OType)); - } - if (colwise) { - CUtensorMap *modified_tensor_map_output_colwise = &g_tensor_maps.output_colwise[tensor_id]; - const uintptr_t global_data_ptr = - reinterpret_cast(output_colwise_data_ptr + offset_elts); - modify_base_tensor_map(base_tensor_map_output_colwise, modified_tensor_map_output_colwise, - global_data_ptr, rows, cols, sizeof(OType)); + if constexpr (NVFP4_CAST) { + const size_t rowwise_scale_base = + nvfp4::core::get_rowwise_scale_base(shape_rep, tensor_id, offset_elts, rows, cols); + const size_t colwise_scale_base = + nvfp4::core::get_colwise_scale_base(shape_rep, tensor_id, offset_elts, rows, cols); + g_tensor_metadata[tensor_id] = + TensorMetadata(rows, cols, offset_elts, rowwise_scale_base, colwise_scale_base); + } + + constexpr size_t input_type_bits = TypeInfo::size; + constexpr size_t output_type_bits = []() constexpr { + if constexpr (NVFP4_CAST) { + return size_t{4}; + } else { + return TypeInfo::size; } + }(); + const size_t output_colwise_dim_Y = NVFP4_CAST ? cols : rows; + const size_t output_colwise_dim_X = NVFP4_CAST ? rows : cols; + + { + CUtensorMap *modified_tensor_map_input = &g_tensor_maps.input[tensor_id]; + const uintptr_t global_data_ptr = get_pointer_with_offset_bits( + reinterpret_cast(input_data_ptr), offset_elts, input_type_bits); + modify_base_tensor_map_bits(base_tensor_map_input, modified_tensor_map_input, global_data_ptr, + rows, cols, input_type_bits); + } + if (compute_dactivations) { + CUtensorMap *modified_tensor_map_act_input = &g_tensor_maps.act_input[tensor_id]; + const uintptr_t global_data_ptr = get_pointer_with_offset_bits( + reinterpret_cast(act_input_data_ptr), offset_elts, input_type_bits); + modify_base_tensor_map_bits(base_tensor_map_act_input, modified_tensor_map_act_input, + global_data_ptr, rows, cols, input_type_bits); + } + if (rowwise) { + CUtensorMap *modified_tensor_map_output_rowwise = &g_tensor_maps.output_rowwise[tensor_id]; + const uintptr_t global_data_ptr = get_pointer_with_offset_bits( + reinterpret_cast(output_rowwise_data_ptr), offset_elts, output_type_bits); + modify_base_tensor_map_bits(base_tensor_map_output_rowwise, modified_tensor_map_output_rowwise, + global_data_ptr, rows, cols, output_type_bits); + } + if (colwise) { + CUtensorMap *modified_tensor_map_output_colwise = &g_tensor_maps.output_colwise[tensor_id]; + const uintptr_t global_data_ptr = get_pointer_with_offset_bits( + reinterpret_cast(output_colwise_data_ptr), offset_elts, output_type_bits); + modify_base_tensor_map_bits(base_tensor_map_output_colwise, modified_tensor_map_output_colwise, + global_data_ptr, output_colwise_dim_Y, output_colwise_dim_X, + output_type_bits); } } diff --git a/transformer_engine/common/cast/dispatch/quantize.cuh b/transformer_engine/common/cast/dispatch/quantize.cuh index 0dd2aebce8..5ad0893299 100644 --- a/transformer_engine/common/cast/dispatch/quantize.cuh +++ b/transformer_engine/common/cast/dispatch/quantize.cuh @@ -26,6 +26,7 @@ #include "../nvfp4/group_quantize_transpose_nvfp4.cuh" #include "../nvfp4/quantize_4over6_nvfp4.cuh" #include "../nvfp4/quantize_transpose_nvfp4.cuh" +#include "../nvfp4/specialized/group_quantize_transpose_nvfp4_tuned_1D.cuh" namespace transformer_engine { namespace dispatch { @@ -502,6 +503,19 @@ void group_quantize_fwd_helper(const NVTEGroupedTensor input, NVTEGroupedTensor workspace_tensor, &quant_config_cpp, stream); break; } + case NVTE_NVFP4_1D_SCALING: { + NVTE_CHECK(!IS_ACT, "IS_ACT is not supported by FWD NVTE_NVFP4_1D_SCALING"); + + const bool is_bf16_input_type = input_tensor->dtype() == DType::kBFloat16; + NVTE_CHECK(is_bf16_input_type, "Optimized grouped NVFP4 kernel supports only BF16 input."); + + const bool is_2D_quantization = quant_config_cpp.nvfp4_2d_quantization; + NVTE_CHECK(!is_2D_quantization, "2D quantization is not supported for group quantize."); + + nvfp4::group_quantize_transpose(input_tensor, noop_tensor, output_tensor, &quant_config_cpp, + stream); + break; + } case NVTE_BLOCK_SCALING_1D: { NVTE_CHECK(!IS_ACT, "IS_ACT is not implemented for grouped NVTE_BLOCK_SCALING_1D."); fp8_blockwise::group_quantize_blockwise_1d(input_tensor, output_tensor, noop_tensor, @@ -597,6 +611,19 @@ void group_quantize_bwd_helper(const NVTEGroupedTensor grad, const NVTEGroupedTe &quant_config_cpp, stream); break; } + case NVTE_NVFP4_1D_SCALING: { + NVTE_CHECK((!IS_DBIAS && !IS_DACT), + "IS_DBIAS and IS_DACT are not supported by BWD NVTE_NVFP4_1D_SCALING"); + const bool is_bf16_input_type = grad_tensor->dtype() == DType::kBFloat16; + NVTE_CHECK(is_bf16_input_type, "Optimized grouped NVFP4 kernel supports only BF16 input."); + + const bool is_2D_quantization = quant_config_cpp.nvfp4_2d_quantization; + NVTE_CHECK(!is_2D_quantization, "2D quantization is not supported for group quantize."); + + nvfp4::group_quantize_transpose(grad_tensor, noop_tensor, output_tensor, &quant_config_cpp, + stream); + break; + } case NVTE_BLOCK_SCALING_1D: case NVTE_BLOCK_SCALING_2D: { NVTE_CHECK(!IS_DACT, "IS_DACT is not implemented for grouped FP8 block scaling."); diff --git a/transformer_engine/common/cast/nvfp4/core_nvfp4.cuh b/transformer_engine/common/cast/nvfp4/core_nvfp4.cuh index 3820430d5b..1fa23f9e0c 100644 --- a/transformer_engine/common/cast/nvfp4/core_nvfp4.cuh +++ b/transformer_engine/common/cast/nvfp4/core_nvfp4.cuh @@ -32,6 +32,7 @@ namespace dispatch { namespace nvfp4 { using nvfp4_scale_t = fp8e4m3; +constexpr int NVFP4_SCALE_DIM = 16; // NVFP4 block (x16 elts) namespace quantization_and_transposition_SF { #if FP4_TYPE_SUPPORTED @@ -72,6 +73,44 @@ __device__ __forceinline__ fp8e4m3 compute_decoding_scaling_factor(const float b namespace core { +__device__ __forceinline__ size_t get_nvfp4_scale_stride(const size_t block_scaled_dim) { + return DIVUP_TO_MULTIPLE(DIVUP(block_scaled_dim, static_cast(NVFP4_SCALE_DIM)), 4); +} + +// Scale buffers are compact per-tensor concatenations. Same-dim paths keep the padded +// scale stride explicit; fully varying supported shapes are 128-aligned, so base = elts / 16. +__device__ __forceinline__ size_t get_rowwise_scale_base(const ShapeRepresentation shape_rep, + const size_t tensor_id, + const size_t tensor_base, + const size_t rows, const size_t cols) { + switch (shape_rep) { + case ShapeRepresentation::SAME_BOTH_DIMS: + return tensor_id * rows * get_nvfp4_scale_stride(cols); + case ShapeRepresentation::VARYING_FIRST_DIM: + return (tensor_base / cols) * get_nvfp4_scale_stride(cols); + case ShapeRepresentation::VARYING_LAST_DIM: + case ShapeRepresentation::VARYING_BOTH_DIMS: + return tensor_base / static_cast(NVFP4_SCALE_DIM); + } + return 0; +} + +__device__ __forceinline__ size_t get_colwise_scale_base(const ShapeRepresentation shape_rep, + const size_t tensor_id, + const size_t tensor_base, + const size_t rows, const size_t cols) { + switch (shape_rep) { + case ShapeRepresentation::SAME_BOTH_DIMS: + return tensor_id * cols * get_nvfp4_scale_stride(rows); + case ShapeRepresentation::VARYING_LAST_DIM: + return (tensor_base / rows) * get_nvfp4_scale_stride(rows); + case ShapeRepresentation::VARYING_FIRST_DIM: + case ShapeRepresentation::VARYING_BOTH_DIMS: + return tensor_base / static_cast(NVFP4_SCALE_DIM); + } + return 0; +} + #if FP4_TYPE_SUPPORTED using namespace ptx; @@ -94,6 +133,30 @@ __device__ __forceinline__ float compute_global_encode_scaling_factor_FP4(const return global_encode_scale; } +// Compute "correct" per-block encoding scaling factor +template +__device__ __forceinline__ SF_TYPE compute_scaling_coefficient(const nvfp4_scale_t S_dec_block, + const float S_enc) { + NVTE_DEVICE_ERROR("Unsupported scaling-factor type. Only FP32 and BF16 are supported."); +} + +template <> +__device__ __forceinline__ float compute_scaling_coefficient(const nvfp4_scale_t S_dec_block, + const float S_enc) { + const float S_dec = 1.0f / S_enc; + const float scale_rcp = + fminf(1.0f / (static_cast(S_dec_block) * S_dec), detail::TypeExtrema::max); + return scale_rcp; +} + +template <> +__device__ __forceinline__ bf16 compute_scaling_coefficient(const nvfp4_scale_t S_dec_block, + const float S_enc) { + const float scale_rcp = + fminf(S_enc / (static_cast(S_dec_block)), detail::TypeExtrema::max); + return static_cast(scale_rcp); +} + __device__ __forceinline__ uint32_t get_rbits( transformer_engine::curanddx::detail::philox4x32_native_state &rng, diff --git a/transformer_engine/common/cast/nvfp4/group_quantize_transpose_nvfp4.cuh b/transformer_engine/common/cast/nvfp4/group_quantize_transpose_nvfp4.cuh index 3c6d9585e4..ebbd9a694b 100644 --- a/transformer_engine/common/cast/nvfp4/group_quantize_transpose_nvfp4.cuh +++ b/transformer_engine/common/cast/nvfp4/group_quantize_transpose_nvfp4.cuh @@ -331,6 +331,17 @@ __global__ void __launch_bounds__(THREADS_NUM) const size_t buff_offset_out = buff * BUFF_OUT_SIZE; const size_t buff_offset_out_t = buff * BUFF_OUT_T_SIZE; + // Wait for TMA transfer to have finished reading shared memory. + // I.e. the buffer is ready to be written to + if (stage >= BUFFS_NUM) { + if (is_master_thread) { + ptx::cp_async_bulk_wait_group_read(); + } + // Bulk async-groups are thread-local. Publish the master's completion before any thread + // overwrites this reused output buffer. + __syncthreads(); + } + // for stages from 1 to STAGES - 1, we need to update the tensor id // skip updating tensor id if it's the last CTA, and some stages will be out of bounds if (need_update_tensor_id && stage > 0 && (block_offset_Y + stage_offset_Y < rows)) { @@ -349,10 +360,6 @@ __global__ void __launch_bounds__(THREADS_NUM) } if (next_stage < STAGES) { - // Wait for TMA transfer to have finished reading shared memory. - // I.e. the buffer is ready to be written to - ptx::cp_async_bulk_wait_group_read<1>(); - const size_t next_buff = next_stage % BUFFS_NUM; const size_t next_stage_offset_Y = next_stage * BUFF_DIM_Y; const size_t global_offset_Y = block_offset_Y + next_stage_offset_Y; @@ -726,6 +733,11 @@ __global__ void __launch_bounds__(THREADS_NUM) // } // } + if (is_master_thread) { + ptx::cp_async_bulk_wait_group(); + } + __syncthreads(); + destroy_barriers(mbar, is_master_thread); #else NVTE_DEVICE_ERROR("sm_100 or higher is required."); diff --git a/transformer_engine/common/cast/nvfp4/quantize_transpose_nvfp4.cuh b/transformer_engine/common/cast/nvfp4/quantize_transpose_nvfp4.cuh index 2734cf3fef..12d31b249f 100644 --- a/transformer_engine/common/cast/nvfp4/quantize_transpose_nvfp4.cuh +++ b/transformer_engine/common/cast/nvfp4/quantize_transpose_nvfp4.cuh @@ -20,7 +20,7 @@ #include "../../common.h" #include "../../util/math.h" -#include "../../util/ptx_arch_spec.cuh" +#include "../../util/ptx.cuh" #include "../../utils.cuh" #include "core_nvfp4.cuh" #include "specialized/quantize_transpose_nvfp4_tuned_1D.cuh" @@ -236,16 +236,15 @@ using namespace ptx; #if FP4_TYPE_SUPPORTED -constexpr size_t SCALE_DIM = 16; // NVFP4 block (x16 elts) - constexpr size_t CHUNK_DIM_Y = 128; constexpr size_t CHUNK_DIM_X = 128; constexpr size_t THREADS_NUM = 128; -constexpr size_t SCALES_PER_CHUNK_Y = CHUNK_DIM_Y / SCALE_DIM; -constexpr size_t SCALES_PER_CHUNK_X = CHUNK_DIM_X / SCALE_DIM; +constexpr size_t SCALES_PER_CHUNK_Y = CHUNK_DIM_Y / NVFP4_SCALE_DIM; +constexpr size_t SCALES_PER_CHUNK_X = CHUNK_DIM_X / NVFP4_SCALE_DIM; -constexpr size_t SCALES_PER_THREAD = 2 * (CHUNK_DIM_Y * CHUNK_DIM_X) / SCALE_DIM / THREADS_NUM; +constexpr size_t SCALES_PER_THREAD = + 2 * (CHUNK_DIM_Y * CHUNK_DIM_X) / NVFP4_SCALE_DIM / THREADS_NUM; // Each call generates 4x uint32_t random numbers constexpr size_t RNG_GENS_PER_THREAD = SCALES_PER_THREAD / 4; @@ -253,9 +252,9 @@ constexpr size_t RNG_GENS_PER_THREAD = SCALES_PER_THREAD / 4; constexpr size_t TILE_DIM_Y = 32; constexpr size_t TILE_DIM_X = 128; -// SHould this be SCALE_DIM or BLOCK_DIM? Both are 16, should work for both 1D and 2D -constexpr size_t SCALES_PER_TILE_Y = TILE_DIM_Y / SCALE_DIM; -constexpr size_t SCALES_PER_TILE_X = TILE_DIM_X / SCALE_DIM; // 128 / 16 = 8 +// SHould this be NVFP4_SCALE_DIM or BLOCK_DIM? Both are 16, should work for both 1D and 2D +constexpr size_t SCALES_PER_TILE_Y = TILE_DIM_Y / NVFP4_SCALE_DIM; +constexpr size_t SCALES_PER_TILE_X = TILE_DIM_X / NVFP4_SCALE_DIM; // 128 / 16 = 8 constexpr size_t TILES_Y = CHUNK_DIM_Y / TILE_DIM_Y; constexpr size_t TILES_X = CHUNK_DIM_X / TILE_DIM_X; @@ -284,17 +283,17 @@ constexpr size_t BUFF_OUT_T_SIZE = BUFF_OUT_T_DIM_Y * BUFF_OUT_T_DIM_X; // Manual swizzling parameters to reduce SHMEM bank conflicts constexpr size_t PACK_SIZE = 8; -constexpr size_t WAVES = SCALE_DIM / PACK_SIZE; +constexpr size_t WAVES = NVFP4_SCALE_DIM / PACK_SIZE; -constexpr size_t SCALING_FACTORS_PER_TILE_X = TILE_DIM_X / SCALE_DIM; +constexpr size_t SCALING_FACTORS_PER_TILE_X = TILE_DIM_X / NVFP4_SCALE_DIM; constexpr size_t THREADS_X_ROWWISE = SCALING_FACTORS_PER_TILE_X; // 128 / 16 = 8 constexpr size_t THREADS_Y_ROWWISE = THREADS_NUM / THREADS_X_ROWWISE; // 128 / 8 = 16 constexpr size_t ITERATIONS_NORMAL = BUFF_DIM_Y / THREADS_Y_ROWWISE; // 32/ 16 = 2 -constexpr size_t ITERATIONS_TRANSPOSE = BUFF_IN_DIM_Y / SCALE_DIM; +constexpr size_t ITERATIONS_TRANSPOSE = BUFF_IN_DIM_Y / NVFP4_SCALE_DIM; constexpr size_t BUFF_OUT_IT_OFFSET = BUFF_OUT_T_DIM_X / ITERATIONS_TRANSPOSE; -static_assert(BUFF_DIM_Y >= SCALE_DIM && +static_assert(BUFF_DIM_Y >= NVFP4_SCALE_DIM && "Number of buffer rows must be greater or equal to the size of the columwise " "scaling block\0"); static_assert(CHUNK_DIM_Y >= BUFF_DIM_Y); @@ -306,7 +305,7 @@ static_assert(BUFF_DIM_Y >= THREADS_Y_ROWWISE && constexpr size_t TOTAL_BANKS_WIDTH = (32 * 4 * 8) / 4; // 256 // Number of threads (rowwise scaling) that span 32 banks (4-byte banks) of shared memory -constexpr size_t THREADS_PER_BANK = TOTAL_BANKS_WIDTH / SCALE_DIM; // 8 = 128 / 16 +constexpr size_t THREADS_PER_BANK = TOTAL_BANKS_WIDTH / NVFP4_SCALE_DIM; // 8 = 128 / 16 template (0.0f); #pragma unroll - for (int i = 0; i < SCALE_DIM; ++i) { + for (int i = 0; i < NVFP4_SCALE_DIM; ++i) { const int shmem_offset_colwise = shmem_offset_base_colwise_in + i * BUFF_IN_DIM_X; in_colwise_IType[i] = in_sh[shmem_offset_colwise]; block_amax_f16 = __hmax(block_amax_f16, __habs(in_colwise_IType[i])); @@ -507,7 +506,7 @@ __global__ void __launch_bounds__(THREADS_NUM) block_amax = static_cast(block_amax_f16); } else { #pragma unroll - for (int i = 0; i < SCALE_DIM; ++i) { + for (int i = 0; i < NVFP4_SCALE_DIM; ++i) { const int shmem_offset_colwise = shmem_offset_base_colwise_in + i * BUFF_IN_DIM_X; float elt = static_cast(in_sh[shmem_offset_colwise]); if constexpr (COMPUTE_ACTIVATIONS) { @@ -551,10 +550,10 @@ __global__ void __launch_bounds__(THREADS_NUM) const float2 block_scale_inverse_2x{block_scale_inverse, block_scale_inverse}; // 3. Scale elements - fp4e2m1x4 regs[SCALE_DIM / 4]; + fp4e2m1x4 regs[NVFP4_SCALE_DIM / 4]; #pragma unroll - for (int e = 0; e < SCALE_DIM / 4; ++e) { + for (int e = 0; e < NVFP4_SCALE_DIM / 4; ++e) { const uint32_t rbits = get_rbits(rng, random_uint4, rnd_idx); if constexpr (NO_ACTIVATIONS_NOT_FP32_INPUT) { const uint64_t elts = *reinterpret_cast(&in_colwise_IType[4 * e]); @@ -606,7 +605,7 @@ __global__ void __launch_bounds__(THREADS_NUM) const size_t it_offset_Y = stage_offset_Y + it * THREADS_Y_ROWWISE; block_amax = 0.0f; - float in_compute_rowwise[SCALE_DIM]; + float in_compute_rowwise[NVFP4_SCALE_DIM]; Vec in_cached[WAVES]; // used as an IType container for BF16/FP16 --> NVFP4 CAST ONLY @@ -617,7 +616,7 @@ __global__ void __launch_bounds__(THREADS_NUM) IType2 thread_amax_2x = {static_cast(0.0f), static_cast(0.0f)}; #pragma unroll for (int w = 0; w < WAVES; ++w) { - const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % SCALE_DIM; + const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % NVFP4_SCALE_DIM; const size_t swizzled_thread_idx = thread_offset_X_rowwise + swizzled_group_idx; const size_t shmem_offset_rowwise = shmem_offset_base_rowwise_in + swizzled_thread_idx; // Load elements @@ -635,7 +634,7 @@ __global__ void __launch_bounds__(THREADS_NUM) IType2 thread_amax_2x = {static_cast(0.0f), static_cast(0.0f)}; #pragma unroll for (int w = 0; w < WAVES; ++w) { - const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % SCALE_DIM; + const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % NVFP4_SCALE_DIM; const size_t swizzled_thread_idx = thread_offset_X_rowwise + swizzled_group_idx; const size_t shmem_offset_rowwise = shmem_offset_base_rowwise_in + swizzled_thread_idx; @@ -670,7 +669,7 @@ __global__ void __launch_bounds__(THREADS_NUM) } else { #pragma unroll for (int w = 0; w < WAVES; ++w) { - const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % SCALE_DIM; + const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % NVFP4_SCALE_DIM; const size_t swizzled_thread_idx = thread_offset_X_rowwise + swizzled_group_idx; const size_t shmem_offset_rowwise = shmem_offset_base_rowwise_in + swizzled_thread_idx; @@ -785,7 +784,7 @@ __global__ void __launch_bounds__(THREADS_NUM) in01, in23, block_scale_inverse_2x, rbits); } } - const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % SCALE_DIM; + const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % NVFP4_SCALE_DIM; const size_t swizzled_idx = swizzled_group_idx + thread_offset_X_rowwise; const size_t shmem_offset_rowwise = shmem_offset_base_rowwise_out + swizzled_idx / 2; out.store_to(&out_data_sh[shmem_offset_rowwise]); @@ -831,7 +830,7 @@ __global__ void __launch_bounds__(THREADS_NUM) ScalesVec &scales_vec = *reinterpret_cast(&out_colwise_scales_sh[scale_idx_sh]); const size_t scale_idx_global = scales_offset_Y_t * scale_stride_t + scales_offset_X_t; const size_t count = // number of scales in Y dimension of this chunk - (chunk_rows >= CHUNK_DIM_Y) ? SCALES_PER_CHUNK_Y : (chunk_rows / SCALE_DIM); + (chunk_rows >= CHUNK_DIM_Y) ? SCALES_PER_CHUNK_Y : (chunk_rows / NVFP4_SCALE_DIM); nvfp4_scale_t *dst = &scales_t_ptr[scale_idx_global]; constexpr size_t vec_bytes = SCALES_PER_CHUNK_Y * sizeof(nvfp4_scale_t); if (count == SCALES_PER_CHUNK_Y && (reinterpret_cast(dst) % vec_bytes == 0)) { @@ -911,7 +910,7 @@ __global__ void __launch_bounds__(THREADS_NUM) const size_t tid_Y_t = tid_X_colwise; const size_t thread_offset_Y_rowwise = tid_Y_rowwise; - const size_t thread_offset_X_rowwise = tid_X_rowwise * SCALE_DIM; + const size_t thread_offset_X_rowwise = tid_X_rowwise * NVFP4_SCALE_DIM; const size_t thread_offset_X_colwise = tid_X_colwise; const size_t scales_offset_Y_rowwise = scales_block_offset_Y_rowwise + tid_Y_rowwise; @@ -919,7 +918,7 @@ __global__ void __launch_bounds__(THREADS_NUM) const size_t scales_offset_Y_t = scales_block_offset_Y_t + tid_Y_t; const size_t scales_offset_X_t = scales_block_offset_X_t; - const size_t SFs_per_row = cols / SCALE_DIM; + const size_t SFs_per_row = cols / NVFP4_SCALE_DIM; const bool rowwise_scale_is_within_bounds_X = scales_offset_X_rowwise < SFs_per_row; const bool colwise_scale_is_within_bounds_Y = scales_offset_Y_t < cols; @@ -1101,7 +1100,7 @@ __global__ void __launch_bounds__(THREADS_NUM) const size_t block_in_tile_y = it; const size_t block_in_tile_x = threadIdx.x / BLOCK_DIM; - const size_t in_thread_offset_Y = 0 + it * SCALE_DIM; + const size_t in_thread_offset_Y = 0 + it * NVFP4_SCALE_DIM; const size_t in_thread_offset_X = thread_offset_X_colwise; const size_t out_t_thread_offset_Y = thread_offset_X_colwise; @@ -1113,19 +1112,19 @@ __global__ void __launch_bounds__(THREADS_NUM) buff_offset_out_t + out_t_thread_offset_Y * BUFF_OUT_T_DIM_X + out_t_thread_offset_X; block_amax = block_amax_matrix[block_in_tile_y][block_in_tile_x]; - float in_compute_colwise[SCALE_DIM]; - IType in_colwise_IType[SCALE_DIM]; + float in_compute_colwise[NVFP4_SCALE_DIM]; + IType in_colwise_IType[NVFP4_SCALE_DIM]; // 3. Scale elements // Load data in if constexpr (NO_ACTIVATIONS_NOT_FP32_INPUT) { #pragma unroll - for (int i = 0; i < SCALE_DIM; ++i) { + for (int i = 0; i < NVFP4_SCALE_DIM; ++i) { const int shmem_offset_colwise = shmem_offset_base_colwise_in + i * BUFF_IN_DIM_X; in_colwise_IType[i] = in_sh[shmem_offset_colwise]; } } else { - for (int i = 0; i < SCALE_DIM; ++i) { + for (int i = 0; i < NVFP4_SCALE_DIM; ++i) { const int shmem_offset_colwise = shmem_offset_base_colwise_in + i * BUFF_IN_DIM_X; float elt = static_cast(in_sh[shmem_offset_colwise]); if constexpr (COMPUTE_ACTIVATIONS) { @@ -1159,9 +1158,9 @@ __global__ void __launch_bounds__(THREADS_NUM) 1.0f / (static_cast(S_dec_b_fp8) * S_dec_colwise), float_max); // S_enc_b_fp8 const float2 block_scale_inverse_2x{block_scale_inverse, block_scale_inverse}; - fp4e2m1x4 regs[SCALE_DIM / 4]; + fp4e2m1x4 regs[NVFP4_SCALE_DIM / 4]; #pragma unroll - for (int e = 0; e < SCALE_DIM / 4; ++e) { + for (int e = 0; e < NVFP4_SCALE_DIM / 4; ++e) { const uint32_t rbits = get_rbits(rng, random_uint4, rnd_idx); if constexpr (NO_ACTIVATIONS_NOT_FP32_INPUT) { const uint64_t elts = *reinterpret_cast(&in_colwise_IType[4 * e]); @@ -1212,7 +1211,7 @@ __global__ void __launch_bounds__(THREADS_NUM) buff_offset_out + it_thread_offset_Y_rowwise * BUFF_OUT_DIM_X; block_amax = block_amax_matrix[block_in_tile_y][block_in_tile_x]; - float in_compute_rowwise[SCALE_DIM]; + float in_compute_rowwise[NVFP4_SCALE_DIM]; Vec in_cached[WAVES]; // used as an IType container for BF16/FP16 --> NVFP4 CAST ONLY @@ -1223,7 +1222,7 @@ __global__ void __launch_bounds__(THREADS_NUM) IType2 thread_amax_2x = {static_cast(0.0f), static_cast(0.0f)}; #pragma unroll for (int w = 0; w < WAVES; ++w) { - const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % SCALE_DIM; + const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % NVFP4_SCALE_DIM; const size_t swizzled_thread_idx = thread_offset_X_rowwise + swizzled_group_idx; const size_t shmem_offset_rowwise = shmem_offset_base_rowwise_in + swizzled_thread_idx; // Load elements @@ -1234,7 +1233,7 @@ __global__ void __launch_bounds__(THREADS_NUM) __syncthreads(); #pragma unroll for (int w = 0; w < WAVES; ++w) { - const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % SCALE_DIM; + const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % NVFP4_SCALE_DIM; const size_t swizzled_thread_idx = thread_offset_X_rowwise + swizzled_group_idx; const size_t shmem_offset_rowwise = shmem_offset_base_rowwise_in + swizzled_thread_idx; @@ -1244,7 +1243,7 @@ __global__ void __launch_bounds__(THREADS_NUM) } else { #pragma unroll for (int w = 0; w < WAVES; ++w) { - const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % SCALE_DIM; + const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % NVFP4_SCALE_DIM; const size_t swizzled_thread_idx = thread_offset_X_rowwise + swizzled_group_idx; const size_t shmem_offset_rowwise = shmem_offset_base_rowwise_in + swizzled_thread_idx; @@ -1280,7 +1279,7 @@ __global__ void __launch_bounds__(THREADS_NUM) size_t scale_idx_global; if constexpr (WITH_GEMM_SWIZZLED_SCALES) { // Write the scale directly into the cuBLAS GEMM-swizzled layout so no - // separate swizzle pass is needed. SFs_per_row (= cols / SCALE_DIM) is + // separate swizzle pass is needed. SFs_per_row (= cols / NVFP4_SCALE_DIM) is // the number of compact scale columns. scale_idx_global = swizzle::gemm_swizzled_scale_idx(scales_offset_Y, scales_offset_X, SFs_per_row); @@ -1326,7 +1325,7 @@ __global__ void __launch_bounds__(THREADS_NUM) } } - const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % SCALE_DIM; + const size_t swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % NVFP4_SCALE_DIM; const size_t swizzled_idx = swizzled_group_idx + thread_offset_X_rowwise; const size_t shmem_offset_rowwise = shmem_offset_base_rowwise_out + swizzled_idx / 2; out.store_to(&out_data_sh[shmem_offset_rowwise]); @@ -1371,15 +1370,15 @@ __global__ void __launch_bounds__(THREADS_NUM) if (RETURN_TRANSPOSE && colwise_scale_is_within_bounds_Y) { const size_t scale_idx_sh = tid_Y_t * SCALES_PER_CHUNK_Y; const size_t count = // number of scales in Y dimension of this chunk - (chunk_rows >= CHUNK_DIM_Y) ? SCALES_PER_CHUNK_Y : (chunk_rows / SCALE_DIM); + (chunk_rows >= CHUNK_DIM_Y) ? SCALES_PER_CHUNK_Y : (chunk_rows / NVFP4_SCALE_DIM); if constexpr (WITH_GEMM_SWIZZLED_SCALES) { // The swizzled layout scatters the contiguous columnwise scales, so the // vectorized store cannot be used. Emit each scale at its swizzled offset. - // The transposed scale matrix has `rows / SCALE_DIM` (= M/16) column tiles + // The transposed scale matrix has `rows / NVFP4_SCALE_DIM` (= M/16) column tiles // (exact because the swizzled path requires 128-aligned dims). Read - // SCALE_DIM by value; passing the namespace-scope constexpr by reference + // NVFP4_SCALE_DIM by value; passing the namespace-scope constexpr by reference // (e.g. via DIVUP) would ODR-use it and fail to compile in device code. - const size_t col_length_t = rows / SCALE_DIM; + const size_t col_length_t = rows / NVFP4_SCALE_DIM; for (size_t k = 0; k < count; ++k) { const size_t off = swizzle::gemm_swizzled_scale_idx(scales_offset_Y_t, scales_offset_X_t + k, col_length_t); diff --git a/transformer_engine/common/cast/nvfp4/specialized/group_quantize_transpose_nvfp4_tuned_1D.cuh b/transformer_engine/common/cast/nvfp4/specialized/group_quantize_transpose_nvfp4_tuned_1D.cuh new file mode 100644 index 0000000000..c4c8db73e7 --- /dev/null +++ b/transformer_engine/common/cast/nvfp4/specialized/group_quantize_transpose_nvfp4_tuned_1D.cuh @@ -0,0 +1,845 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +/*! \file group_quantize_transpose_nvfp4_tuned_1D.cuh + * \brief Tuned grouped kernel to cast to NVFP4 and transpose. + */ + +#ifndef TRANSFORMER_ENGINE_GROUP_QUANTIZE_TRANSPOSE_NVFP4_TUNED_1D_CUH_ +#define TRANSFORMER_ENGINE_GROUP_QUANTIZE_TRANSPOSE_NVFP4_TUNED_1D_CUH_ + +#include +#include +#include +#include + +#include "../../../common.h" +#include "../../../util/cuda_runtime.h" +#include "../../../util/math.h" +#include "../../../util/ptx_arch_spec.cuh" +#include "../../../utils.cuh" +#include "../../core/common.cuh" +#include "../core_nvfp4.cuh" +#include "scaling_nvfp4_tuned_1D.cuh" + +namespace transformer_engine { +namespace dispatch { +namespace nvfp4 { + +namespace group_quantize_transpose_tuned_kernel { + +using namespace quantization_and_transposition_SF; +using namespace core; +using namespace ptx; +using namespace dispatch::common; + +#if FP4_TYPE_SUPPORTED + +using tuned_1D_scaling_common::colwise_scaling; +using tuned_1D_scaling_common::rowwise_scaling; + +struct DefaultCastConfig : tuned_1D_scaling_common::DefaultGroupedScalingConfig { + static constexpr int STATIC_PERSISTENT_BLOCKS_PER_SM = 128; +}; + +template +struct CastConfig; + +// Keep every layout independently specializable while sharing the current defaults. +template <> +struct CastConfig : DefaultCastConfig {}; + +template <> +struct CastConfig : DefaultCastConfig {}; + +template <> +struct CastConfig : DefaultCastConfig {}; + +template <> +struct CastConfig : DefaultCastConfig {}; + +template +struct CastTraitsImpl : tuned_1D_scaling_common::KernelTraits { + static constexpr ShapeRepresentation SHAPE_REPRESENTATION = SHAPE_REP; + static constexpr int STATIC_PERSISTENT_BLOCKS_PER_SM = Config::STATIC_PERSISTENT_BLOCKS_PER_SM; + + static_assert(STATIC_PERSISTENT_BLOCKS_PER_SM > 0, + "STATIC_PERSISTENT_BLOCKS_PER_SM must be greater than zero."); +}; + +template +struct CastTraits : CastTraitsImpl> {}; + +using RNG_t = typename transformer_engine::curanddx::detail::philox4x32_native_state< + NVTE_BUILD_NUM_PHILOX_ROUNDS>; + +template +struct WorkProvider { + using ActiveCastTraits = CastTraits; + + static constexpr bool FIXED_X_DIM = SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS || + SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM; + static constexpr int CHUNK_DIM_Y = ActiveCastTraits::CHUNK_DIM_Y; + static constexpr int CHUNK_DIM_X = ActiveCastTraits::CHUNK_DIM_X; + + TensorMetadata metadata_; + int tensor_id_; + int rows_; + int cols_; + int scale_stride_; + int scale_stride_t_; + int block_id_X_; + size_t blocks_X_per_tensor_; + size_t blocks_Y_per_tensor_; + size_t tiles_per_tensor_; + size_t current_block_id_Y_; + size_t next_block_id_Y_; + size_t current_tile_id_; + size_t next_tile_id_; + size_t work_stride_; + size_t rowwise_scale_base_; + size_t colwise_scale_base_; + size_t launch_block_id_; + bool valid_; + + __device__ __forceinline__ WorkProvider(const size_t num_tensors, const int common_rows, + const int common_cols, const int common_scale_stride, + const int common_scale_stride_t, + const size_t common_blocks_Y_per_tensor) + : valid_(false) { + if constexpr (FIXED_X_DIM) { + tensor_id_ = static_cast(blockIdx.z); + rows_ = common_rows; + cols_ = common_cols; + scale_stride_ = common_scale_stride; + scale_stride_t_ = common_scale_stride_t; + blocks_X_per_tensor_ = 0; + blocks_Y_per_tensor_ = common_blocks_Y_per_tensor; + tiles_per_tensor_ = 0; + current_tile_id_ = 0; + next_tile_id_ = 0; + + if constexpr (SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM) { + metadata_ = g_tensor_metadata[tensor_id_]; + if (metadata_.rows == 0 || metadata_.cols == 0) { + return; + } + rows_ = static_cast(metadata_.rows); + scale_stride_t_ = DIVUP_TO_MULTIPLE(DIVUP(rows_, static_cast(NVFP4_SCALE_DIM)), 4); + blocks_Y_per_tensor_ = DIVUP(rows_, static_cast(CHUNK_DIM_Y)); + } + + block_id_X_ = static_cast(blockIdx.x); + current_block_id_Y_ = static_cast(blockIdx.y); + if (current_block_id_Y_ >= blocks_Y_per_tensor_) { + return; + } + + launch_block_id_ = + (static_cast(blockIdx.z) * gridDim.y + blockIdx.y) * gridDim.x + blockIdx.x; + work_stride_ = static_cast(gridDim.y); + + const size_t tensor_id_u = static_cast(tensor_id_); + rowwise_scale_base_ = tensor_id_u * rows_ * scale_stride_; + colwise_scale_base_ = tensor_id_u * cols_ * scale_stride_t_; + if constexpr (SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM) { + rowwise_scale_base_ = metadata_.rowwise_scale_base; + colwise_scale_base_ = metadata_.colwise_scale_base; + } + } else { + const size_t tensor_id_u = static_cast(blockIdx.y); + if (tensor_id_u >= num_tensors) { + return; + } + + tensor_id_ = static_cast(tensor_id_u); + metadata_ = g_tensor_metadata[tensor_id_]; + const size_t rows_u = metadata_.rows; + const size_t cols_u = metadata_.cols; + if (rows_u == 0 || cols_u == 0) { + return; + } + + rows_ = static_cast(rows_u); + cols_ = static_cast(cols_u); + scale_stride_ = DIVUP_TO_MULTIPLE(DIVUP(cols_u, static_cast(NVFP4_SCALE_DIM)), 4); + scale_stride_t_ = DIVUP_TO_MULTIPLE(DIVUP(rows_u, static_cast(NVFP4_SCALE_DIM)), 4); + block_id_X_ = 0; + blocks_X_per_tensor_ = DIVUP(cols_u, static_cast(CHUNK_DIM_X)); + blocks_Y_per_tensor_ = DIVUP(rows_u, static_cast(CHUNK_DIM_Y)); + tiles_per_tensor_ = blocks_X_per_tensor_ * blocks_Y_per_tensor_; + current_block_id_Y_ = 0; + next_block_id_Y_ = 0; + current_tile_id_ = static_cast(blockIdx.x); + if (current_tile_id_ >= tiles_per_tensor_) { + return; + } + + launch_block_id_ = tensor_id_u * gridDim.x + blockIdx.x; + work_stride_ = static_cast(gridDim.x); + rowwise_scale_base_ = metadata_.rowwise_scale_base; + colwise_scale_base_ = metadata_.colwise_scale_base; + } + + valid_ = true; + } + + __device__ __forceinline__ bool is_valid() const { return valid_; } + __device__ __forceinline__ int tensor_id() const { return tensor_id_; } + __device__ __forceinline__ int rows() const { return rows_; } + __device__ __forceinline__ int cols() const { return cols_; } + __device__ __forceinline__ int scale_stride() const { return scale_stride_; } + __device__ __forceinline__ int scale_stride_t() const { return scale_stride_t_; } + __device__ __forceinline__ size_t rowwise_scale_base() const { return rowwise_scale_base_; } + __device__ __forceinline__ size_t colwise_scale_base() const { return colwise_scale_base_; } + __device__ __forceinline__ size_t launch_block_id() const { return launch_block_id_; } + + __device__ __forceinline__ void current_block_ids(int &block_id_Y, int &block_id_X) const { + if constexpr (FIXED_X_DIM) { + block_id_Y = static_cast(current_block_id_Y_); + block_id_X = block_id_X_; + } else { + const size_t block_id_Y_u = current_tile_id_ / blocks_X_per_tensor_; + const size_t block_id_X_u = current_tile_id_ - block_id_Y_u * blocks_X_per_tensor_; + block_id_Y = static_cast(block_id_Y_u); + block_id_X = static_cast(block_id_X_u); + } + } + + __device__ __forceinline__ void prepare_next(bool &job_finished, int &prefetch_block_offset_Y, + int &prefetch_block_offset_X) { + if constexpr (FIXED_X_DIM) { + next_block_id_Y_ = current_block_id_Y_ + work_stride_; + job_finished = (next_block_id_Y_ >= blocks_Y_per_tensor_); + if (!job_finished) { + prefetch_block_offset_Y = static_cast(next_block_id_Y_ * CHUNK_DIM_Y); + prefetch_block_offset_X = block_id_X_ * CHUNK_DIM_X; + } + } else { + next_tile_id_ = current_tile_id_ + work_stride_; + job_finished = (next_tile_id_ >= tiles_per_tensor_); + if (!job_finished) { + const size_t prefetch_block_id_Y = next_tile_id_ / blocks_X_per_tensor_; + const size_t prefetch_block_id_X = + next_tile_id_ - prefetch_block_id_Y * blocks_X_per_tensor_; + prefetch_block_offset_Y = static_cast(prefetch_block_id_Y * CHUNK_DIM_Y); + prefetch_block_offset_X = static_cast(prefetch_block_id_X * CHUNK_DIM_X); + } + } + } + + __device__ __forceinline__ void commit_next() { + if constexpr (FIXED_X_DIM) { + current_block_id_Y_ = next_block_id_Y_; + } else { + current_tile_id_ = next_tile_id_; + } + } +}; + +template +__global__ void __launch_bounds__(CastTraits::THREADS_NUM) + group_quantize_transpose_nvfp4_tuned_1D_kernel( + const size_t num_tensors, nvfp4_scale_t *const scales_ptr, + nvfp4_scale_t *const scales_t_ptr, const float *noop, const float *const amax_rowwise_ptr, + const float *const amax_colwise_ptr, const size_t amax_rowwise_numel, + const size_t amax_colwise_numel, const int common_rows, const int common_cols, + const int common_scale_stride, const int common_scale_stride_t, + const size_t common_blocks_Y_per_tensor, const size_t *rng_state) { +#if (defined __CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + using ScalingTraits = CastTraits; + using IType = typename ScalingTraits::IType; + using IType3D = typename ScalingTraits::IType3D; + using OType2x3D = typename ScalingTraits::OType2x3D; + using OType2xt3D = typename ScalingTraits::OType2xt3D; + using ScalesType2D = typename ScalingTraits::ScalesType2D; + using ScalesTypeTr2D = typename ScalingTraits::ScalesTypeTr2D; + + if (noop != nullptr && noop[0] == 1.0f) { + return; + } + + WorkProvider work(num_tensors, common_rows, common_cols, common_scale_stride, + common_scale_stride_t, common_blocks_Y_per_tensor); + if (!work.is_valid()) { + return; + } + + extern __shared__ char dynamic_shmem[]; + __shared__ uint64_t IN_buff_readable_mbar[ScalingTraits::BUFFS_NUM]; + + constexpr int CHUNK_DIM_Y = ScalingTraits::CHUNK_DIM_Y; + constexpr int CHUNK_DIM_X = ScalingTraits::CHUNK_DIM_X; + constexpr int PREFETCH_STAGES = ScalingTraits::PREFETCH_STAGES; + constexpr int THREADS_NUM = ScalingTraits::THREADS_NUM; + constexpr int TILE_DIM_Y = ScalingTraits::TILE_DIM_Y; + constexpr int TILE_DIM_X = ScalingTraits::TILE_DIM_X; + constexpr int STAGES_X = ScalingTraits::STAGES_X; + constexpr int STAGES = ScalingTraits::STAGES; + constexpr int BUFFS_NUM = ScalingTraits::BUFFS_NUM; + constexpr int BUFFS_NUM_IN = ScalingTraits::BUFFS_NUM_IN; + constexpr int BUFFS_NUM_OUT = ScalingTraits::BUFFS_NUM_OUT; + constexpr int BUFFS_NUM_OUT_TR = ScalingTraits::BUFFS_NUM_OUT_TR; + constexpr int BUFF_SIZE_ALIGNED_IN = ScalingTraits::BUFF_SIZE_ALIGNED_IN; + constexpr int BUFF_SIZE_ALIGNED_OUT = ScalingTraits::BUFF_SIZE_ALIGNED_OUT; + constexpr int BUFF_SIZE_ALIGNED_OUT_TR = ScalingTraits::BUFF_SIZE_ALIGNED_OUT_TR; + constexpr int BUFF_SIZE_ROWWISE_SCALES = ScalingTraits::BUFF_SIZE_ROWWISE_SCALES; + constexpr int SCALES_PER_CHUNK_X = ScalingTraits::SCALES_PER_CHUNK_X; + constexpr int SCALES_PER_CHUNK_Y = ScalingTraits::SCALES_PER_CHUNK_Y; + + const int tensor_id = work.tensor_id(); + const int rows = work.rows(); + const int cols = work.cols(); + const int scale_stride = work.scale_stride(); + const int scale_stride_t = work.scale_stride_t(); + + const size_t rng_sequence = threadIdx.x + work.launch_block_id() * THREADS_NUM; + const size_t rng_seed = rng_state != nullptr ? rng_state[0] : 0; + const size_t rng_offset = rng_state != nullptr ? rng_state[1] : 0; + RNG_t rng; + rng.init(rng_seed, rng_sequence, rng_offset); + uint4 random_uint4 = USE_STOCHASTIC_ROUNDING ? rng.generate4() : uint4{0, 0, 0, 0}; + int rnd_idx = 0; + + const bool leading_thread = (threadIdx.x == 0); + + const int amax_rowwise_idx = (amax_rowwise_numel > 1) ? tensor_id : 0; + const float S_enc_rowwise = + (amax_rowwise_ptr == nullptr || amax_rowwise_numel == 0) + ? 1.0f + : core::compute_global_encode_scaling_factor_FP4(amax_rowwise_ptr[amax_rowwise_idx]); + const int amax_colwise_idx = (amax_colwise_numel > 1) ? tensor_id : 0; + const float S_enc_colwise = + (amax_colwise_ptr == nullptr || amax_colwise_numel == 0) + ? S_enc_rowwise + : core::compute_global_encode_scaling_factor_FP4(amax_colwise_ptr[amax_colwise_idx]); + + nvfp4_scale_t *const scales_rowwise = scales_ptr + work.rowwise_scale_base(); + nvfp4_scale_t *const scales_colwise = + RETURN_TRANSPOSE ? (scales_t_ptr + work.colwise_scale_base()) : nullptr; + + const CUtensorMap &tensor_map_input = g_tensor_maps.input[tensor_id]; + const CUtensorMap &tensor_map_output = g_tensor_maps.output_rowwise[tensor_id]; + const CUtensorMap &tensor_map_output_t = g_tensor_maps.output_colwise[tensor_id]; + + constexpr int in_mem = BUFF_SIZE_ALIGNED_IN; + constexpr int out_mem_rowwise_data = BUFF_SIZE_ALIGNED_OUT; + constexpr int out_mem_colwise_data = RETURN_TRANSPOSE ? BUFF_SIZE_ALIGNED_OUT_TR : 0; + constexpr int out_mem_rowwise_scales = BUFF_SIZE_ROWWISE_SCALES; + + char *dshmem = align_up(dynamic_shmem, TMA_SHMEM_ALIGNMENT); + + IType *sIn_ptr = reinterpret_cast(dshmem); + fp4e2m1x2 *sOut_ptr = reinterpret_cast(dshmem + in_mem); + fp4e2m1x2 *sOut_tr_ptr = reinterpret_cast(dshmem + in_mem + out_mem_rowwise_data); + + auto &sIn = *reinterpret_cast(sIn_ptr); + auto &sOut = *reinterpret_cast(sOut_ptr); + auto &sOut_tr = *reinterpret_cast(sOut_tr_ptr); + + nvfp4_scale_t *sSFrowwise_ptr = reinterpret_cast( + dshmem + in_mem + out_mem_rowwise_data + out_mem_colwise_data); + nvfp4_scale_t *sSFcolwise_ptr = reinterpret_cast( + dshmem + in_mem + out_mem_rowwise_data + out_mem_colwise_data + out_mem_rowwise_scales); + auto &sSFrowwise = *reinterpret_cast(sSFrowwise_ptr); + auto &sSFcolwise = *reinterpret_cast(sSFcolwise_ptr); + + constexpr int shmem_buff_size = BUFF_SIZE_ALIGNED_IN / BUFFS_NUM; + + if (leading_thread) { +#pragma unroll + for (int buff = 0; buff < BUFFS_NUM; ++buff) { + ptx::mbarrier_init(&IN_buff_readable_mbar[buff], 1); + } + ptx::fence_proxy_async_shared_cta(); + } + __syncthreads(); + + if (leading_thread) { + fence_acquire_tensormap(&tensor_map_input); + fence_acquire_tensormap(&tensor_map_output); + if constexpr (RETURN_TRANSPOSE) { + fence_acquire_tensormap(&tensor_map_output_t); + } + } + + { + int first_block_id_Y = 0; + int first_block_id_X = 0; + work.current_block_ids(first_block_id_Y, first_block_id_X); + const int first_block_offset_Y = first_block_id_Y * CHUNK_DIM_Y; + const int first_block_offset_X = first_block_id_X * CHUNK_DIM_X; +#pragma unroll + for (int stage = 0; stage < PREFETCH_STAGES; ++stage) { + const int stage_Y = stage / STAGES_X; + const int stage_X = stage % STAGES_X; + const int stage_offset_Y = stage_Y * TILE_DIM_Y; + const int stage_offset_X = stage_X * TILE_DIM_X; + const int global_offset_Y = first_block_offset_Y + stage_offset_Y; + const int global_offset_X = first_block_offset_X + stage_offset_X; + if (leading_thread) { + uint64_t *dst = reinterpret_cast(&sIn[stage]); + const uint64_t *src = reinterpret_cast(&tensor_map_input); + uint64_t *barrier = &IN_buff_readable_mbar[stage]; + ptx::mbarrier_arrive_expect_tx(barrier, shmem_buff_size); + ptx::cp_async_bulk_tensor_2d_global_to_shared(dst, src, global_offset_X, global_offset_Y, + barrier); + } + } + } + + int buff_in = 0; + int buff_out = 0; + int buff_out_tr = 0; + int IN_buff_readable_parity[BUFFS_NUM] = {0}; + + bool job_finished = false; + while (!job_finished) { + int block_id_Y = 0; + int block_id_X = 0; + work.current_block_ids(block_id_Y, block_id_X); + + const int block_offset_Y = block_id_Y * CHUNK_DIM_Y; + const int block_offset_X = block_id_X * CHUNK_DIM_X; + const int block_offset_Y_tr = block_offset_X; + const int block_offset_X_tr = block_offset_Y; + const int chunk_rows = rows - block_offset_Y; + const int chunk_cols = cols - block_offset_X; + const int scales_block_offset_Y_rowwise = block_id_Y * CHUNK_DIM_Y; + const int scales_block_offset_X_rowwise = block_id_X * SCALES_PER_CHUNK_X; + const int scales_block_offset_Y_tr = block_id_X * CHUNK_DIM_X; + const int scales_block_offset_X_tr = block_id_Y * SCALES_PER_CHUNK_Y; + + int prefetch_block_offset_Y = block_offset_Y; + int prefetch_block_offset_X = block_offset_X; + +#pragma unroll + for (int stage = 0; stage < STAGES; ++stage) { + const int stage_Y = stage / STAGES_X; + const int stage_X = stage % STAGES_X; + const int stage_offset_Y = stage_Y * TILE_DIM_Y; + const int stage_offset_X = stage_X * TILE_DIM_X; + + if (stage == STAGES - PREFETCH_STAGES) { + work.prepare_next(job_finished, prefetch_block_offset_Y, prefetch_block_offset_X); + } + + if ((stage < STAGES - PREFETCH_STAGES) || !job_finished) { + const int next_prefetch_buff = (buff_in + PREFETCH_STAGES) % BUFFS_NUM; + const int next_prefetch_stage = (stage + PREFETCH_STAGES) % STAGES; + const int next_prefetch_stage_Y = next_prefetch_stage / STAGES_X; + const int next_prefetch_stage_X = next_prefetch_stage % STAGES_X; + const int next_prefetch_stage_offset_Y = next_prefetch_stage_Y * TILE_DIM_Y; + const int next_prefetch_stage_offset_X = next_prefetch_stage_X * TILE_DIM_X; + const bool prefetch_next_tile = (stage >= STAGES - PREFETCH_STAGES); + const int prefetch_base_offset_Y = + prefetch_next_tile ? prefetch_block_offset_Y : block_offset_Y; + const int prefetch_base_offset_X = + prefetch_next_tile ? prefetch_block_offset_X : block_offset_X; + const int global_offset_Y = prefetch_base_offset_Y + next_prefetch_stage_offset_Y; + const int global_offset_X = prefetch_base_offset_X + next_prefetch_stage_offset_X; + + if (leading_thread) { + uint64_t *dst = reinterpret_cast(&sIn[next_prefetch_buff]); + const uint64_t *src = reinterpret_cast(&tensor_map_input); + uint64_t *barrier = &IN_buff_readable_mbar[next_prefetch_buff]; + ptx::mbarrier_arrive_expect_tx(barrier, shmem_buff_size); + ptx::cp_async_bulk_tensor_2d_global_to_shared(dst, src, global_offset_X, global_offset_Y, + barrier); + } + ptx::fence_proxy_async_shared_cta(); + } + + ptx::mbarrier_wait_parity_acquire_cta_shared_cta(&IN_buff_readable_mbar[buff_in], + IN_buff_readable_parity[buff_in]); + IN_buff_readable_parity[buff_in] ^= 1; + + // Bulk async-groups are per-thread. Only the leading thread issues and commits the TMA + // stores, so it is also the only thread whose wait observes their completion. Hand that + // completion off to every cooperative writer before the output ring buffer is reused. + if (leading_thread) { + ptx::cp_async_bulk_wait_group_read(); + } + __syncthreads(); + + rowwise_scaling( + sIn_ptr, sOut_ptr, sSFrowwise_ptr, S_enc_rowwise, stage_Y, stage_X, buff_in, buff_out, + rng, random_uint4, rnd_idx); + if constexpr (RETURN_TRANSPOSE) { + colwise_scaling( + sIn_ptr, sOut_tr_ptr, sSFcolwise_ptr, S_enc_colwise, stage_Y, stage_X, buff_in, + buff_out_tr, rng, random_uint4, rnd_idx); + } + + ptx::fence_proxy_async_shared_cta(); + __syncthreads(); + + if (leading_thread) { + const int global_offset_Y = block_offset_Y + stage_offset_Y; + const int global_offset_X = block_offset_X + stage_offset_X; + ptx::cp_async_bulk_tensor_2d_shared_to_global( + reinterpret_cast(&tensor_map_output), global_offset_X, + global_offset_Y, reinterpret_cast(&sOut[buff_out])); + + if constexpr (RETURN_TRANSPOSE) { + const int global_offset_Y_tr = block_offset_Y_tr + stage_offset_X; + const int global_offset_X_tr = block_offset_X_tr + stage_offset_Y; + ptx::cp_async_bulk_tensor_2d_shared_to_global( + reinterpret_cast(&tensor_map_output_t), global_offset_X_tr, + global_offset_Y_tr, reinterpret_cast(&sOut_tr[buff_out_tr])); + } + ptx::cp_async_bulk_commit_group(); + } + + buff_in = (buff_in + 1) % BUFFS_NUM_IN; + buff_out = (buff_out + 1) % BUFFS_NUM_OUT; + buff_out_tr = (buff_out_tr + 1) % BUFFS_NUM_OUT_TR; + } + + { + using RowwiseScalesVec = Vec; + const int rowwise_count = + min(SCALES_PER_CHUNK_X, chunk_cols / static_cast(NVFP4_SCALE_DIM)); + for (int row = threadIdx.x; row < CHUNK_DIM_Y; row += THREADS_NUM) { + const int row_global_i = scales_block_offset_Y_rowwise + row; + if (row_global_i < rows) { + const size_t row_global = static_cast(row_global_i); + RowwiseScalesVec &scales_vec = *reinterpret_cast(sSFrowwise[row]); + const size_t scale_idx_global = row_global * scale_stride + scales_block_offset_X_rowwise; + scales_vec.store_to_elts(&scales_rowwise[scale_idx_global], 0, rowwise_count); + } + } + + if constexpr (RETURN_TRANSPOSE) { + using ColwiseScalesVec = Vec; + const int colwise_count = + min(SCALES_PER_CHUNK_Y, chunk_rows / static_cast(NVFP4_SCALE_DIM)); + for (int row_tr = threadIdx.x; row_tr < CHUNK_DIM_X; row_tr += THREADS_NUM) { + const int row_tr_global_i = scales_block_offset_Y_tr + row_tr; + if (row_tr_global_i < cols) { + const size_t row_tr_global = static_cast(row_tr_global_i); + ColwiseScalesVec &scales_vec = + *reinterpret_cast(sSFcolwise[row_tr]); + const size_t scale_idx_global = + row_tr_global * scale_stride_t + scales_block_offset_X_tr; + scales_vec.store_to_elts(&scales_colwise[scale_idx_global], 0, colwise_count); + } + } + } + + if (!job_finished) { + work.commit_next(); + __syncthreads(); + } + } + } + + // Drain every TMA store before the CTA releases its shared-memory source buffers. + if (leading_thread) { + ptx::cp_async_bulk_wait_group(); + } + __syncthreads(); + + if (leading_thread) { +#pragma unroll + for (int buff = 0; buff < BUFFS_NUM; ++buff) { + ptx::mbarrier_invalid(&IN_buff_readable_mbar[buff]); + } + } +#else + NVTE_DEVICE_ERROR("sm_100 or higher is required."); +#endif +} + +template +inline void launch_group_quantize_transpose_kernel( + const size_t num_tensors, const size_t first_logical_dim, const size_t last_logical_dim, + nvfp4_scale_t *const scales_ptr, nvfp4_scale_t *const scales_t_ptr, const float *const noop_ptr, + const float *const amax_rowwise_ptr, const float *const amax_colwise_ptr, + const size_t amax_rowwise_numel, const size_t amax_colwise_numel, const size_t work_blocks_X, + const size_t work_blocks_Y, const size_t *const rng_state, const int dshmem_size, + cudaStream_t stream) { + using ScalingTraits = CastTraits; + constexpr int CHUNK_DIM_Y = ScalingTraits::CHUNK_DIM_Y; + constexpr int THREADS_NUM = ScalingTraits::THREADS_NUM; + + const int block_size = THREADS_NUM; + auto kernel = group_quantize_transpose_nvfp4_tuned_1D_kernel; + NVTE_CHECK_CUDA( + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, dshmem_size)); + + int active_blocks_per_sm = 0; + NVTE_CHECK_CUDA(cudaOccupancyMaxActiveBlocksPerMultiprocessor(&active_blocks_per_sm, kernel, + block_size, dshmem_size)); + NVTE_CHECK(active_blocks_per_sm > 0, + "Grouped NVFP4 optimized kernel has zero active blocks per SM."); + + const size_t sm_num = static_cast(transformer_engine::cuda::sm_count()); + + if constexpr (SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS) { + const size_t rows_per_tensor = first_logical_dim / num_tensors; + const size_t blocks_Y_per_tensor = DIVUP(rows_per_tensor, static_cast(CHUNK_DIM_Y)); + const int rows = static_cast(rows_per_tensor); + const int cols = static_cast(last_logical_dim); + const int scale_stride = + DIVUP_TO_MULTIPLE(DIVUP(last_logical_dim, static_cast(NVFP4_SCALE_DIM)), 4); + const int scale_stride_t = + DIVUP_TO_MULTIPLE(DIVUP(rows_per_tensor, static_cast(NVFP4_SCALE_DIM)), 4); + const size_t requested_grid_size = sm_num * static_cast(active_blocks_per_sm); + const size_t workers_X_total = num_tensors * work_blocks_X; + const size_t requested_workers_Y_per_tensor = + std::max(size_t{1}, requested_grid_size / workers_X_total); + const size_t workers_Y_per_tensor = + std::min(blocks_Y_per_tensor, requested_workers_Y_per_tensor); + const dim3 grid(work_blocks_X, workers_Y_per_tensor, num_tensors); + + kernel<<>>( + num_tensors, scales_ptr, scales_t_ptr, noop_ptr, amax_rowwise_ptr, amax_colwise_ptr, + amax_rowwise_numel, amax_colwise_numel, rows, cols, scale_stride, scale_stride_t, + blocks_Y_per_tensor, rng_state); + } else if constexpr (SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM) { + const int cols = static_cast(last_logical_dim); + const int scale_stride = + DIVUP_TO_MULTIPLE(DIVUP(last_logical_dim, static_cast(NVFP4_SCALE_DIM)), 4); + const size_t requested_grid_size = sm_num * static_cast(active_blocks_per_sm); + const size_t workers_X_total = num_tensors * work_blocks_X; + const size_t requested_workers_Y_per_tensor = + std::max(size_t{1}, requested_grid_size / workers_X_total); + const size_t avg_blocks_Y_per_tensor = + std::max(size_t{1}, DIVUP(work_blocks_Y, num_tensors)); + const size_t workers_Y_per_tensor = + std::min(avg_blocks_Y_per_tensor, requested_workers_Y_per_tensor); + NVTE_CHECK(workers_Y_per_tensor > 0, + "VARYING_FIRST_DIM persistent grid size must be greater than zero."); + const dim3 grid(work_blocks_X, workers_Y_per_tensor, num_tensors); + + kernel<<>>( + num_tensors, scales_ptr, scales_t_ptr, noop_ptr, amax_rowwise_ptr, amax_colwise_ptr, + amax_rowwise_numel, amax_colwise_numel, 0, cols, scale_stride, 0, 0, rng_state); + } else { + const size_t total_work_blocks = work_blocks_X * work_blocks_Y; + const size_t persistent_blocks_per_sm = + std::min(active_blocks_per_sm, ScalingTraits::STATIC_PERSISTENT_BLOCKS_PER_SM); + const size_t requested_workers_per_tensor = + std::max(size_t{1}, (sm_num * persistent_blocks_per_sm) / num_tensors); + const size_t workers_per_tensor = std::min(total_work_blocks, requested_workers_per_tensor); + NVTE_CHECK(workers_per_tensor > 0, + "Tensor-local persistent grid size must be greater than zero."); + const dim3 grid(workers_per_tensor, num_tensors); + + kernel<<>>( + num_tensors, scales_ptr, scales_t_ptr, noop_ptr, amax_rowwise_ptr, amax_colwise_ptr, + amax_rowwise_numel, amax_colwise_numel, 0, 0, 0, 0, 0, rng_state); + } + + NVTE_CHECK_CUDA(cudaGetLastError()); +} + +#endif // FP4_TYPE_SUPPORTED +} // namespace group_quantize_transpose_tuned_kernel + +inline void group_quantize_transpose(const GroupedTensor *input, const Tensor *noop, + GroupedTensor *output, const QuantizationConfig *quant_config, + cudaStream_t stream) { +#if FP4_TYPE_SUPPORTED + using namespace group_quantize_transpose_tuned_kernel; + using namespace ptx; + + const bool use_stochastic_rounding = quant_config ? quant_config->stochastic_rounding : false; + const bool use_fast_math = quant_config ? quant_config->use_fast_math : false; + const bool return_transpose = output->has_columnwise_data(); + + checkCuDriverContext(stream); + CheckNoopTensor(*noop, "cast_noop"); + + NVTE_CHECK(input->num_tensors == output->num_tensors, + "Number of input and output tensors must be same."); + NVTE_CHECK(input->has_data(), "Cannot quantize tensor without rowwise data."); + NVTE_CHECK(input->dtype() == DType::kBFloat16, + "Optimized grouped NVFP4 kernel supports only BF16 input."); + NVTE_CHECK(output->has_data(), "Grouped NVFP4 output tensor must be allocated."); + NVTE_CHECK(is_fp4_dtype(output->dtype()), "Output must have FP4 type."); + NVTE_CHECK(output->scale_inv.dptr != nullptr, "Scaling tensor must be allocated."); + NVTE_CHECK(!output->with_gemm_swizzled_scales, "Output must have scales in compact format."); + if (return_transpose) { + NVTE_CHECK(is_fp4_dtype(output->columnwise_data.dtype), + "Transposed output must have FP4 type."); + NVTE_CHECK(output->columnwise_scale_inv.dptr != nullptr, + "Transposed scaling tensor must be allocated."); + } + + ShapeRepresentation shape_rep = ShapeRepresentation::SAME_BOTH_DIMS; + if (output->all_same_shape()) { + shape_rep = ShapeRepresentation::SAME_BOTH_DIMS; + } else if (output->all_same_first_dim()) { + shape_rep = ShapeRepresentation::VARYING_LAST_DIM; + } else if (output->all_same_last_dim()) { + shape_rep = ShapeRepresentation::VARYING_FIRST_DIM; + } else if (output->varying_both_dims()) { + shape_rep = ShapeRepresentation::VARYING_BOTH_DIMS; + } + + const size_t first_logical_dim = input->logical_shape.data[0]; + const size_t last_logical_dim = input->logical_shape.data[1]; + const size_t elts_total = first_logical_dim * last_logical_dim; + const size_t num_tensors = input->num_tensors; + + NVTE_CHECK(num_tensors <= MAX_SUPPORTED_TENSOR_DESCRIPTORS, + "Number of tensors in a group is larger than the MAX number of supported " + "descriptors (64)."); + const int64_t *const offsets_ptr = reinterpret_cast(output->tensor_offsets.dptr); + const int64_t *const first_dims_ptr = reinterpret_cast(output->first_dims.dptr); + const int64_t *const last_dims_ptr = reinterpret_cast(output->last_dims.dptr); + + nvfp4_scale_t *const scales_ptr = reinterpret_cast(output->scale_inv.dptr); + nvfp4_scale_t *const scales_t_ptr = + reinterpret_cast(output->columnwise_scale_inv.dptr); + + const float *noop_ptr = reinterpret_cast(noop->data.dptr); + const float *const amax_rowwise_ptr = reinterpret_cast(input->amax.dptr); + const float *const amax_colwise_ptr = + reinterpret_cast(input->columnwise_amax.dptr); + const size_t amax_rowwise_numel = input->amax.has_data() ? input->amax.numel() : 0; + const size_t amax_colwise_numel = + input->columnwise_amax.has_data() ? input->columnwise_amax.numel() : 0; + + if (input->amax.has_data()) { + NVTE_CHECK(amax_rowwise_numel == 1 || amax_rowwise_numel == num_tensors, + "Rowwise amax must contain either 1 value or num_tensors values, found ", + amax_rowwise_numel, " values for num_tensors=", num_tensors, "."); + } + if (input->columnwise_amax.has_data()) { + NVTE_CHECK(amax_colwise_numel == 1 || amax_colwise_numel == num_tensors, + "Columnwise amax must contain either 1 value or num_tensors values, found ", + amax_colwise_numel, " values for num_tensors=", num_tensors, "."); + } + + const NVTETensor rng_state_tensor = (quant_config != nullptr) ? quant_config->rng_state : nullptr; + const size_t *rng_state = nullptr; + if (rng_state_tensor != nullptr) { + Tensor &rng_state_te_tensor = *convertNVTETensor(rng_state_tensor); + NVTE_CHECK(rng_state_te_tensor.dtype() == DType::kInt64, + "RNG state should contain 2 64-bit values."); + NVTE_CHECK(rng_state_te_tensor.data.shape == std::vector{2}, + "Shape of the RNG state should be [2], but got ", rng_state_te_tensor.data.shape); + rng_state = reinterpret_cast(rng_state_te_tensor.data.dptr); + } + + TRANSFORMER_ENGINE_GROUP_TENSOR_SHAPE_REPRESENTATION_SWITCH(shape_rep, SHAPE_REP, { + using ActiveCastTraits = CastTraits; + using IType = typename ActiveCastTraits::IType; + + constexpr int CHUNK_DIM_Y = ActiveCastTraits::CHUNK_DIM_Y; + constexpr int CHUNK_DIM_X = ActiveCastTraits::CHUNK_DIM_X; + constexpr size_t ELTS_PER_CHUNK = ActiveCastTraits::ELTS_PER_CHUNK; + constexpr int BUFF_DIM_Y = ActiveCastTraits::BUFF_DIM_Y; + constexpr int BUFF_DIM_X = ActiveCastTraits::BUFF_DIM_X; + constexpr int BUFF_SIZE_ALIGNED_IN = ActiveCastTraits::BUFF_SIZE_ALIGNED_IN; + constexpr int BUFF_SIZE_ALIGNED_OUT = ActiveCastTraits::BUFF_SIZE_ALIGNED_OUT; + constexpr int BUFF_SIZE_ALIGNED_OUT_TR = ActiveCastTraits::BUFF_SIZE_ALIGNED_OUT_TR; + constexpr int BUFF_SIZE_ROWWISE_SCALES = ActiveCastTraits::BUFF_SIZE_ROWWISE_SCALES; + constexpr int BUFF_SIZE_COLWISE_SCALES = ActiveCastTraits::BUFF_SIZE_COLWISE_SCALES; + constexpr bool USE_SINGLE_WORK_GRID = SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS || + SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM; + + if constexpr (SHAPE_REP == ShapeRepresentation::SAME_BOTH_DIMS) { + NVTE_CHECK(first_logical_dim % num_tensors == 0, + "First logical dimension of a grouped tensor must be divisible by the number of " + "tensors."); + NVTE_CHECK((first_logical_dim / num_tensors) % CHUNK_DIM_Y == 0, + "First dimension of each tensor in a group must be divisible by ", CHUNK_DIM_Y, + "."); + } else if constexpr (SHAPE_REP == ShapeRepresentation::VARYING_FIRST_DIM) { + NVTE_CHECK(first_logical_dim % CHUNK_DIM_Y == 0, + "First logical dimension of a grouped tensor must be divisible by ", CHUNK_DIM_Y, + "."); + } else if constexpr (SHAPE_REP == ShapeRepresentation::VARYING_LAST_DIM) { + NVTE_CHECK(first_logical_dim % CHUNK_DIM_Y == 0, + "First logical dimension of a grouped tensor must be divisible by ", CHUNK_DIM_Y, + "."); + NVTE_CHECK(last_logical_dim % CHUNK_DIM_Y == 0, + "Last logical dimension of a grouped tensor must be divisible by ", CHUNK_DIM_Y, + "."); + } else { + NVTE_CHECK(last_logical_dim % ELTS_PER_CHUNK == 0, + "Last logical dimension of a grouped tensor must be divisible by ", CHUNK_DIM_Y, + "x", CHUNK_DIM_X, "."); + } + + size_t work_blocks_X = 0; + size_t work_blocks_Y = 0; + if constexpr (USE_SINGLE_WORK_GRID) { + work_blocks_Y = DIVUP(first_logical_dim, static_cast(CHUNK_DIM_Y)); + work_blocks_X = DIVUP(last_logical_dim, static_cast(CHUNK_DIM_X)); + } else { + work_blocks_Y = 1; + work_blocks_X = DIVUP(elts_total, ELTS_PER_CHUNK); + } + + alignas(64) CUtensorMap tensor_map_input{}; + alignas(64) CUtensorMap tensor_map_act_input{}; + alignas(64) CUtensorMap tensor_map_output{}; + alignas(64) CUtensorMap tensor_map_output_transpose{}; + + const size_t dummy_first_logical_dim = 32; + const size_t dummy_last_logical_dim = 32; + create_2D_tensor_map(tensor_map_input, input->data, dummy_first_logical_dim, + dummy_last_logical_dim, BUFF_DIM_Y, BUFF_DIM_X, dummy_last_logical_dim, 0, + sizeof(IType) * 8); + create_2D_tensor_map(tensor_map_output, output->data, dummy_first_logical_dim, + dummy_last_logical_dim, BUFF_DIM_Y, BUFF_DIM_X, dummy_last_logical_dim, 0, + 4); + if (return_transpose) { + create_2D_tensor_map(tensor_map_output_transpose, output->columnwise_data, + dummy_last_logical_dim, dummy_first_logical_dim, BUFF_DIM_X, BUFF_DIM_Y, + dummy_first_logical_dim, 0, 4); + } + + const int in_mem = BUFF_SIZE_ALIGNED_IN; + const int out_data_mem = BUFF_SIZE_ALIGNED_OUT; + const int out_data_transpose_mem = return_transpose ? BUFF_SIZE_ALIGNED_OUT_TR : 0; + const int out_scales_mem = BUFF_SIZE_ROWWISE_SCALES; + const int out_scales_transpose_mem = return_transpose ? BUFF_SIZE_COLWISE_SCALES : 0; + const int out_mem = out_data_mem + out_data_transpose_mem; + const int dshmem_size = + in_mem + out_mem + out_scales_transpose_mem + out_scales_mem + TMA_SHMEM_ALIGNMENT; + + const IType *const input_dptr = reinterpret_cast(input->data.dptr); + const void *const output_dptr = output->data.dptr; + const void *const output_t_dptr = return_transpose ? output->columnwise_data.dptr : nullptr; + + update_tma_descriptors<<>>( + tensor_map_input, tensor_map_act_input, tensor_map_output, tensor_map_output_transpose, + input_dptr, nullptr, output_dptr, output_t_dptr, shape_rep, num_tensors, first_logical_dim, + last_logical_dim, offsets_ptr, first_dims_ptr, last_dims_ptr, true, return_transpose, + false); + NVTE_CHECK_CUDA(cudaGetLastError()); + + TRANSFORMER_ENGINE_SWITCH_CONDITION( + use_stochastic_rounding, USE_STOCHASTIC_ROUNDING, + TRANSFORMER_ENGINE_SWITCH_CONDITION( + use_fast_math, USE_FAST_MATH, + TRANSFORMER_ENGINE_SWITCH_CONDITION(return_transpose, RETURN_TRANSPOSE, { + launch_group_quantize_transpose_kernel( + num_tensors, first_logical_dim, last_logical_dim, scales_ptr, scales_t_ptr, + noop_ptr, amax_rowwise_ptr, amax_colwise_ptr, amax_rowwise_numel, + amax_colwise_numel, work_blocks_X, work_blocks_Y, rng_state, dshmem_size, stream); + }););); + }); +#else + NVTE_ERROR("FP4 support requires CUDA 12.8+, but compile-time CUDA version is ", CUDA_VERSION); +#endif +} + +} // namespace nvfp4 +} // namespace dispatch +} // namespace transformer_engine + +#endif // TRANSFORMER_ENGINE_GROUP_QUANTIZE_TRANSPOSE_NVFP4_TUNED_1D_CUH_ diff --git a/transformer_engine/common/cast/nvfp4/specialized/quantize_transpose_nvfp4_tuned_1D.cuh b/transformer_engine/common/cast/nvfp4/specialized/quantize_transpose_nvfp4_tuned_1D.cuh index cdd0d4916a..2a12b03080 100644 --- a/transformer_engine/common/cast/nvfp4/specialized/quantize_transpose_nvfp4_tuned_1D.cuh +++ b/transformer_engine/common/cast/nvfp4/specialized/quantize_transpose_nvfp4_tuned_1D.cuh @@ -21,6 +21,7 @@ #include "../../../util/ptx_arch_spec.cuh" #include "../../../utils.cuh" #include "../core_nvfp4.cuh" +#include "scaling_nvfp4_tuned_1D.cuh" namespace transformer_engine { namespace dispatch { @@ -34,348 +35,54 @@ using namespace ptx; #if FP4_TYPE_SUPPORTED +using tuned_1D_scaling_common::colwise_scaling; +using tuned_1D_scaling_common::rowwise_scaling; + struct TunableConfig { - static constexpr int CHUNK_DIM_Y = 128; - static constexpr int CHUNK_DIM_X = 128; - static constexpr int PREFETCH_STAGES = 1; static constexpr bool PERSISTENT = false; }; -constexpr int SCALE_DIM = 16; // NVFP4 block (x16 elts) -constexpr int THREADS_NUM = 128; -constexpr int ELTS_PER_THREAD = 16; -constexpr int TILE_DIM_Y = 64; -constexpr int TILE_DIM_X = 64; - -static_assert(ELTS_PER_THREAD == SCALE_DIM && "Hardcoded and fixed parameter\0"); - -static_assert((THREADS_NUM * ELTS_PER_THREAD <= TILE_DIM_Y * TILE_DIM_X) && - "Unbalanced threads workload\0"); - -static_assert((TunableConfig::CHUNK_DIM_Y % TILE_DIM_Y == 0) && - "Chunk size Y must be evenly divisible by the tile size Y\0"); -static_assert((TunableConfig::CHUNK_DIM_X % TILE_DIM_X == 0) && - "Chunk size X must be evenly divisible by the tile size X\0"); - -static_assert((TILE_DIM_Y % SCALE_DIM == 0) && - "Tile size Y must be evenly divisible by the scale dim\0"); -static_assert((TILE_DIM_X % SCALE_DIM == 0) && - "Tile size X must be evenly divisible by the scale dim\0"); - -constexpr int TILES_Y = TunableConfig::CHUNK_DIM_Y / TILE_DIM_Y; -constexpr int TILES_X = TunableConfig::CHUNK_DIM_X / TILE_DIM_X; - -constexpr int THREADS_PER_SCALE_ROWWISE = SCALE_DIM / ELTS_PER_THREAD; - -constexpr int SCALES_PER_CHUNK_Y = TunableConfig::CHUNK_DIM_Y / SCALE_DIM; -constexpr int SCALES_PER_CHUNK_X = TunableConfig::CHUNK_DIM_X / SCALE_DIM; - -constexpr int SCALES_PER_TILE_Y = TILE_DIM_Y / SCALE_DIM; -constexpr int SCALES_PER_TILE_X = TILE_DIM_X / SCALE_DIM; - -constexpr int STAGES_Y = TILES_Y; -constexpr int STAGES_X = TILES_X; -constexpr int STAGES = STAGES_Y * STAGES_X; - -constexpr int BUFFS_NUM = TunableConfig::PREFETCH_STAGES + 1; -constexpr int BUFFS_NUM_IN = BUFFS_NUM; -constexpr int BUFFS_NUM_OUT = BUFFS_NUM; -constexpr int BUFFS_NUM_OUT_TR = 2; -constexpr int BUFF_DIM_Y = TILE_DIM_Y; -constexpr int BUFF_DIM_X = TILE_DIM_X; -constexpr int BUFF_SIZE = BUFF_DIM_Y * BUFF_DIM_X; -constexpr int BUFF_SIZE_TOTAL = BUFF_SIZE * BUFFS_NUM; - -// Input buffer (BF16) -constexpr int BUFF_IN_DIM_Y = BUFF_DIM_Y; -constexpr int BUFF_IN_DIM_X = BUFF_DIM_X; -constexpr int BUFF_IN_SIZE = BUFF_IN_DIM_Y * BUFF_IN_DIM_X; -constexpr int BUFF_IN_ELTS_NUM = BUFF_IN_DIM_Y * BUFF_IN_DIM_X; - -// Output buffer (NVFP4) -constexpr int BUFF_OUT_DIM_Y = BUFF_DIM_Y; -constexpr int BUFF_OUT_DIM_X = (BUFF_DIM_X * 4) / 8; -constexpr int BUFF_OUT_SIZE = BUFF_OUT_DIM_Y * BUFF_OUT_DIM_X; - -// Output transpose buffer (NVFP4) -constexpr int BUFF_OUT_TR_DIM_Y = BUFF_DIM_X; -constexpr int BUFF_OUT_TR_DIM_X = (BUFF_DIM_Y * 4) / 8; -constexpr int BUFF_OUT_TR_SIZE = BUFF_OUT_TR_DIM_Y * BUFF_OUT_TR_DIM_X; - -// Manual swizzling parameters to reduce SHMEM bank conflicts -constexpr int PACK_SIZE = 8; -constexpr int WAVES = ELTS_PER_THREAD / PACK_SIZE; - -constexpr int THREADS_X_ROWWISE = TILE_DIM_X / ELTS_PER_THREAD; -constexpr int THREADS_Y_ROWWISE = THREADS_NUM / THREADS_X_ROWWISE; - -constexpr int THREADS_X_TR = TILE_DIM_X / 2; -constexpr int THREADS_Y_TR = THREADS_NUM / THREADS_X_TR; - -constexpr int ITERATIONS_NORMAL = BUFF_DIM_Y / THREADS_Y_ROWWISE; -constexpr int ITERATIONS_TR = SCALES_PER_TILE_Y / THREADS_Y_TR; -static_assert(ITERATIONS_TR >= 1 && "Number of transpose iterations should be >=1\0"); -static_assert((SCALES_PER_TILE_Y % THREADS_Y_TR == 0) && - "Partial transpose iterations are not supported\0"); - -constexpr int BUFF_OUT_IT_OFFSET = BUFF_OUT_TR_DIM_X / ITERATIONS_TR / STAGES; - -static_assert(BUFF_DIM_Y >= SCALE_DIM && - "Number of buffer rows must be greater or equal to the size of the columwise " - "scaling block\0"); -static_assert(TunableConfig::CHUNK_DIM_Y >= BUFF_DIM_Y); -static_assert(BUFF_DIM_Y >= THREADS_Y_ROWWISE && - "Number of buffer rows must be greater or equal to the number of rowwise " - "processing threads in Y dimension\0"); - -// Number of 4-bit elements that span 32 banks (4-byte each) of shared memory -constexpr int TOTAL_BANKS_WIDTH = (32 * 4 * 8) / 4; // 256 - -// Number of threads (rowwise scaling) that span 32 banks (4-byte banks) of shared memory -constexpr int THREADS_PER_BANK = TOTAL_BANKS_WIDTH / ELTS_PER_THREAD; - -using IType = bf16; -using IType2 = typename ptx::FPx2; -using IType3D = IType[BUFFS_NUM_IN][BUFF_IN_DIM_Y][BUFF_IN_DIM_X]; -using IType2x3D = IType2[BUFFS_NUM_IN][BUFF_IN_DIM_Y][BUFF_IN_DIM_X / 2]; -using OType2x3D = fp4e2m1x2[BUFFS_NUM_OUT][BUFF_OUT_DIM_Y][BUFF_OUT_DIM_X]; -using OType2xt3D = fp4e2m1x2[BUFFS_NUM_OUT_TR][BUFF_OUT_TR_DIM_Y][BUFF_OUT_TR_DIM_X]; -using ScalesType2D = nvfp4_scale_t[TunableConfig::CHUNK_DIM_Y][SCALES_PER_CHUNK_X]; -using ScalesTypeTr2D = nvfp4_scale_t[TunableConfig::CHUNK_DIM_X][SCALES_PER_CHUNK_Y]; using RNG_t = typename transformer_engine::curanddx::detail::philox4x32_native_state< NVTE_BUILD_NUM_PHILOX_ROUNDS>; -template -struct SCALING_COEFFICIENT_TYPE {}; -template <> -struct SCALING_COEFFICIENT_TYPE { - using type = float; -}; -template <> -struct SCALING_COEFFICIENT_TYPE { - using type = bf16; -}; - -__device__ __forceinline__ float get_amax_of_pair(const IType2 pair) { - return static_cast(__hmax(__habs(pair.x), __habs(pair.y))); -} - -// Compute "correct" per-block encoding scaling factor -template -__device__ __forceinline__ SF_TYPE -compute_nvfp4_scaling_coefficient(const nvfp4_scale_t S_dec_block, const float S_enc) { - NVTE_DEVICE_ERROR("Unsupported scaling-factor type. Only FP32 and BF16 are supported."); -} - -template <> -__device__ __forceinline__ float compute_nvfp4_scaling_coefficient( - const nvfp4_scale_t S_dec_block, const float S_enc) { - const float S_dec = 1.0f / S_enc; - const float scale_rcp = - fminf(1.0f / (static_cast(S_dec_block) * S_dec), detail::TypeExtrema::max); - return scale_rcp; -} - -template <> -__device__ __forceinline__ bf16 -compute_nvfp4_scaling_coefficient(const nvfp4_scale_t S_dec_block, const float S_enc) { - const float scale_rcp = - fminf(S_enc / (static_cast(S_dec_block)), detail::TypeExtrema::max); - return static_cast(scale_rcp); -} - -template -__device__ __forceinline__ void colwise_scaling( - const IType *__restrict__ sIn_ptr, fp4e2m1x2 *__restrict__ sOut_tr_ptr, - nvfp4_scale_t *__restrict__ sSFcolwise_ptr, const float S_enc_colwise, const int stage_Y, - const int stage_X, const int buff_in, const int buff_out_tr, const float *amax_colwise_ptr, - const size_t col_offset, const size_t cols, RNG_t &rng, uint4 &random_uint4, int &rnd_idx) { - using scaling_coeff_type = typename SCALING_COEFFICIENT_TYPE::type; - - const auto &sIn2x = *reinterpret_cast(sIn_ptr); - auto &sOut_tr = *reinterpret_cast(sOut_tr_ptr); - auto &sSFcolwise = *reinterpret_cast(sSFcolwise_ptr); - - const int warp = threadIdx.x / THREADS_PER_WARP; - const int thread_lane = threadIdx.x % THREADS_PER_WARP; - - const int tid_Y_colwise = (thread_lane / 2 + warp) % 4; - const int tid_X_colwise = thread_lane; - - const int thread_offset_Y_colwise = tid_Y_colwise * SCALE_DIM; - const int thread_offset_X_colwise = tid_X_colwise * 2; - - const int in_thread_offset_Y = thread_offset_Y_colwise; - const int in_thread_offset_X = thread_offset_X_colwise / 2; - - const int out_tr_thread_offset_Y = thread_offset_X_colwise; - const int out_tr_thread_offset_X = thread_offset_Y_colwise / 2; - - const int scale_tr_offset_Y = (stage_X * TILE_DIM_X) + 2 * tid_X_colwise; - const int scale_tr_offset_X = (stage_Y * SCALES_PER_TILE_Y) + tid_Y_colwise; - - __align__(8) IType rIn[2][SCALE_DIM]; - // Read (cache) a pair of input elements (S2R). Find NVFP4-block AMAX - IType2 thread_amax_2x = {static_cast(0.0f), static_cast(0.0f)}; -#pragma unroll - for (int i = 0; i < SCALE_DIM; ++i) { - const IType2 elt_pair = - ptx::ld_shared_b32(&sIn2x[buff_in][in_thread_offset_Y + i][in_thread_offset_X]); - rIn[0][i] = elt_pair.x; - rIn[1][i] = elt_pair.y; - ptx::abs_max_2x(thread_amax_2x, thread_amax_2x, elt_pair); - } - const float block_amax[2] = {static_cast(__habs(thread_amax_2x.x)), - static_cast(__habs(thread_amax_2x.y))}; -#pragma unroll - for (int w = 0; w < 2; ++w) { - float S_enc_colwise_block = S_enc_colwise; - if constexpr (ROW_SCALED_NVFP4) { - const size_t col_idx = col_offset + stage_X * TILE_DIM_X + thread_offset_X_colwise + w; - S_enc_colwise_block = - col_idx < cols ? core::compute_global_encode_scaling_factor_FP4(amax_colwise_ptr[col_idx]) - : 1.0f; - } - const nvfp4_scale_t S_dec_b_fp8 = - compute_decoding_scaling_factor(block_amax[w], S_enc_colwise_block); - - // Store scaling factors to SMEM buffer (R2S) - sSFcolwise[scale_tr_offset_Y + w][scale_tr_offset_X] = S_dec_b_fp8; - - const scaling_coeff_type SFcoefficient = - compute_nvfp4_scaling_coefficient(S_dec_b_fp8, S_enc_colwise_block); - - // Scale elements - __align__(8) uint32_t rOut[SCALE_DIM / 8]; -#pragma unroll - for (int e = 0; e < SCALE_DIM / 8; ++e) { - const uint64_t elts03 = *reinterpret_cast(&rIn[w][8 * e]); - const uint64_t elts47 = *reinterpret_cast(&rIn[w][8 * e + 4]); - if constexpr (USE_STOCHASTIC_ROUNDING) { - const uint32_t rbits03 = core::get_rbits(rng, random_uint4, rnd_idx); - const uint32_t rbits47 = core::get_rbits(rng, random_uint4, rnd_idx); - rOut[e] = ptx::mul_cvt_bf16_to_fp4_8x_stochastic_rounding( - elts03, elts47, SFcoefficient, rbits03, rbits47); - } else { - rOut[e] = ptx::mul_cvt_bf16_to_fp4_8x_round_to_nearest(elts03, elts47, - SFcoefficient); - } - } - uint64_t &out_pack_16x = *reinterpret_cast(rOut); - ptx::st_shared_b64(&sOut_tr[buff_out_tr][out_tr_thread_offset_Y + w][out_tr_thread_offset_X], - out_pack_16x); - } -} - -template -__device__ __forceinline__ void rowwise_scaling( - const IType *__restrict__ sIn_ptr, fp4e2m1x2 *__restrict__ sOut_ptr, - nvfp4_scale_t *__restrict__ sSFrowwise_ptr, const float S_enc_rowwise, const int stage_Y, - const int stage_X, const int buff_in, const int buff_out, const float *amax_rowwise_ptr, - const size_t row_offset, const size_t rows, RNG_t &rng, uint4 &random_uint4, int &rnd_idx) { - using scaling_coeff_type = typename SCALING_COEFFICIENT_TYPE::type; - - const auto &sIn = *reinterpret_cast(sIn_ptr); - auto &sOut = *reinterpret_cast(sOut_ptr); - auto &sSFrowwise = *reinterpret_cast(sSFrowwise_ptr); - - const int thread_lane = threadIdx.x % THREADS_PER_WARP; - const int bank_group = thread_lane / THREADS_PER_BANK; - - const int tid_Y_rowwise = threadIdx.x / THREADS_X_ROWWISE; - const int tid_X_rowwise = threadIdx.x % THREADS_X_ROWWISE; - - const int thread_offset_Y_rowwise = tid_Y_rowwise; - const int thread_offset_X_rowwise = tid_X_rowwise * ELTS_PER_THREAD; - - const int SF_thread_offset_rowwise_Y = tid_Y_rowwise; - const int SF_thread_offset_rowwise_X = tid_X_rowwise / THREADS_PER_SCALE_ROWWISE; - - const bool SF_storing_thread = (tid_X_rowwise % THREADS_PER_SCALE_ROWWISE == 0); - - const int stage_rowwise_scales_offset_Y = SF_thread_offset_rowwise_Y + stage_Y * TILE_DIM_Y; - const int stage_rowwise_scales_offset_X = - SF_thread_offset_rowwise_X + stage_X * SCALES_PER_TILE_X; -#pragma unroll - for (int it = 0; it < ITERATIONS_NORMAL; ++it) { - const int it_offset_Y_rowwise = thread_offset_Y_rowwise + it * THREADS_Y_ROWWISE; - - __align__(16) IType2 rIn[WAVES][PACK_SIZE / 2]; - - // Read (cache) input elements (S2R). Find NVFP4-block AMAX - IType2 thread_amax_2x = {static_cast(0.0f), static_cast(0.0f)}; -#pragma unroll - for (int w = 0; w < WAVES; ++w) { - const int swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % ELTS_PER_THREAD; - const int swizzled_thread_idx = thread_offset_X_rowwise + swizzled_group_idx; - - // Load elements - __uint128_t &elts_8x = *reinterpret_cast<__uint128_t *>(&rIn[w]); - elts_8x = ptx::ld_shared_b128(&sIn[buff_in][it_offset_Y_rowwise][swizzled_thread_idx]); -#pragma unroll - for (int e = 0; e < PACK_SIZE / 2; ++e) { - ptx::abs_max_2x(thread_amax_2x, thread_amax_2x, rIn[w][e]); - } - } - const float block_amax = get_amax_of_pair(thread_amax_2x); - - nvfp4_scale_t S_dec_b_fp8; - scaling_coeff_type SFcoefficient; - if constexpr (ROW_SCALED_NVFP4) { - const size_t row_idx = row_offset + stage_Y * TILE_DIM_Y + it_offset_Y_rowwise; - const float S_enc_rowwise_block = - row_idx < rows ? core::compute_global_encode_scaling_factor_FP4(amax_rowwise_ptr[row_idx]) - : 1.0f; - S_dec_b_fp8 = compute_decoding_scaling_factor(block_amax, S_enc_rowwise_block); - SFcoefficient = - compute_nvfp4_scaling_coefficient(S_dec_b_fp8, S_enc_rowwise_block); - } else { - S_dec_b_fp8 = compute_decoding_scaling_factor(block_amax, S_enc_rowwise); - SFcoefficient = - compute_nvfp4_scaling_coefficient(S_dec_b_fp8, S_enc_rowwise); - } - - // Store scaling factors to SMEM buffer (R2S) - if (SF_storing_thread) { - const int scales_offset_Y = stage_rowwise_scales_offset_Y + it * THREADS_Y_ROWWISE; - const int scales_offset_X = stage_rowwise_scales_offset_X; - sSFrowwise[scales_offset_Y][scales_offset_X] = S_dec_b_fp8; - } - -// Scale elements -#pragma unroll - for (int w = 0; w < WAVES; ++w) { - const uint64_t elts03 = *reinterpret_cast(&rIn[w][0]); - const uint64_t elts47 = *reinterpret_cast(&rIn[w][2]); - - uint32_t out_x8; - if constexpr (USE_STOCHASTIC_ROUNDING) { - const uint32_t rbits03 = core::get_rbits(rng, random_uint4, rnd_idx); - const uint32_t rbits47 = core::get_rbits(rng, random_uint4, rnd_idx); - out_x8 = ptx::mul_cvt_bf16_to_fp4_8x_stochastic_rounding( - elts03, elts47, SFcoefficient, rbits03, rbits47); - } else { - out_x8 = ptx::mul_cvt_bf16_to_fp4_8x_round_to_nearest(elts03, elts47, - SFcoefficient); - } - - const int swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % ELTS_PER_THREAD; - const int swizzled_idx = (swizzled_group_idx + thread_offset_X_rowwise) / 2; - ptx::st_shared_b32(&sOut[buff_out][it_offset_Y_rowwise][swizzled_idx], out_x8); - } - } -} +using ScalingTraits = tuned_1D_scaling_common::NonGroupedKernelTraits; +using IType = typename ScalingTraits::IType; +using IType3D = typename ScalingTraits::IType3D; +using OType2x3D = typename ScalingTraits::OType2x3D; +using OType2xt3D = typename ScalingTraits::OType2xt3D; +using ScalesType2D = typename ScalingTraits::ScalesType2D; +using ScalesTypeTr2D = typename ScalingTraits::ScalesTypeTr2D; template -__global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D_kernel( - const __grid_constant__ CUtensorMap tensor_map_input, - const __grid_constant__ CUtensorMap tensor_map_output, - const __grid_constant__ CUtensorMap tensor_map_output_t, nvfp4_scale_t *const scales_ptr, - nvfp4_scale_t *const scales_t_ptr, const float *noop, const float *const amax_rowwise_ptr, - const float *const amax_colwise_ptr, const size_t rows, const size_t cols, - const size_t scale_stride, const size_t scale_stride_t, const size_t *rng_state) { +__global__ void __launch_bounds__(ScalingTraits::THREADS_NUM) + quantize_transpose_nvfp4_tuned_1D_kernel( + const __grid_constant__ CUtensorMap tensor_map_input, + const __grid_constant__ CUtensorMap tensor_map_output, + const __grid_constant__ CUtensorMap tensor_map_output_t, nvfp4_scale_t *const scales_ptr, + nvfp4_scale_t *const scales_t_ptr, const float *noop, const float *const amax_rowwise_ptr, + const float *const amax_colwise_ptr, const size_t rows, const size_t cols, + const size_t scale_stride, const size_t scale_stride_t, const size_t *rng_state) { #if (defined __CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000) + constexpr int CHUNK_DIM_Y = ScalingTraits::CHUNK_DIM_Y; + constexpr int CHUNK_DIM_X = ScalingTraits::CHUNK_DIM_X; + constexpr int PREFETCH_STAGES = ScalingTraits::PREFETCH_STAGES; + constexpr int THREADS_NUM = ScalingTraits::THREADS_NUM; + constexpr int TILE_DIM_Y = ScalingTraits::TILE_DIM_Y; + constexpr int TILE_DIM_X = ScalingTraits::TILE_DIM_X; + constexpr int STAGES_X = ScalingTraits::STAGES_X; + constexpr int STAGES = ScalingTraits::STAGES; + constexpr int BUFFS_NUM = ScalingTraits::BUFFS_NUM; + constexpr int BUFFS_NUM_IN = ScalingTraits::BUFFS_NUM_IN; + constexpr int BUFFS_NUM_OUT = ScalingTraits::BUFFS_NUM_OUT; + constexpr int BUFFS_NUM_OUT_TR = ScalingTraits::BUFFS_NUM_OUT_TR; + constexpr int BUFF_SIZE_ALIGNED_IN = ScalingTraits::BUFF_SIZE_ALIGNED_IN; + constexpr int BUFF_SIZE_ALIGNED_OUT = ScalingTraits::BUFF_SIZE_ALIGNED_OUT; + constexpr int BUFF_SIZE_ALIGNED_OUT_TR = ScalingTraits::BUFF_SIZE_ALIGNED_OUT_TR; + constexpr int BUFF_SIZE_ROWWISE_SCALES = ScalingTraits::BUFF_SIZE_ROWWISE_SCALES; + constexpr int SCALES_PER_CHUNK_X = ScalingTraits::SCALES_PER_CHUNK_X; + constexpr int SCALES_PER_CHUNK_Y = ScalingTraits::SCALES_PER_CHUNK_Y; + if (noop != nullptr && noop[0] == 1.0f) { return; } @@ -392,22 +99,11 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D const bool leading_thread = (threadIdx.x == 0); - constexpr int buff_elems = BUFF_DIM_Y * BUFF_IN_DIM_X; - constexpr int buff_elems_total_in = BUFFS_NUM_IN * buff_elems; - - constexpr int buff_size_aligned_in = - DIVUP_TO_MULTIPLE(buff_elems_total_in * sizeof(IType), TMA_SHMEM_ALIGNMENT); - constexpr int buff_size_aligned_out = - DIVUP_TO_MULTIPLE(BUFFS_NUM_OUT * BUFF_OUT_SIZE, TMA_SHMEM_ALIGNMENT); - constexpr int buff_size_aligned_out_t = - DIVUP_TO_MULTIPLE(BUFFS_NUM_OUT_TR * BUFF_OUT_TR_SIZE, TMA_SHMEM_ALIGNMENT); - - constexpr int in_mem = buff_size_aligned_in; + constexpr int in_mem = BUFF_SIZE_ALIGNED_IN; - constexpr int out_mem_rowwise_data = buff_size_aligned_out; - constexpr int out_mem_colwise_data = RETURN_TRANSPOSE ? buff_size_aligned_out_t : 0; - constexpr int out_mem_rowwise_scales = DIVUP_TO_MULTIPLE( - TunableConfig::CHUNK_DIM_Y * SCALES_PER_CHUNK_X * sizeof(nvfp4_scale_t), TMA_SHMEM_ALIGNMENT); + constexpr int out_mem_rowwise_data = BUFF_SIZE_ALIGNED_OUT; + constexpr int out_mem_colwise_data = RETURN_TRANSPOSE ? BUFF_SIZE_ALIGNED_OUT_TR : 0; + constexpr int out_mem_rowwise_scales = BUFF_SIZE_ROWWISE_SCALES; // The destination shared memory buffer of a bulk tensor operation should be 16-byte aligned extern __shared__ unsigned char dynamic_shmem[]; @@ -429,7 +125,7 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D auto &sSFrowwise = *reinterpret_cast(sSFrowwise_ptr); auto &sSFcolwise = *reinterpret_cast(sSFcolwise_ptr); - constexpr int shmem_buff_size = buff_size_aligned_in / BUFFS_NUM; + constexpr int shmem_buff_size = BUFF_SIZE_ALIGNED_IN / BUFFS_NUM; // Compute a global encoding/decoding scaling factors for all S_dec_b const float S_enc_rowwise = @@ -474,7 +170,7 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D // Prefetch input data only when processing the first chunk, // which enables the one-iteration overlap throughout the entire kernel life #pragma unroll - for (int stage = 0; stage < TunableConfig::PREFETCH_STAGES; ++stage) { + for (int stage = 0; stage < PREFETCH_STAGES; ++stage) { const int buff_in = stage; const int stage_Y = stage / STAGES_X; const int stage_X = stage % STAGES_X; @@ -482,8 +178,8 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D const int stage_offset_Y = stage_Y * TILE_DIM_Y; const int stage_offset_X = stage_X * TILE_DIM_X; - const int block_offset_Y = ctaid_Y * TunableConfig::CHUNK_DIM_Y; - const int block_offset_X = ctaid_X * TunableConfig::CHUNK_DIM_X; + const int block_offset_Y = ctaid_Y * CHUNK_DIM_Y; + const int block_offset_X = ctaid_X * CHUNK_DIM_X; const int global_offset_Y = block_offset_Y + stage_offset_Y; const int global_offset_X = block_offset_X + stage_offset_X; @@ -503,18 +199,18 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D } while (!job_finished) { - const int block_offset_Y = ctaid_Y * TunableConfig::CHUNK_DIM_Y; - const int block_offset_X = ctaid_X * TunableConfig::CHUNK_DIM_X; + const int block_offset_Y = ctaid_Y * CHUNK_DIM_Y; + const int block_offset_X = ctaid_X * CHUNK_DIM_X; - const int block_offset_Y_tr = ctaid_X * TunableConfig::CHUNK_DIM_X; - const int block_offset_X_tr = ctaid_Y * TunableConfig::CHUNK_DIM_Y; + const int block_offset_Y_tr = ctaid_X * CHUNK_DIM_X; + const int block_offset_X_tr = ctaid_Y * CHUNK_DIM_Y; const int chunk_rows = rows - block_offset_Y; const int chunk_cols = cols - block_offset_X; - const int scales_block_offset_Y_rowwise = ctaid_Y * TunableConfig::CHUNK_DIM_Y; + const int scales_block_offset_Y_rowwise = ctaid_Y * CHUNK_DIM_Y; const int scales_block_offset_X_rowwise = ctaid_X * SCALES_PER_CHUNK_X; - const int scales_block_offset_Y_tr = ctaid_X * TunableConfig::CHUNK_DIM_X; + const int scales_block_offset_Y_tr = ctaid_X * CHUNK_DIM_X; const int scales_block_offset_X_tr = ctaid_Y * SCALES_PER_CHUNK_Y; if constexpr (TunableConfig::PERSISTENT) { @@ -532,7 +228,7 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D const int stage_offset_Y = stage_Y * TILE_DIM_Y; const int stage_offset_X = stage_X * TILE_DIM_X; - if (stage == STAGES - TunableConfig::PREFETCH_STAGES) { + if (stage == STAGES - PREFETCH_STAGES) { if constexpr (TunableConfig::PERSISTENT) { ptx::mbarrier_wait_parity_acquire_cta_shared_cta(&workID_mbar, ctaid_parity); ptx::get_cancelled_cta_id_2D(&workID_response, ctaid_X, ctaid_Y); @@ -547,9 +243,9 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D } // Prefetch next stage Input data - if (!job_finished || (stage < STAGES - TunableConfig::PREFETCH_STAGES)) { - const int next_prefetch_buff = (buff_in + TunableConfig::PREFETCH_STAGES) % BUFFS_NUM; - const int next_prefetch_stage = (stage + TunableConfig::PREFETCH_STAGES) % STAGES; + if (!job_finished || (stage < STAGES - PREFETCH_STAGES)) { + const int next_prefetch_buff = (buff_in + PREFETCH_STAGES) % BUFFS_NUM; + const int next_prefetch_stage = (stage + PREFETCH_STAGES) % STAGES; const int next_prefetch_stage_Y = next_prefetch_stage / STAGES_X; const int next_prefetch_stage_X = next_prefetch_stage % STAGES_X; @@ -557,8 +253,8 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D const int next_prefetch_stage_offset_X = next_prefetch_stage_X * TILE_DIM_X; // Offsets change, because coordinates of the next "to-be-prefetched" CTA do also chage - const int block_offset_Y = ctaid_Y * TunableConfig::CHUNK_DIM_Y; - const int block_offset_X = ctaid_X * TunableConfig::CHUNK_DIM_X; + const int block_offset_Y = ctaid_Y * CHUNK_DIM_Y; + const int block_offset_X = ctaid_X * CHUNK_DIM_X; const int global_offset_Y = block_offset_Y + next_prefetch_stage_offset_Y; const int global_offset_X = block_offset_X + next_prefetch_stage_offset_X; @@ -585,15 +281,20 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D // Wait for TMA transfer to have finished reading shared memory // I.e. the OUT buffer is ready to be written to - ptx::cp_async_bulk_wait_group_read(); + if (leading_thread) { + ptx::cp_async_bulk_wait_group_read(); + } + // Bulk async-groups are thread-local. Publish the leading thread's completion to all + // threads before they cooperatively overwrite a reused output buffer. + __syncthreads(); // NVFP4 Quantization - rowwise_scaling( + rowwise_scaling( sIn_ptr, sOut_ptr, sSFrowwise_ptr, S_enc_rowwise, stage_Y, stage_X, buff_in, buff_out, amax_rowwise_ptr, block_offset_Y, rows, rng, random_uint4, rnd_idx); if constexpr (RETURN_TRANSPOSE) { - colwise_scaling( + colwise_scaling( sIn_ptr, sOut_tr_ptr, sSFcolwise_ptr, S_enc_colwise, stage_Y, stage_X, buff_in, buff_out_tr, amax_colwise_ptr, block_offset_X, cols, rng, random_uint4, rnd_idx); } @@ -635,9 +336,9 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D { using ScalesVec = Vec; // number of scales in X dimension of this chunk - const int count = min(SCALES_PER_CHUNK_X, chunk_cols / SCALE_DIM); + const int count = min(SCALES_PER_CHUNK_X, chunk_cols / NVFP4_SCALE_DIM); - for (size_t row = threadIdx.x; row < TunableConfig::CHUNK_DIM_Y; row += THREADS_NUM) { + for (size_t row = threadIdx.x; row < CHUNK_DIM_Y; row += THREADS_NUM) { const size_t row_global = scales_block_offset_Y_rowwise + row; if (row_global < rows) { ScalesVec &scales_vec = *reinterpret_cast(sSFrowwise[row]); @@ -652,10 +353,9 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D if constexpr (RETURN_TRANSPOSE) { using ScalesVec = Vec; // number of scales in Y dimension of this chunk - const int count = min(SCALES_PER_CHUNK_Y, chunk_rows / SCALE_DIM); + const int count = min(SCALES_PER_CHUNK_Y, chunk_rows / NVFP4_SCALE_DIM); - for (size_t row_tr = threadIdx.x; row_tr < TunableConfig::CHUNK_DIM_X; - row_tr += THREADS_NUM) { + for (size_t row_tr = threadIdx.x; row_tr < CHUNK_DIM_X; row_tr += THREADS_NUM) { const size_t row_tr_global = scales_block_offset_Y_tr + row_tr; if (row_tr_global < cols) { ScalesVec &scales_vec = *reinterpret_cast(sSFcolwise[row_tr]); @@ -673,6 +373,11 @@ __global__ void __launch_bounds__(THREADS_NUM) quantize_transpose_nvfp4_tuned_1D } } + if (leading_thread) { + ptx::cp_async_bulk_wait_group(); + } + __syncthreads(); + if (leading_thread) { #pragma unroll for (int buff = 0; buff < BUFFS_NUM; ++buff) { @@ -695,6 +400,17 @@ inline void quantize_transpose_tuned_1D(const Tensor &input, const Tensor *noop, using namespace quantize_transpose_tuned_kernel; using namespace ptx; + constexpr int CHUNK_DIM_Y = ScalingTraits::CHUNK_DIM_Y; + constexpr int CHUNK_DIM_X = ScalingTraits::CHUNK_DIM_X; + constexpr int THREADS_NUM = ScalingTraits::THREADS_NUM; + constexpr int BUFF_DIM_Y = ScalingTraits::BUFF_DIM_Y; + constexpr int BUFF_DIM_X = ScalingTraits::BUFF_DIM_X; + constexpr int BUFF_SIZE_ALIGNED_IN = ScalingTraits::BUFF_SIZE_ALIGNED_IN; + constexpr int BUFF_SIZE_ALIGNED_OUT = ScalingTraits::BUFF_SIZE_ALIGNED_OUT; + constexpr int BUFF_SIZE_ALIGNED_OUT_TR = ScalingTraits::BUFF_SIZE_ALIGNED_OUT_TR; + constexpr int BUFF_SIZE_ROWWISE_SCALES = ScalingTraits::BUFF_SIZE_ROWWISE_SCALES; + constexpr int BUFF_SIZE_COLWISE_SCALES = ScalingTraits::BUFF_SIZE_COLWISE_SCALES; + const bool use_stochastic_rounding = quant_config ? quant_config->stochastic_rounding : false; const bool use_fast_math = quant_config ? quant_config->use_fast_math : false; const bool row_scaled_nvfp4 = output->row_scaled_nvfp4; @@ -731,8 +447,8 @@ inline void quantize_transpose_tuned_1D(const Tensor &input, const Tensor *noop, NVTE_CHECK(cols % 32 == 0, "Number of tensor cols must be a multiple of 32"); // 16B alignment for TMA - const int blocks_Y = DIVUP(rows, static_cast(TunableConfig::CHUNK_DIM_Y)); - const int blocks_X = DIVUP(cols, static_cast(TunableConfig::CHUNK_DIM_X)); + const int blocks_Y = DIVUP(rows, static_cast(CHUNK_DIM_Y)); + const int blocks_X = DIVUP(cols, static_cast(CHUNK_DIM_X)); const dim3 grid(blocks_X, blocks_Y); const int block_size = THREADS_NUM; @@ -774,26 +490,12 @@ inline void quantize_transpose_tuned_1D(const Tensor &input, const Tensor *noop, BUFF_DIM_X, BUFF_DIM_Y, rows, 0, 4); } - constexpr int buff_elems = BUFF_DIM_Y * BUFF_DIM_X; - constexpr int buff_elems_total_in = BUFFS_NUM_IN * buff_elems; - constexpr int buff_size_aligned_in = - DIVUP_TO_MULTIPLE(buff_elems_total_in * sizeof(IType), TMA_SHMEM_ALIGNMENT); - constexpr int buff_size_aligned_out = - DIVUP_TO_MULTIPLE(BUFFS_NUM_OUT * BUFF_OUT_SIZE, TMA_SHMEM_ALIGNMENT); - constexpr int buff_size_aligned_out_t = - DIVUP_TO_MULTIPLE(BUFFS_NUM_OUT_TR * BUFF_OUT_TR_SIZE, TMA_SHMEM_ALIGNMENT); - - constexpr int buff_size_scales = DIVUP_TO_MULTIPLE( - TunableConfig::CHUNK_DIM_Y * SCALES_PER_CHUNK_X * sizeof(nvfp4_scale_t), TMA_SHMEM_ALIGNMENT); - constexpr int buff_size_scales_transpose = DIVUP_TO_MULTIPLE( - TunableConfig::CHUNK_DIM_X * SCALES_PER_CHUNK_Y * sizeof(nvfp4_scale_t), TMA_SHMEM_ALIGNMENT); - - const int in_mem = buff_size_aligned_in; - - const int out_data_mem = buff_size_aligned_out; - const int out_data_transpose_mem = return_transpose ? buff_size_aligned_out_t : 0; - const int out_scales_mem = buff_size_scales; - const int out_scales_transpose_mem = return_transpose ? buff_size_scales_transpose : 0; + const int in_mem = BUFF_SIZE_ALIGNED_IN; + + const int out_data_mem = BUFF_SIZE_ALIGNED_OUT; + const int out_data_transpose_mem = return_transpose ? BUFF_SIZE_ALIGNED_OUT_TR : 0; + const int out_scales_mem = BUFF_SIZE_ROWWISE_SCALES; + const int out_scales_transpose_mem = return_transpose ? BUFF_SIZE_COLWISE_SCALES : 0; const int out_mem = out_data_mem + out_data_transpose_mem; diff --git a/transformer_engine/common/cast/nvfp4/specialized/scaling_nvfp4_tuned_1D.cuh b/transformer_engine/common/cast/nvfp4/specialized/scaling_nvfp4_tuned_1D.cuh new file mode 100644 index 0000000000..542cde6c64 --- /dev/null +++ b/transformer_engine/common/cast/nvfp4/specialized/scaling_nvfp4_tuned_1D.cuh @@ -0,0 +1,420 @@ +/************************************************************************* + * Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * + * See LICENSE for license information. + ************************************************************************/ + +/*! \file scaling_nvfp4_tuned_1D.cuh + * \brief Common scaling functions for tuned NVFP4 transpose kernels. + */ + +#ifndef TRANSFORMER_ENGINE_SCALING_NVFP4_TUNED_1D_CUH_ +#define TRANSFORMER_ENGINE_SCALING_NVFP4_TUNED_1D_CUH_ + +#include "../../../util/ptx_arch_spec.cuh" +#include "../core_nvfp4.cuh" + +namespace transformer_engine { +namespace dispatch { +namespace nvfp4 { +namespace tuned_1D_scaling_common { + +#if FP4_TYPE_SUPPORTED + +struct DefaultScalingConfig { + static constexpr int CHUNK_DIM_Y = 128; + static constexpr int CHUNK_DIM_X = 128; + static constexpr int PREFETCH_STAGES = 1; + static constexpr int THREADS_NUM = 128; + static constexpr int ELTS_PER_THREAD = 16; + static constexpr int TILE_DIM_Y = 64; + static constexpr int TILE_DIM_X = 64; +}; + +struct DefaultGroupedScalingConfig : DefaultScalingConfig { + static constexpr int CHUNK_DIM_X = 256; +}; + +template +struct KernelTraits { + static constexpr int CHUNK_DIM_Y = Config::CHUNK_DIM_Y; + static constexpr int CHUNK_DIM_X = Config::CHUNK_DIM_X; + static constexpr int PREFETCH_STAGES = Config::PREFETCH_STAGES; + + static constexpr int THREADS_NUM = Config::THREADS_NUM; + static constexpr int ELTS_PER_THREAD = Config::ELTS_PER_THREAD; + static constexpr int TILE_DIM_Y = Config::TILE_DIM_Y; + static constexpr int TILE_DIM_X = Config::TILE_DIM_X; + + static_assert(ELTS_PER_THREAD == NVFP4_SCALE_DIM, "Hardcoded and fixed parameter\0"); + static_assert(THREADS_NUM * ELTS_PER_THREAD <= TILE_DIM_Y * TILE_DIM_X, + "Unbalanced threads workload\0"); + static_assert(CHUNK_DIM_Y % TILE_DIM_Y == 0, + "Chunk size Y must be evenly divisible by the tile size Y\0"); + static_assert(CHUNK_DIM_X % TILE_DIM_X == 0, + "Chunk size X must be evenly divisible by the tile size X\0"); + static_assert(TILE_DIM_Y % NVFP4_SCALE_DIM == 0, + "Tile size Y must be evenly divisible by the scale dim\0"); + static_assert(TILE_DIM_X % NVFP4_SCALE_DIM == 0, + "Tile size X must be evenly divisible by the scale dim\0"); + + static constexpr int TILES_Y = CHUNK_DIM_Y / TILE_DIM_Y; + static constexpr int TILES_X = CHUNK_DIM_X / TILE_DIM_X; + static constexpr int THREADS_PER_SCALE_ROWWISE = NVFP4_SCALE_DIM / ELTS_PER_THREAD; + + static constexpr int SCALES_PER_CHUNK_Y = CHUNK_DIM_Y / NVFP4_SCALE_DIM; + static constexpr int SCALES_PER_CHUNK_X = CHUNK_DIM_X / NVFP4_SCALE_DIM; + static constexpr int SCALES_PER_TILE_Y = TILE_DIM_Y / NVFP4_SCALE_DIM; + static constexpr int SCALES_PER_TILE_X = TILE_DIM_X / NVFP4_SCALE_DIM; + + static constexpr int STAGES_Y = TILES_Y; + static constexpr int STAGES_X = TILES_X; + static constexpr int STAGES = STAGES_Y * STAGES_X; + + static_assert(PREFETCH_STAGES > 0, "At least one prefetch stage is required"); + static_assert(PREFETCH_STAGES <= STAGES, + "The number of prefetch stages cannot exceed the number of compute stages"); + + static constexpr int BUFFS_NUM = PREFETCH_STAGES + 1; + static constexpr int BUFFS_NUM_IN = BUFFS_NUM; + static constexpr int BUFFS_NUM_OUT = BUFFS_NUM; + static constexpr int BUFFS_NUM_OUT_TR = 2; + static constexpr int BUFF_DIM_Y = TILE_DIM_Y; + static constexpr int BUFF_DIM_X = TILE_DIM_X; + static constexpr int BUFF_SIZE = BUFF_DIM_Y * BUFF_DIM_X; + static constexpr int BUFF_SIZE_TOTAL = BUFF_SIZE * BUFFS_NUM; + + static constexpr int BUFF_IN_DIM_Y = BUFF_DIM_Y; + static constexpr int BUFF_IN_DIM_X = BUFF_DIM_X; + static constexpr int BUFF_IN_SIZE = BUFF_IN_DIM_Y * BUFF_IN_DIM_X; + static constexpr int BUFF_IN_ELTS_NUM = BUFF_IN_DIM_Y * BUFF_IN_DIM_X; + + static constexpr int BUFF_OUT_DIM_Y = BUFF_DIM_Y; + static constexpr int BUFF_OUT_DIM_X = (BUFF_DIM_X * 4) / 8; + static constexpr int BUFF_OUT_SIZE = BUFF_OUT_DIM_Y * BUFF_OUT_DIM_X; + + static constexpr int BUFF_OUT_TR_DIM_Y = BUFF_DIM_X; + static constexpr int BUFF_OUT_TR_DIM_X = (BUFF_DIM_Y * 4) / 8; + static constexpr int BUFF_OUT_TR_SIZE = BUFF_OUT_TR_DIM_Y * BUFF_OUT_TR_DIM_X; + + using IType = bf16; + using IType2 = typename ptx::FPx2; + + static constexpr int BUFF_ELEMS = BUFF_DIM_Y * BUFF_IN_DIM_X; + static constexpr int BUFF_ELEMS_TOTAL_IN = BUFFS_NUM_IN * BUFF_ELEMS; + static constexpr int BUFF_SIZE_ALIGNED_IN = + DIVUP_TO_MULTIPLE(BUFF_ELEMS_TOTAL_IN * sizeof(IType), TMA_SHMEM_ALIGNMENT); + static constexpr int BUFF_SIZE_ALIGNED_OUT = + DIVUP_TO_MULTIPLE(BUFFS_NUM_OUT * BUFF_OUT_SIZE, TMA_SHMEM_ALIGNMENT); + static constexpr int BUFF_SIZE_ALIGNED_OUT_TR = + DIVUP_TO_MULTIPLE(BUFFS_NUM_OUT_TR * BUFF_OUT_TR_SIZE, TMA_SHMEM_ALIGNMENT); + static constexpr int BUFF_SIZE_ROWWISE_SCALES = DIVUP_TO_MULTIPLE( + CHUNK_DIM_Y * SCALES_PER_CHUNK_X * sizeof(nvfp4_scale_t), TMA_SHMEM_ALIGNMENT); + static constexpr int BUFF_SIZE_COLWISE_SCALES = DIVUP_TO_MULTIPLE( + CHUNK_DIM_X * SCALES_PER_CHUNK_Y * sizeof(nvfp4_scale_t), TMA_SHMEM_ALIGNMENT); + + static constexpr int PACK_SIZE = 8; + static constexpr int WAVES = ELTS_PER_THREAD / PACK_SIZE; + + static constexpr int THREADS_X_ROWWISE = TILE_DIM_X / ELTS_PER_THREAD; + static constexpr int THREADS_Y_ROWWISE = THREADS_NUM / THREADS_X_ROWWISE; + + static constexpr int THREADS_X_TR = TILE_DIM_X / 2; + static constexpr int THREADS_Y_TR = THREADS_NUM / THREADS_X_TR; + + static constexpr int ITERATIONS_NORMAL = BUFF_DIM_Y / THREADS_Y_ROWWISE; + static constexpr int ITERATIONS_TR = SCALES_PER_TILE_Y / THREADS_Y_TR; + static_assert(ITERATIONS_TR >= 1, "Number of transpose iterations should be >=1\0"); + static_assert(SCALES_PER_TILE_Y % THREADS_Y_TR == 0, + "Partial transpose iterations are not supported\0"); + + static constexpr int BUFF_OUT_IT_OFFSET = BUFF_OUT_TR_DIM_X / ITERATIONS_TR / STAGES; + + static_assert(BUFF_DIM_Y >= NVFP4_SCALE_DIM, + "Number of buffer rows must be greater or equal to the size of the columwise " + "scaling block\0"); + static_assert(CHUNK_DIM_Y >= BUFF_DIM_Y); + static_assert(BUFF_DIM_Y >= THREADS_Y_ROWWISE, + "Number of buffer rows must be greater or equal to the number of rowwise " + "processing threads in Y dimension\0"); + + // Number of 4-bit elements that span 32 banks (4-byte each) of shared memory. + static constexpr int TOTAL_BANKS_WIDTH = (32 * 4 * 8) / 4; // 256 + + // Number of threads (rowwise scaling) that span 32 banks (4-byte banks) of shared memory. + static constexpr int THREADS_PER_BANK = TOTAL_BANKS_WIDTH / ELTS_PER_THREAD; + + static constexpr size_t ELTS_PER_CHUNK = + static_cast(CHUNK_DIM_Y) * static_cast(CHUNK_DIM_X); + + using IType3D = IType[BUFFS_NUM_IN][BUFF_IN_DIM_Y][BUFF_IN_DIM_X]; + using IType2x3D = IType2[BUFFS_NUM_IN][BUFF_IN_DIM_Y][BUFF_IN_DIM_X / 2]; + using OType2x3D = fp4e2m1x2[BUFFS_NUM_OUT][BUFF_OUT_DIM_Y][BUFF_OUT_DIM_X]; + using OType2xt3D = fp4e2m1x2[BUFFS_NUM_OUT_TR][BUFF_OUT_TR_DIM_Y][BUFF_OUT_TR_DIM_X]; + using ScalesType2D = nvfp4_scale_t[CHUNK_DIM_Y][SCALES_PER_CHUNK_X]; + using ScalesTypeTr2D = nvfp4_scale_t[CHUNK_DIM_X][SCALES_PER_CHUNK_Y]; +}; + +using NonGroupedKernelTraits = KernelTraits; +using GroupedKernelTraits = KernelTraits; + +template +struct ScalingCoefficientType {}; +template <> +struct ScalingCoefficientType { + using type = float; +}; +template <> +struct ScalingCoefficientType { + using type = bf16; +}; + +template +__device__ __forceinline__ float get_amax_of_pair(const PairType pair) { + return static_cast(__hmax(__habs(pair.x), __habs(pair.y))); +} + +template +__device__ __forceinline__ void colwise_scaling( + const typename Traits::IType *__restrict__ sIn_ptr, fp4e2m1x2 *__restrict__ sOut_tr_ptr, + nvfp4_scale_t *__restrict__ sSFcolwise_ptr, const float S_enc_colwise, const int stage_Y, + const int stage_X, const int buff_in, const int buff_out_tr, const float *amax_colwise_ptr, + const size_t col_offset, const size_t cols, RngType &rng, uint4 &random_uint4, int &rnd_idx) { + using IType = typename Traits::IType; + using IType2 = typename Traits::IType2; + using IType2x3D = typename Traits::IType2x3D; + using OType2xt3D = typename Traits::OType2xt3D; + using ScalesTypeTr2D = typename Traits::ScalesTypeTr2D; + using scaling_coeff_type = typename ScalingCoefficientType::type; + + constexpr int TILE_DIM_X = Traits::TILE_DIM_X; + constexpr int SCALES_PER_TILE_Y = Traits::SCALES_PER_TILE_Y; + + const auto &sIn2x = *reinterpret_cast(sIn_ptr); + auto &sOut_tr = *reinterpret_cast(sOut_tr_ptr); + auto &sSFcolwise = *reinterpret_cast(sSFcolwise_ptr); + + const int warp = threadIdx.x / THREADS_PER_WARP; + const int thread_lane = threadIdx.x % THREADS_PER_WARP; + + const int tid_Y_colwise = (thread_lane / 2 + warp) % SCALES_PER_TILE_Y; + const int tid_X_colwise = thread_lane; + + const int thread_offset_Y_colwise = tid_Y_colwise * NVFP4_SCALE_DIM; + const int thread_offset_X_colwise = tid_X_colwise * 2; + + const int in_thread_offset_Y = thread_offset_Y_colwise; + const int in_thread_offset_X = thread_offset_X_colwise / 2; + + const int out_tr_thread_offset_Y = thread_offset_X_colwise; + const int out_tr_thread_offset_X = thread_offset_Y_colwise / 2; + + const int scale_tr_offset_Y = (stage_X * TILE_DIM_X) + 2 * tid_X_colwise; + const int scale_tr_offset_X = (stage_Y * SCALES_PER_TILE_Y) + tid_Y_colwise; + + __align__(8) IType rIn[2][NVFP4_SCALE_DIM]; + // Read (cache) a pair of input elements (S2R). Find NVFP4-block AMAX. + IType2 thread_amax_2x = {static_cast(0.0f), static_cast(0.0f)}; +#pragma unroll + for (int i = 0; i < NVFP4_SCALE_DIM; ++i) { + const IType2 elt_pair = + ptx::ld_shared_b32(&sIn2x[buff_in][in_thread_offset_Y + i][in_thread_offset_X]); + rIn[0][i] = elt_pair.x; + rIn[1][i] = elt_pair.y; + ptx::abs_max_2x(thread_amax_2x, thread_amax_2x, elt_pair); + } + const float block_amax[2] = {static_cast(__habs(thread_amax_2x.x)), + static_cast(__habs(thread_amax_2x.y))}; +#pragma unroll + for (int w = 0; w < 2; ++w) { + float S_enc_colwise_block = S_enc_colwise; + if constexpr (ROW_SCALED_NVFP4) { + const size_t col_idx = col_offset + stage_X * TILE_DIM_X + thread_offset_X_colwise + w; + S_enc_colwise_block = + col_idx < cols ? core::compute_global_encode_scaling_factor_FP4(amax_colwise_ptr[col_idx]) + : 1.0f; + } + const nvfp4_scale_t S_dec_b_fp8 = + quantization_and_transposition_SF::compute_decoding_scaling_factor(block_amax[w], + S_enc_colwise_block); + + // Store scaling factors to SMEM buffer (R2S). + sSFcolwise[scale_tr_offset_Y + w][scale_tr_offset_X] = S_dec_b_fp8; + + const scaling_coeff_type SFcoefficient = + core::compute_scaling_coefficient(S_dec_b_fp8, S_enc_colwise_block); + + // Scale elements. + __align__(8) uint32_t rOut[NVFP4_SCALE_DIM / 8]; +#pragma unroll + for (int e = 0; e < NVFP4_SCALE_DIM / 8; ++e) { + const uint64_t elts03 = *reinterpret_cast(&rIn[w][8 * e]); + const uint64_t elts47 = *reinterpret_cast(&rIn[w][8 * e + 4]); + if constexpr (USE_STOCHASTIC_ROUNDING) { + const uint32_t rbits03 = core::get_rbits(rng, random_uint4, rnd_idx); + const uint32_t rbits47 = core::get_rbits(rng, random_uint4, rnd_idx); + rOut[e] = ptx::mul_cvt_bf16_to_fp4_8x_stochastic_rounding( + elts03, elts47, SFcoefficient, rbits03, rbits47); + } else { + rOut[e] = ptx::mul_cvt_bf16_to_fp4_8x_round_to_nearest(elts03, elts47, + SFcoefficient); + } + } + uint64_t &out_pack_16x = *reinterpret_cast(rOut); + ptx::st_shared_b64(&sOut_tr[buff_out_tr][out_tr_thread_offset_Y + w][out_tr_thread_offset_X], + out_pack_16x); + } +} + +template +__device__ __forceinline__ void colwise_scaling(const typename Traits::IType *__restrict__ sIn_ptr, + fp4e2m1x2 *__restrict__ sOut_tr_ptr, + nvfp4_scale_t *__restrict__ sSFcolwise_ptr, + const float S_enc_colwise, const int stage_Y, + const int stage_X, const int buff_in, + const int buff_out_tr, RngType &rng, + uint4 &random_uint4, int &rnd_idx) { + colwise_scaling( + sIn_ptr, sOut_tr_ptr, sSFcolwise_ptr, S_enc_colwise, stage_Y, stage_X, buff_in, buff_out_tr, + nullptr, 0, 0, rng, random_uint4, rnd_idx); +} + +template +__device__ __forceinline__ void rowwise_scaling( + const typename Traits::IType *__restrict__ sIn_ptr, fp4e2m1x2 *__restrict__ sOut_ptr, + nvfp4_scale_t *__restrict__ sSFrowwise_ptr, const float S_enc_rowwise, const int stage_Y, + const int stage_X, const int buff_in, const int buff_out, const float *amax_rowwise_ptr, + const size_t row_offset, const size_t rows, RngType &rng, uint4 &random_uint4, int &rnd_idx) { + using IType = typename Traits::IType; + using IType2 = typename Traits::IType2; + using IType3D = typename Traits::IType3D; + using OType2x3D = typename Traits::OType2x3D; + using ScalesType2D = typename Traits::ScalesType2D; + using scaling_coeff_type = typename ScalingCoefficientType::type; + + constexpr int TILE_DIM_Y = Traits::TILE_DIM_Y; + constexpr int ELTS_PER_THREAD = Traits::ELTS_PER_THREAD; + constexpr int PACK_SIZE = Traits::PACK_SIZE; + constexpr int WAVES = Traits::WAVES; + constexpr int THREADS_PER_BANK = Traits::THREADS_PER_BANK; + constexpr int THREADS_X_ROWWISE = Traits::THREADS_X_ROWWISE; + constexpr int THREADS_Y_ROWWISE = Traits::THREADS_Y_ROWWISE; + constexpr int THREADS_PER_SCALE_ROWWISE = Traits::THREADS_PER_SCALE_ROWWISE; + constexpr int SCALES_PER_TILE_X = Traits::SCALES_PER_TILE_X; + constexpr int ITERATIONS_NORMAL = Traits::ITERATIONS_NORMAL; + + const auto &sIn = *reinterpret_cast(sIn_ptr); + auto &sOut = *reinterpret_cast(sOut_ptr); + auto &sSFrowwise = *reinterpret_cast(sSFrowwise_ptr); + + const int thread_lane = threadIdx.x % THREADS_PER_WARP; + const int bank_group = thread_lane / THREADS_PER_BANK; + + const int tid_Y_rowwise = threadIdx.x / THREADS_X_ROWWISE; + const int tid_X_rowwise = threadIdx.x % THREADS_X_ROWWISE; + + const int thread_offset_Y_rowwise = tid_Y_rowwise; + const int thread_offset_X_rowwise = tid_X_rowwise * ELTS_PER_THREAD; + + const int SF_thread_offset_rowwise_Y = tid_Y_rowwise; + const int SF_thread_offset_rowwise_X = tid_X_rowwise / THREADS_PER_SCALE_ROWWISE; + + const bool SF_storing_thread = (tid_X_rowwise % THREADS_PER_SCALE_ROWWISE == 0); + + const int stage_rowwise_scales_offset_Y = SF_thread_offset_rowwise_Y + stage_Y * TILE_DIM_Y; + const int stage_rowwise_scales_offset_X = + SF_thread_offset_rowwise_X + stage_X * SCALES_PER_TILE_X; +#pragma unroll + for (int it = 0; it < ITERATIONS_NORMAL; ++it) { + const int it_offset_Y_rowwise = thread_offset_Y_rowwise + it * THREADS_Y_ROWWISE; + + __align__(16) IType2 rIn[WAVES][PACK_SIZE / 2]; + + // Read (cache) input elements (S2R). Find NVFP4-block AMAX. + IType2 thread_amax_2x = {static_cast(0.0f), static_cast(0.0f)}; +#pragma unroll + for (int w = 0; w < WAVES; ++w) { + const int swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % ELTS_PER_THREAD; + const int swizzled_thread_idx = thread_offset_X_rowwise + swizzled_group_idx; + + __uint128_t &elts_8x = *reinterpret_cast<__uint128_t *>(&rIn[w]); + elts_8x = ptx::ld_shared_b128(&sIn[buff_in][it_offset_Y_rowwise][swizzled_thread_idx]); +#pragma unroll + for (int e = 0; e < PACK_SIZE / 2; ++e) { + ptx::abs_max_2x(thread_amax_2x, thread_amax_2x, rIn[w][e]); + } + } + const float block_amax = get_amax_of_pair(thread_amax_2x); + + nvfp4_scale_t S_dec_b_fp8; + scaling_coeff_type SFcoefficient; + if constexpr (ROW_SCALED_NVFP4) { + const size_t row_idx = row_offset + stage_Y * TILE_DIM_Y + it_offset_Y_rowwise; + const float S_enc_rowwise_block = + row_idx < rows ? core::compute_global_encode_scaling_factor_FP4(amax_rowwise_ptr[row_idx]) + : 1.0f; + S_dec_b_fp8 = quantization_and_transposition_SF::compute_decoding_scaling_factor( + block_amax, S_enc_rowwise_block); + SFcoefficient = + core::compute_scaling_coefficient(S_dec_b_fp8, S_enc_rowwise_block); + } else { + S_dec_b_fp8 = quantization_and_transposition_SF::compute_decoding_scaling_factor( + block_amax, S_enc_rowwise); + SFcoefficient = + core::compute_scaling_coefficient(S_dec_b_fp8, S_enc_rowwise); + } + + // Store scaling factors to SMEM buffer (R2S). + if (SF_storing_thread) { + const int scales_offset_Y = stage_rowwise_scales_offset_Y + it * THREADS_Y_ROWWISE; + const int scales_offset_X = stage_rowwise_scales_offset_X; + sSFrowwise[scales_offset_Y][scales_offset_X] = S_dec_b_fp8; + } + +// Scale elements. +#pragma unroll + for (int w = 0; w < WAVES; ++w) { + const uint64_t elts03 = *reinterpret_cast(&rIn[w][0]); + const uint64_t elts47 = *reinterpret_cast(&rIn[w][2]); + + uint32_t out_x8; + if constexpr (USE_STOCHASTIC_ROUNDING) { + const uint32_t rbits03 = core::get_rbits(rng, random_uint4, rnd_idx); + const uint32_t rbits47 = core::get_rbits(rng, random_uint4, rnd_idx); + out_x8 = ptx::mul_cvt_bf16_to_fp4_8x_stochastic_rounding( + elts03, elts47, SFcoefficient, rbits03, rbits47); + } else { + out_x8 = ptx::mul_cvt_bf16_to_fp4_8x_round_to_nearest(elts03, elts47, + SFcoefficient); + } + + const int swizzled_group_idx = ((w + bank_group) * PACK_SIZE) % ELTS_PER_THREAD; + const int swizzled_idx = (swizzled_group_idx + thread_offset_X_rowwise) / 2; + ptx::st_shared_b32(&sOut[buff_out][it_offset_Y_rowwise][swizzled_idx], out_x8); + } + } +} + +template +__device__ __forceinline__ void rowwise_scaling(const typename Traits::IType *__restrict__ sIn_ptr, + fp4e2m1x2 *__restrict__ sOut_ptr, + nvfp4_scale_t *__restrict__ sSFrowwise_ptr, + const float S_enc_rowwise, const int stage_Y, + const int stage_X, const int buff_in, + const int buff_out, RngType &rng, + uint4 &random_uint4, int &rnd_idx) { + rowwise_scaling( + sIn_ptr, sOut_ptr, sSFrowwise_ptr, S_enc_rowwise, stage_Y, stage_X, buff_in, buff_out, + nullptr, 0, 0, rng, random_uint4, rnd_idx); +} + +#endif // FP4_TYPE_SUPPORTED + +} // namespace tuned_1D_scaling_common +} // namespace nvfp4 +} // namespace dispatch +} // namespace transformer_engine + +#endif // TRANSFORMER_ENGINE_SCALING_NVFP4_TUNED_1D_CUH_