From 2044d1d59c415bc86ddc400ec334ca4b4712b92b Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Mon, 12 Jan 2026 09:42:03 +0000 Subject: [PATCH 01/23] Rename the parameters of load_interleaved_pk_type and load_and_convert_tile --- .../ck_tile/ops/common/load_and_convert_tile.hpp | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/projects/composablekernel/include/ck_tile/ops/common/load_and_convert_tile.hpp b/projects/composablekernel/include/ck_tile/ops/common/load_and_convert_tile.hpp index 0748c5fb49e8..9fb26e7638f8 100644 --- a/projects/composablekernel/include/ck_tile/ops/common/load_and_convert_tile.hpp +++ b/projects/composablekernel/include/ck_tile/ops/common/load_and_convert_tile.hpp @@ -14,11 +14,11 @@ template struct ConverterLoader { template - CK_TILE_DEVICE static void load_interleaved_pk_type(WarpTile& dst, const WarpWindow& src) + CK_TILE_DEVICE static void load_interleaved_pk_type(WarpTile& dst, const WarpWindow& src_window) { static_assert(WarpTile::get_thread_buffer_size() % UnaryOpSize == 0); constexpr index_t thread_buffer_size = WarpTile::get_thread_buffer_size() / UnaryOpSize; - const auto tmp = load_tile(src); + const auto src = load_tile(src_window); // NOTE: we rely on types packing neatly here using RawSrcType = typename SrcDataType::type; @@ -30,13 +30,13 @@ struct ConverterLoader const element_wise::PassThroughPack8 elementwise_op{}; elementwise_op(dst.get_thread_buffer().template get_as()(i), - tmp.get_thread_buffer().template get_as()[i]); + src.get_thread_buffer().template get_as()[i]); }); } }; template -CK_TILE_DEVICE void load_and_convert_tile(WarpTile& dst, const WarpWindow& src) +CK_TILE_DEVICE void load_and_convert_tile(WarpTile& dst, const WarpWindow& src_window) { using SrcDataType = typename WarpWindow::Base::DataType; using DstDataType = typename WarpTile::DataType; @@ -44,15 +44,16 @@ CK_TILE_DEVICE void load_and_convert_tile(WarpTile& dst, const WarpWindow& src) if constexpr(is_packed_type_v && !is_packed_type_v) { static_assert(!LoadTranspose, "LoadTranspose not supported with pk_int4_t or pk_fp4_t"); - ConverterLoader::load_interleaved_pk_type(dst, src); + ConverterLoader::load_interleaved_pk_type( + dst, src_window); } else if constexpr(LoadTranspose) { - load_tile_transpose(dst, src); + load_tile_transpose(dst, src_window); } else { - load_tile(dst, src); + load_tile(dst, src_window); } } From f8b027bc7acaa8612311bf24c4ca3833b72d7c74 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Mon, 26 Jan 2026 09:26:59 +0000 Subject: [PATCH 02/23] Add load_tile_transpose_convert for mixed precision transpose loading --- .../core/tensor/load_tile_transpose.hpp | 115 ++++++++++++++++++ .../unary_element_wise_operation.hpp | 38 ++++++ 2 files changed, 153 insertions(+) diff --git a/projects/composablekernel/include/ck_tile/core/tensor/load_tile_transpose.hpp b/projects/composablekernel/include/ck_tile/core/tensor/load_tile_transpose.hpp index 5f73d4934a1f..3636447f82f3 100644 --- a/projects/composablekernel/include/ck_tile/core/tensor/load_tile_transpose.hpp +++ b/projects/composablekernel/include/ck_tile/core/tensor/load_tile_transpose.hpp @@ -14,6 +14,7 @@ #include "ck_tile/core/container/statically_indexed_array.hpp" #include "ck_tile/core/numeric/math.hpp" #include "ck_tile/core/utility/type_traits.hpp" +#include "ck_tile/ops/elementwise/unary_element_wise_operation.hpp" namespace ck_tile { @@ -529,4 +530,118 @@ load_tile_transpose(const tile_window_with_static_distribution, + typename = std::enable_if_t::distr_encoding_valid, + Policy>> +CK_TILE_DEVICE void load_tile_transpose_convert_with_offset( + DistributedTensor_& out_tensor, + const tile_window_with_static_distribution& __restrict__ tile_window, + const index_t offset, + number = {}) +{ + using SrcDataType = typename BottomTensorView_::DataType; + using DstDataType = typename DistributedTensor_::DataType; + + auto trans_tensor = tile_window.template load_transpose_with_offset(offset); + constexpr auto input_distr = TileDistribution_{}; + constexpr auto output_distr = typename DistributedTensor_::StaticTileDistribution{}; + + constexpr auto y_in_desc = input_distr.get_ys_to_d_descriptor(); + constexpr auto y_out_desc = output_distr.get_ys_to_d_descriptor(); + + constexpr auto y_in_lengths = to_sequence(y_in_desc.get_lengths()); + constexpr auto y_out_lengths = to_sequence(y_out_desc.get_lengths()); + + constexpr auto y_in_element_space_size = y_in_desc.get_element_space_size(); + constexpr auto y_out_element_space_size = y_out_desc.get_element_space_size(); + + // For mixed precision: element space size must be the same (total bytes match) + static_assert(y_in_element_space_size == y_out_element_space_size, + "For mixed precision transpose, input and output element space size must match!"); + + // Ensure total element counts are consistent and divisible by the input vector length. + constexpr index_t total_elems_in = + reduce_on_sequence(y_in_lengths, multiplies<>{}, number<1>{}); + constexpr index_t total_elems_out = + reduce_on_sequence(y_out_lengths, multiplies<>{}, number<1>{}); + static_assert(total_elems_in == total_elems_out, + "For mixed precision transpose, input/output element counts must match!"); + static_assert(total_elems_in % number{} == 0, + "Input vector length must evenly divide total elements."); + + constexpr index_t num_of_access = total_elems_in / number{}; + + // Read as input type, convert to output type + using SrcDataVec = ext_vector_t{}>; + using DstDataVec = ext_vector_t{}>; + static_for<0, num_of_access, 1>{}([&](auto i) { + static_assert(number{} == 8, "Only PassThroughPack8 is supported for now."); + const element_wise::PassThroughPack8 elementwise_op{}; + + elementwise_op(out_tensor.get_thread_buffer().template get_as()(i), + trans_tensor.get_thread_buffer().template get_as()[i]); + }); +} + +/** + * @brief Mixed-precision transpose load with zero offset. + * + * Convenience wrapper for load_tile_transpose_convert_with_offset with offset=0. + */ +template < + typename DistributedTensor_, + typename BottomTensorView_, + typename WindowLengths_, + typename TileDistribution_, + index_t NumCoord_, + index_t UnaryOpSize_, + typename Policy = DefaultTranspose, + typename = std::enable_if_t::distr_encoding_valid, + Policy>> +CK_TILE_DEVICE void load_tile_transpose_convert( + DistributedTensor_& out_tensor, + const tile_window_with_static_distribution& __restrict__ tile_window, + number = {}) +{ + load_tile_transpose_convert_with_offset(out_tensor, tile_window, 0, number{}); +} + } // namespace ck_tile diff --git a/projects/composablekernel/include/ck_tile/ops/elementwise/unary_element_wise_operation.hpp b/projects/composablekernel/include/ck_tile/ops/elementwise/unary_element_wise_operation.hpp index 4ad699629c02..eea1eb6acd62 100644 --- a/projects/composablekernel/include/ck_tile/ops/elementwise/unary_element_wise_operation.hpp +++ b/projects/composablekernel/include/ck_tile/ops/elementwise/unary_element_wise_operation.hpp @@ -447,6 +447,15 @@ CK_TILE_HOST_DEVICE bf16x8_t fp8x8_to_bf16x8_scale(const fp8x8_t& src, const flo return y; } +CK_TILE_HOST_DEVICE fp8x8_t bf16x8_to_fp8x8_scale(const bf16x8_t& src, const float& scale) +{ + fp8x8_t y; + static_for<0, 8, 1>{}([&](auto i) { + y[i.value] = type_convert(type_convert(src[i.value]) * scale); + }); + return y; +} + CK_TILE_HOST_DEVICE fp16x8_t fp8x8_to_fp16x8_scale(const fp8x8_t& src, const float& scale) { fp16x8_t y; @@ -491,6 +500,15 @@ CK_TILE_HOST_DEVICE fp16x8_t fp8x8_to_fp16x8_scale(const fp8x8_t& src, const flo return y; } +CK_TILE_HOST_DEVICE fp8x8_t fp16x8_to_fp8x8_scale(const fp16x8_t& src, const float& scale) +{ + fp8x8_t y; + static_for<0, 8, 1>{}([&](auto i) { + y[i.value] = type_convert(type_convert(src[i.value]) * scale); + }); + return y; +} + CK_TILE_HOST_DEVICE fp16x8_t bf8x8_to_fp16x8_scale(const bf8x8_t& src, const float& scale) { fp16x8_t y; @@ -620,12 +638,32 @@ struct PassThroughPack8 template CK_TILE_HOST_DEVICE void operator()(Y& y, const X& x) const; + CK_TILE_HOST_DEVICE constexpr void operator()(fp16x8_t& y, const fp8x8_t& x) const + { + y = fp8x8_to_fp16x8_scale(x, 1.0f); + } + + CK_TILE_HOST_DEVICE constexpr void operator()(fp8x8_t& y, const fp16x8_t& x) const + { + y = fp16x8_to_fp8x8_scale(x, 1.0f); + } + CK_TILE_HOST_DEVICE constexpr void operator()(fp16x8_t& y, const pk_int4x4_t& x) const { y.lo = i4_to_half4(bit_cast(x)); y.hi = i4_to_half4(bit_cast(x) >> 8); } + CK_TILE_HOST_DEVICE constexpr void operator()(bf16x8_t& y, const fp8x8_t& x) const + { + y = fp8x8_to_bf16x8_scale(x, 1.0f); + } + + CK_TILE_HOST_DEVICE constexpr void operator()(fp8x8_t& y, const bf16x8_t& x) const + { + y = bf16x8_to_fp8x8_scale(x, 1.0f); + } + CK_TILE_HOST_DEVICE constexpr void operator()(bf16x8_t& y, const pk_int4x4_t& x) const { y.lo = i4_to_bhalf4(bit_cast(x)); From 89c358f7d47a8c84966693379fff3da044e099d3 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Wed, 12 Nov 2025 09:04:15 +0000 Subject: [PATCH 03/23] Add and use load_with_type_convert --- .../ops/common/load_and_convert_tile.hpp | 49 +++++++++++++++---- 1 file changed, 39 insertions(+), 10 deletions(-) diff --git a/projects/composablekernel/include/ck_tile/ops/common/load_and_convert_tile.hpp b/projects/composablekernel/include/ck_tile/ops/common/load_and_convert_tile.hpp index 9fb26e7638f8..6dc2ff35c0a1 100644 --- a/projects/composablekernel/include/ck_tile/ops/common/load_and_convert_tile.hpp +++ b/projects/composablekernel/include/ck_tile/ops/common/load_and_convert_tile.hpp @@ -10,12 +10,16 @@ namespace ck_tile { -template +template struct ConverterLoader { template CK_TILE_DEVICE static void load_interleaved_pk_type(WarpTile& dst, const WarpWindow& src_window) { + static_assert(!LoadTranspose, "LoadTranspose not supported with pk_int4_t or pk_fp4_t"); static_assert(WarpTile::get_thread_buffer_size() % UnaryOpSize == 0); constexpr index_t thread_buffer_size = WarpTile::get_thread_buffer_size() / UnaryOpSize; const auto src = load_tile(src_window); @@ -33,6 +37,36 @@ struct ConverterLoader src.get_thread_buffer().template get_as()[i]); }); } + + template + CK_TILE_DEVICE static void load_with_type_convert(WarpTile& dst, const WarpWindow& src_window) + { + if constexpr(LoadTranspose) + { + if constexpr(std::is_same_v) + { + load_tile_transpose(dst, src_window); + } + else + { + load_tile_transpose_convert(dst, src_window, number{}); + } + } + else + { + if constexpr(std::is_same_v) + { + load_tile(dst, src_window); + } + else + { + auto tmp = load_tile(src_window); + sweep_tile([&](auto i) { + dst(i) = type_convert(type_convert(tmp(i))); + }); + } + } + } }; template @@ -43,18 +77,13 @@ CK_TILE_DEVICE void load_and_convert_tile(WarpTile& dst, const WarpWindow& src_w if constexpr(is_packed_type_v && !is_packed_type_v) { - static_assert(!LoadTranspose, "LoadTranspose not supported with pk_int4_t or pk_fp4_t"); - ConverterLoader::load_interleaved_pk_type( - dst, src_window); - } - else if constexpr(LoadTranspose) - { - load_tile_transpose(dst, src_window); + ConverterLoader:: + load_interleaved_pk_type(dst, src_window); } else { - load_tile(dst, src_window); + ConverterLoader:: + load_with_type_convert(dst, src_window); } } - } // namespace ck_tile From ab97b13bf62ceaf687f33cf8304cb4386e38f6d0 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 17 Mar 2026 09:04:33 +0000 Subject: [PATCH 04/23] Add FillIdentity for host tensors --- .../include/ck_tile/host/fill.hpp | 38 +++++++++++++++++++ 1 file changed, 38 insertions(+) diff --git a/projects/composablekernel/include/ck_tile/host/fill.hpp b/projects/composablekernel/include/ck_tile/host/fill.hpp index 44d191303387..2fe7470f353d 100644 --- a/projects/composablekernel/include/ck_tile/host/fill.hpp +++ b/projects/composablekernel/include/ck_tile/host/fill.hpp @@ -496,6 +496,44 @@ struct FillConstant } }; +template +struct FillIdentity +{ + std::size_t rows_{0}; + std::size_t cols_{0}; + T zero_{type_convert(0)}; + T one_{type_convert(1)}; + + template + void operator()(ForwardIter first, ForwardIter last) const + { + if(rows_ == 0 || cols_ == 0 || first == last) + return; + + const auto total = static_cast(std::distance(first, last)); + if(total < rows_ * cols_) + { + throw std::runtime_error("FillIdentity requires range size >= rows_ * cols_."); + } + + std::fill(first, first + rows_ * cols_, zero_); + + const auto min_dim = std::min(rows_, cols_); + for(std::size_t i = 0; i < min_dim; ++i) + *(first + i * cols_ + i) = one_; + } + + template + auto operator()(ForwardRange&& range) const + -> std::void_t()( + std::begin(std::forward(range)), + std::end(std::forward(range))))> + { + (*this)(std::begin(std::forward(range)), + std::end(std::forward(range))); + } +}; + //---------------------------------------------------------------------------------------------- /// @brief Transforms given input to fit 2:4 structured sparsity pattern so /// every subgroup of 4 elements contain at most 2 non-zero elements From 1eb1a031b68c02bf3dca5ba23c2e86877e247a5f Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 13 Jan 2026 09:04:21 +0000 Subject: [PATCH 05/23] Add test_load_and_convert_tile --- .../test/ck_tile/CMakeLists.txt | 1 + .../load_and_convert_tile/CMakeLists.txt | 9 + .../ck_tile/load_and_convert_tile/kernel.hpp | 247 ++++++++++++++++++ .../test_load_and_convert_tile.cpp | 139 ++++++++++ 4 files changed, 396 insertions(+) create mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt create mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp create mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp diff --git a/projects/composablekernel/test/ck_tile/CMakeLists.txt b/projects/composablekernel/test/ck_tile/CMakeLists.txt index 320e5b1e91c9..a0e35261fdef 100644 --- a/projects/composablekernel/test/ck_tile/CMakeLists.txt +++ b/projects/composablekernel/test/ck_tile/CMakeLists.txt @@ -70,3 +70,4 @@ add_subdirectory(gemm_tile_engine) add_subdirectory(pooling) add_subdirectory(grouped_conv) add_subdirectory(gemm_streamk_tile_engine) +add_subdirectory(load_and_convert_tile) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt new file mode 100644 index 000000000000..e07fc05fc814 --- /dev/null +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt @@ -0,0 +1,9 @@ +# Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +# SPDX-License-Identifier: MIT + +set(LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) +if(GPU_TARGETS MATCHES "gfx9") + add_gtest_executable(test_load_and_convert_tile test_load_and_convert_tile.cpp) + list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8 -fverbose-asm --save-temps -Wno-gnu-line-marker) + target_compile_options(test_load_and_convert_tile PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) +endif() diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp new file mode 100644 index 000000000000..bc8d6b582e41 --- /dev/null +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp @@ -0,0 +1,247 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT + +#pragma once + +#include "ck_tile/core.hpp" +#include "ck_tile/core/algorithm/coordinate_transform.hpp" +#include "ck_tile/ops/common.hpp" +#include "ck_tile/ops/gemm/warp/warp_gemm.hpp" + +namespace ck_tile { + +template +struct LoadAndConvertShape +{ + static constexpr index_t Block_M = BlockTile::at(number<0>{}); + static constexpr index_t Block_N = BlockTile::at(number<1>{}); + static constexpr index_t Block_K = BlockTile::at(number<2>{}); + + static constexpr index_t Warp_M = WarpTile::at(number<0>{}); + static constexpr index_t Warp_N = WarpTile::at(number<1>{}); + static constexpr index_t Warp_K = WarpTile::at(number<2>{}); + + static constexpr index_t Vector_N = Vector::at(number<1>{}); + + static constexpr index_t WarpPerBlock_M = BlockWarps::at(number<0>{}); + static constexpr index_t WarpPerBlock_N = BlockWarps::at(number<1>{}); + static constexpr index_t WarpPerBlock_K = BlockWarps::at(number<2>{}); + + static constexpr index_t Repeat_M = Block_M / (WarpPerBlock_M * Warp_M); + static constexpr index_t Repeat_N = Block_N / (WarpPerBlock_N * Warp_N); + static constexpr index_t Repeat_K = Block_K / (WarpPerBlock_K * Warp_K); + + static constexpr index_t BlockSize = + ck_tile::get_warp_size() * reduce_on_sequence(BlockWarps{}, multiplies<>{}, number<1>{}); +}; + +template +struct LoadAndConvertProblem +{ + using XDataType = remove_cvref_t; + using YDataType = remove_cvref_t; + using BlockShape = remove_cvref_t; + using LoadTranspose = remove_cvref_t; +}; + +template +struct LoadAndConvertKernel +{ + using Problem = ck_tile::remove_cvref_t; + using XDataType = ck_tile::remove_cvref_t; + using YDataType = ck_tile::remove_cvref_t; + using LoadTranspose = ck_tile::remove_cvref_t; + + static constexpr index_t kBlockSize = Problem::BlockShape::BlockSize; + + template + static constexpr auto get_warp_dstr_encoding() + { + using S = typename Problem::BlockShape; + + if constexpr(NumAccess == 1) + return tile_distribution_encoding, + tuple, sequence<2, S::Vector_N>>, + tuple>, + tuple>, + sequence<2>, + sequence<1>>{}; + else + return tile_distribution_encoding< + sequence<>, + tuple, sequence>, + tuple>, + tuple>, + sequence<2, 2>, + sequence<0, 2>>{}; + } + + template + CK_TILE_DEVICE static constexpr auto GetVectorSize() + { + return DS_READ_TR_SIZE() / sizeof(DataType); + } + + template + CK_TILE_DEVICE static constexpr auto MakeDRAMDistribution() + { + using S = typename Problem::BlockShape; + constexpr index_t thread_elements = S::Warp_N * S::Warp_K / get_warp_size(); + constexpr index_t NumAccess = + LoadTranspose::value ? thread_elements / GetVectorSize() : 1; + + constexpr auto a_block_outer_dstr_encode = tile_distribution_encoding< + sequence, + tuple, sequence>, + tuple>, + tuple>, + sequence<1, 2>, + sequence<0, 0>>{}; + + constexpr auto a_block_dstr_encode = detail::make_embed_tile_distribution_encoding( + a_block_outer_dstr_encode, get_warp_dstr_encoding()); + + return make_static_tile_distribution(a_block_dstr_encode); + } + + template + CK_TILE_DEVICE static constexpr auto MakeDRAMTransposedDistribution() + { + return make_static_tile_distribution( + typename InputTileDistributionTraits< + typename decltype(MakeDRAMDistribution())::DstrEncode, + DataType>::TransposedDstrEncode{}); + } + + CK_TILE_DEVICE void + operator()(const XDataType* a, YDataType* c, index_t M, index_t N, index_t K) const + { + using S = typename Problem::BlockShape; + + const index_t kMPerBlock = S::WarpPerBlock_M * S::Repeat_M * S::Block_M; + const index_t kNPerBlock = S::WarpPerBlock_N * S::Repeat_N * S::Block_N; + + constexpr auto block_dims = make_tuple(number{}, number{}); + constexpr auto block_strides = make_tuple(number<1>{}, number{}); + const index_t num_blocks_n = N / kNPerBlock; + const index_t block_m = get_block_id() / num_blocks_n; + + const index_t m_block_base = block_m * kMPerBlock; + + // LDS buffer + __shared__ XDataType a_lds[kMPerBlock * S::Block_K]; + + auto a_lds_write_view = make_naive_tensor_view( + a_lds, block_dims, block_strides, number<1>{}, number<1>{}); + + auto a_block_lds_write_window = make_tile_window(a_lds_write_view, block_dims, {0, 0}); + + auto a_block_lds_read_window = [&] { + if constexpr(LoadTranspose::value) + { + constexpr auto block_dims_t = + make_tuple(number{}, number{}); + constexpr auto block_strides_t = make_tuple(number{}, number<1>{}); + + auto view = make_naive_tensor_view( + a_lds, + block_dims_t, + block_strides_t, + number()>{}, + number<1>{}); + + return make_tile_window( + view, block_dims_t, {0, 0}, MakeDRAMTransposedDistribution()); + } + else + { + auto view = make_naive_tensor_view( + a_lds, block_dims, block_strides, number<1>{}, number<1>{}); + + return make_tile_window( + view, block_dims, {0, 0}, MakeDRAMDistribution()); + } + }(); + + // Input tensor + const auto a_tensor = make_naive_tensor_view( + a, make_tuple(M, K), make_tuple(1, M), number<1>{}, number<1>{}); + + auto a_block_window = make_tile_window( + a_tensor, block_dims, {m_block_base, 0}, MakeDRAMDistribution()); + + // Output tensor + auto c_tensor = [&]() { + if constexpr(LoadTranspose::value && !std::is_same_v) + { + // Similar to QuantGemmKernel with PermuteB: reinterpret the output logical layout + // via descriptor transform so YDataType distribution can be used without row-group + // permutation artifacts in mixed-precision transpose loads. + using TransposeGroupType = std:: + conditional_t<(sizeof(XDataType) >= sizeof(YDataType)), XDataType, YDataType>; + constexpr index_t thread_elements = S::Warp_N * S::Warp_K / get_warp_size(); + constexpr index_t n_group_0 = thread_elements / GetVectorSize(); + constexpr index_t n_group_1 = 2; + constexpr index_t n_group_2 = S::Vector_N / n_group_0; + constexpr index_t n_perm_group = n_group_0 * n_group_1 * n_group_2; + + static_assert(n_group_0 > 0, "Invalid derived transpose grouping factor"); + static_assert(S::Vector_N % n_group_0 == 0, + "Vector_N must be divisible by derived grouping factor"); + + const auto c_m_n_desc = make_naive_tensor_descriptor( + make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); + + const auto c_m_n0_b0_b1_n4_desc = transform_tensor_descriptor( + c_m_n_desc, + make_tuple(make_pass_through_transform(M), + make_unmerge_transform(make_tuple(N / n_perm_group, + number{}, + number{}, + number{}))), + make_tuple(sequence<0>{}, sequence<1>{}), + make_tuple(sequence<0>{}, sequence<1, 2, 3, 4>{})); + + const auto c_perm_m_n_desc = transform_tensor_descriptor( + c_m_n0_b0_b1_n4_desc, + make_tuple(make_pass_through_transform(M), + make_merge_transform(make_tuple(N / n_perm_group, + number{}, + number{}, + number{}))), + make_tuple(sequence<0>{}, sequence<1, 3, 2, 4>{}), + make_tuple(sequence<0>{}, sequence<1>{})); + + return make_tensor_view(c, c_perm_m_n_desc); + } + else + { + return make_naive_tensor_view( + c, make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); + } + }(); + + auto c_block_window = make_tile_window( + c_tensor, block_dims, {m_block_base, 0}, MakeDRAMDistribution()); + + const index_t num_k_loops = K / S::Block_K; + for(index_t k_iter = 0; k_iter < num_k_loops; ++k_iter) + { + auto dram_tile = load_tile(a_block_window); + store_tile(a_block_lds_write_window, dram_tile); + block_sync_lds(); + + decltype(load_tile(c_block_window)) c_tile; + load_and_convert_tile<8, LoadTranspose::value>(c_tile, a_block_lds_read_window); + store_tile(c_block_window, c_tile); + + if(k_iter < num_k_loops - 1) + { + move_tile_window(a_block_window, {0, S::Block_K}); + move_tile_window(c_block_window, {0, S::Block_K}); + } + } + } +}; + +} // namespace ck_tile diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp new file mode 100644 index 000000000000..460f0b0aba1f --- /dev/null +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp @@ -0,0 +1,139 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT + +#include + +#include "ck_tile/host.hpp" +#include "ck_tile/ops/common.hpp" +#include "kernel.hpp" + +#define MONOTONIC_SEQUENCE 0 +#define IDENTITY 1 +#define UNIFORM_DISTRIBUTION 2 +#define MATRIX_TYPE UNIFORM_DISTRIBUTION + +#define PRINT_MATRICES 0 + +// Helper to print matrix (for debugging) +template +void print_matrix(const ck_tile::HostTensor& mat, const std::string& name = "Matrix") +{ + const auto lens = mat.get_lengths(); + assert(len(lens) == 2); + const ck_tile::index_t rows = lens[0]; + const ck_tile::index_t cols = lens[1]; + const ck_tile::index_t limit = 32; + + std::cout << name << " (" << rows << "×" << cols << "):\n"; + for(ck_tile::index_t i = 0; i < std::min(rows, ck_tile::index_t(limit)); ++i) + { + for(ck_tile::index_t j = 0; j < std::min(cols, ck_tile::index_t(limit)); ++j) + { +#if MATRIX_TYPE == MONOTONIC_SEQUENCE + std::cout << std::setw(3) << std::setprecision(3) +#else + std::cout << std::fixed << std::setprecision(2) << std::setw(6) +#endif + << ck_tile::type_convert(mat(i, j)) << " "; + } + if(cols > limit) + std::cout << "..."; + std::cout << "\n"; + } + if(rows > limit) + std::cout << "...\n"; + std::cout << "\n"; +} + +template +class TestLoadAndConvert : public ::testing::Test +{ + public: + using XDataType = std::tuple_element_t<0, Tuple>; + using YDataType = std::tuple_element_t<1, Tuple>; + using LoadTranspose = std::tuple_element_t<2, Tuple>; + + protected: + void RunTest() + { + constexpr ck_tile::index_t M = 32; + constexpr ck_tile::index_t N = 32; + constexpr ck_tile::index_t K = 32; + + ck_tile::HostTensor h_a({M, K}); + ck_tile::HostTensor h_c({M, K}); +#if MATRIX_TYPE == MONOTONIC_SEQUENCE + ck_tile::HostTensor h_a_tmp({M, K}); + ck_tile::FillMonotonicSeq{0.0, 0.1}(h_a_tmp); + ck_tile::reference_unary_elementwise( + h_a_tmp, h_a, [](const auto& x) { return x; }); +#elif MATRIX_TYPE == IDENTITY + ck_tile::FillIdentity{M, K}(h_a); +#else + ck_tile::FillUniformDistributionIntegerValue{-5.0, 5.0, 11939}(h_a); +#endif + + ck_tile::DeviceMem d_a(h_a.get_element_space_size_in_bytes()); + ck_tile::DeviceMem d_c(h_c.get_element_space_size_in_bytes()); + + d_a.ToDevice(h_a.data()); + d_c.ToDevice(h_c.data()); + + using BlockWarps = ck_tile::sequence<1, 1, 1>; + using BlockTile = ck_tile::sequence<32, 32, 16>; + using WarpTile = ck_tile::sequence<32, 32, 16>; + using Vector = ck_tile::sequence<1, 8>; + + using Shape = ck_tile::LoadAndConvertShape; + using Problem = ck_tile::LoadAndConvertProblem; + using Kernel = ck_tile::LoadAndConvertKernel; + + constexpr ck_tile::index_t block_size = Kernel::kBlockSize; + const ck_tile::index_t grid_size = ck_tile::integer_divide_ceil(M, Shape::Block_M) * + ck_tile::integer_divide_ceil(N, Shape::Block_N); + + launch_kernel(ck_tile::stream_config{nullptr, true}, + make_kernel(Kernel{}, + dim3(grid_size), + dim3(block_size), + 0, + static_cast(d_a.GetDeviceBuffer()), + static_cast(d_c.GetDeviceBuffer()), + M, + N, + K)); + + ck_tile::hip_check_error(hipDeviceSynchronize()); + d_c.FromDevice(h_c.data()); + ck_tile::HostTensor h_a_ref({M, K}); + ck_tile::reference_unary_elementwise( + h_a, h_a_ref, [](const auto& x) { return x; }); + bool pass = ck_tile::check_err(h_c, h_a_ref); + +#if PRINT_MATRICES + print_matrix(h_a, "Matrix A"); + print_matrix(h_c, "Matrix C"); +#endif + + EXPECT_TRUE(pass); + } +}; + +using TestTypes = ::testing::Types, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple>; + +TYPED_TEST_SUITE(TestLoadAndConvert, TestTypes); + +TYPED_TEST(TestLoadAndConvert, Test) { this->RunTest(); } From a6309fa632fb71f84087f6b3dbcfa44266072a79 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Thu, 19 Mar 2026 08:48:06 +0000 Subject: [PATCH 06/23] In LoadAndConvertKernel, modify the input tensor view instead of the output tensor view in the mixed precision case --- .../ck_tile/load_and_convert_tile/kernel.hpp | 30 +++++++++---------- 1 file changed, 15 insertions(+), 15 deletions(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp index bc8d6b582e41..f283d44e8ff9 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp @@ -164,14 +164,7 @@ struct LoadAndConvertKernel }(); // Input tensor - const auto a_tensor = make_naive_tensor_view( - a, make_tuple(M, K), make_tuple(1, M), number<1>{}, number<1>{}); - - auto a_block_window = make_tile_window( - a_tensor, block_dims, {m_block_base, 0}, MakeDRAMDistribution()); - - // Output tensor - auto c_tensor = [&]() { + const auto a_tensor = [&]() { if constexpr(LoadTranspose::value && !std::is_same_v) { // Similar to QuantGemmKernel with PermuteB: reinterpret the output logical layout @@ -189,11 +182,11 @@ struct LoadAndConvertKernel static_assert(S::Vector_N % n_group_0 == 0, "Vector_N must be divisible by derived grouping factor"); - const auto c_m_n_desc = make_naive_tensor_descriptor( + const auto a_m_n_desc = make_naive_tensor_descriptor( make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); - const auto c_m_n0_b0_b1_n4_desc = transform_tensor_descriptor( - c_m_n_desc, + const auto a_m_n0_b0_b1_n4_desc = transform_tensor_descriptor( + a_m_n_desc, make_tuple(make_pass_through_transform(M), make_unmerge_transform(make_tuple(N / n_perm_group, number{}, @@ -202,8 +195,8 @@ struct LoadAndConvertKernel make_tuple(sequence<0>{}, sequence<1>{}), make_tuple(sequence<0>{}, sequence<1, 2, 3, 4>{})); - const auto c_perm_m_n_desc = transform_tensor_descriptor( - c_m_n0_b0_b1_n4_desc, + const auto a_perm_m_n_desc = transform_tensor_descriptor( + a_m_n0_b0_b1_n4_desc, make_tuple(make_pass_through_transform(M), make_merge_transform(make_tuple(N / n_perm_group, number{}, @@ -212,15 +205,22 @@ struct LoadAndConvertKernel make_tuple(sequence<0>{}, sequence<1, 3, 2, 4>{}), make_tuple(sequence<0>{}, sequence<1>{})); - return make_tensor_view(c, c_perm_m_n_desc); + return make_tensor_view(a, a_perm_m_n_desc); } else { return make_naive_tensor_view( - c, make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); + a, make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); } }(); + auto a_block_window = make_tile_window( + a_tensor, block_dims, {m_block_base, 0}, MakeDRAMDistribution()); + + // Output tensor + auto c_tensor = make_naive_tensor_view( + c, make_tuple(M, K), make_tuple(1, M), number<1>{}, number<1>{}); + auto c_block_window = make_tile_window( c_tensor, block_dims, {m_block_base, 0}, MakeDRAMDistribution()); From cf86511971427d6e2e5de15f895b6f13793340ed Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Fri, 20 Mar 2026 11:31:02 +0000 Subject: [PATCH 07/23] Only run test-load-tile-transpose on gfx950 --- .../test/ck_tile/load_and_convert_tile/CMakeLists.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt index e07fc05fc814..ce6b1061bf00 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt @@ -2,7 +2,7 @@ # SPDX-License-Identifier: MIT set(LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) -if(GPU_TARGETS MATCHES "gfx9") +if(GPU_TARGETS MATCHES "gfx950") add_gtest_executable(test_load_and_convert_tile test_load_and_convert_tile.cpp) list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8 -fverbose-asm --save-temps -Wno-gnu-line-marker) target_compile_options(test_load_and_convert_tile PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) From 1918d45f4de5dd4085b575271d4f7a0c0b1800c7 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Mon, 23 Mar 2026 15:50:14 +0000 Subject: [PATCH 08/23] Remove an unnecessary seed in a call to FillUniformDistributionIntegerValue --- .../load_and_convert_tile/test_load_and_convert_tile.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp index 460f0b0aba1f..219149e145a0 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp @@ -70,7 +70,7 @@ class TestLoadAndConvert : public ::testing::Test #elif MATRIX_TYPE == IDENTITY ck_tile::FillIdentity{M, K}(h_a); #else - ck_tile::FillUniformDistributionIntegerValue{-5.0, 5.0, 11939}(h_a); + ck_tile::FillUniformDistributionIntegerValue{-5.0, 5.0}(h_a); #endif ck_tile::DeviceMem d_a(h_a.get_element_space_size_in_bytes()); From 5c2c2ae765dbefa6ec5e7050ae178b79632c452f Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 24 Mar 2026 09:26:36 +0000 Subject: [PATCH 09/23] Replace the matrix type macros with constants from an enum struct --- .../test_load_and_convert_tile.cpp | 39 ++++++++++++------- 1 file changed, 25 insertions(+), 14 deletions(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp index 219149e145a0..de4e0b1a0940 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp @@ -7,10 +7,15 @@ #include "ck_tile/ops/common.hpp" #include "kernel.hpp" -#define MONOTONIC_SEQUENCE 0 -#define IDENTITY 1 -#define UNIFORM_DISTRIBUTION 2 -#define MATRIX_TYPE UNIFORM_DISTRIBUTION +// Enum struct specifying what kind of test matrix to use +enum struct TestMatrixType +{ + MonotonicSequence = 0, + Identity = 1, + UniformDistribution = 2 +}; + +static constexpr auto matrix_type = TestMatrixType::UniformDistribution; #define PRINT_MATRICES 0 @@ -62,16 +67,22 @@ class TestLoadAndConvert : public ::testing::Test ck_tile::HostTensor h_a({M, K}); ck_tile::HostTensor h_c({M, K}); -#if MATRIX_TYPE == MONOTONIC_SEQUENCE - ck_tile::HostTensor h_a_tmp({M, K}); - ck_tile::FillMonotonicSeq{0.0, 0.1}(h_a_tmp); - ck_tile::reference_unary_elementwise( - h_a_tmp, h_a, [](const auto& x) { return x; }); -#elif MATRIX_TYPE == IDENTITY - ck_tile::FillIdentity{M, K}(h_a); -#else - ck_tile::FillUniformDistributionIntegerValue{-5.0, 5.0}(h_a); -#endif + + if constexpr(matrix_type == TestMatrixType::MonotonicSequence) + { + ck_tile::HostTensor h_a_tmp({M, K}); + ck_tile::FillMonotonicSeq{0.0, 0.1}(h_a_tmp); + ck_tile::reference_unary_elementwise( + h_a_tmp, h_a, [](const auto& x) { return x; }); + } + else if constexpr(matrix_type == TestMatrixType::Identity) + { + ck_tile::FillIdentity{M, K}(h_a); + } + else + { + ck_tile::FillUniformDistributionIntegerValue{-5.0, 5.0}(h_a); + } ck_tile::DeviceMem d_a(h_a.get_element_space_size_in_bytes()); ck_tile::DeviceMem d_c(h_c.get_element_space_size_in_bytes()); From f0e5e5750925f893567d08e908a64f53e302f9b5 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 24 Mar 2026 09:32:45 +0000 Subject: [PATCH 10/23] Add width and precision parameters to print_matrix --- .../test_load_and_convert_tile.cpp | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp index de4e0b1a0940..e90de113b241 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp @@ -21,7 +21,10 @@ static constexpr auto matrix_type = TestMatrixType::UniformDistribution; // Helper to print matrix (for debugging) template -void print_matrix(const ck_tile::HostTensor& mat, const std::string& name = "Matrix") +void print_matrix(const ck_tile::HostTensor& mat, + const std::string& name = "Matrix", + const int width = 3, + const int precision = 3) { const auto lens = mat.get_lengths(); assert(len(lens) == 2); @@ -34,11 +37,7 @@ void print_matrix(const ck_tile::HostTensor& mat, const std::string& name = " { for(ck_tile::index_t j = 0; j < std::min(cols, ck_tile::index_t(limit)); ++j) { -#if MATRIX_TYPE == MONOTONIC_SEQUENCE - std::cout << std::setw(3) << std::setprecision(3) -#else - std::cout << std::fixed << std::setprecision(2) << std::setw(6) -#endif + std::cout << std::setw(width) << std::setprecision(precision) << ck_tile::type_convert(mat(i, j)) << " "; } if(cols > limit) @@ -122,8 +121,11 @@ class TestLoadAndConvert : public ::testing::Test bool pass = ck_tile::check_err(h_c, h_a_ref); #if PRINT_MATRICES - print_matrix(h_a, "Matrix A"); - print_matrix(h_c, "Matrix C"); + auto [width, precision] = matrix_type == TestMatrixType::MonotonicSequence + ? std::make_pair(3, 3) + : std::make_pair(2, 6); + print_matrix(h_a, "Matrix A", width, precision); + print_matrix(h_c, "Matrix C", width, precision); #endif EXPECT_TRUE(pass); From 8c0d3fecab5866264098f279324fa232889e865e Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 24 Mar 2026 11:22:25 +0000 Subject: [PATCH 11/23] Move print_matrix to a separate header file under include/ck_tile/host --- .../composablekernel/include/ck_tile/host.hpp | 1 + .../include/ck_tile/host/print_matrix.hpp | 33 +++++++++++++++++++ .../test_load_and_convert_tile.cpp | 30 ----------------- 3 files changed, 34 insertions(+), 30 deletions(-) create mode 100644 projects/composablekernel/include/ck_tile/host/print_matrix.hpp diff --git a/projects/composablekernel/include/ck_tile/host.hpp b/projects/composablekernel/include/ck_tile/host.hpp index 995d8545364f..fbe083335dd8 100644 --- a/projects/composablekernel/include/ck_tile/host.hpp +++ b/projects/composablekernel/include/ck_tile/host.hpp @@ -17,6 +17,7 @@ #include "ck_tile/host/joinable_thread.hpp" #include "ck_tile/host/kernel_launch.hpp" #include "ck_tile/host/permute_pk_int4.hpp" +#include "ck_tile/host/print_matrix.hpp" #include "ck_tile/host/ranges.hpp" #include "ck_tile/host/reference/reference_batched_contraction.hpp" #include "ck_tile/host/reference/reference_batched_dropout.hpp" diff --git a/projects/composablekernel/include/ck_tile/host/print_matrix.hpp b/projects/composablekernel/include/ck_tile/host/print_matrix.hpp new file mode 100644 index 000000000000..8cf5588e619a --- /dev/null +++ b/projects/composablekernel/include/ck_tile/host/print_matrix.hpp @@ -0,0 +1,33 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT +#pragma once + +// Helper to print matrix (for debugging) +template +void print_matrix(const ck_tile::HostTensor& mat, + const std::string& name = "Matrix", + const int width = 3, + const int precision = 3) +{ + const auto lens = mat.get_lengths(); + assert(len(lens) == 2); + const ck_tile::index_t rows = lens[0]; + const ck_tile::index_t cols = lens[1]; + const ck_tile::index_t limit = 32; + + std::cout << name << " (" << rows << "×" << cols << "):\n"; + for(ck_tile::index_t i = 0; i < std::min(rows, ck_tile::index_t(limit)); ++i) + { + for(ck_tile::index_t j = 0; j < std::min(cols, ck_tile::index_t(limit)); ++j) + { + std::cout << std::setw(width) << std::setprecision(precision) + << ck_tile::type_convert(mat(i, j)) << " "; + } + if(cols > limit) + std::cout << "..."; + std::cout << "\n"; + } + if(rows > limit) + std::cout << "...\n"; + std::cout << "\n"; +} diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp index e90de113b241..109daf5a85db 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp @@ -19,36 +19,6 @@ static constexpr auto matrix_type = TestMatrixType::UniformDistribution; #define PRINT_MATRICES 0 -// Helper to print matrix (for debugging) -template -void print_matrix(const ck_tile::HostTensor& mat, - const std::string& name = "Matrix", - const int width = 3, - const int precision = 3) -{ - const auto lens = mat.get_lengths(); - assert(len(lens) == 2); - const ck_tile::index_t rows = lens[0]; - const ck_tile::index_t cols = lens[1]; - const ck_tile::index_t limit = 32; - - std::cout << name << " (" << rows << "×" << cols << "):\n"; - for(ck_tile::index_t i = 0; i < std::min(rows, ck_tile::index_t(limit)); ++i) - { - for(ck_tile::index_t j = 0; j < std::min(cols, ck_tile::index_t(limit)); ++j) - { - std::cout << std::setw(width) << std::setprecision(precision) - << ck_tile::type_convert(mat(i, j)) << " "; - } - if(cols > limit) - std::cout << "..."; - std::cout << "\n"; - } - if(rows > limit) - std::cout << "...\n"; - std::cout << "\n"; -} - template class TestLoadAndConvert : public ::testing::Test { From cd42fd7fac13092f822ef4973485925934e108df Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 24 Mar 2026 11:37:47 +0000 Subject: [PATCH 12/23] Move TestLoadAndConvert to a separate utility file --- .../test_load_and_convert_tile.cpp | 104 +----------------- .../test_load_and_convert_tile_util.hpp | 103 +++++++++++++++++ 2 files changed, 104 insertions(+), 103 deletions(-) create mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp index 109daf5a85db..e636d8f2205e 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp @@ -1,106 +1,4 @@ -// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. -// SPDX-License-Identifier: MIT - -#include - -#include "ck_tile/host.hpp" -#include "ck_tile/ops/common.hpp" -#include "kernel.hpp" - -// Enum struct specifying what kind of test matrix to use -enum struct TestMatrixType -{ - MonotonicSequence = 0, - Identity = 1, - UniformDistribution = 2 -}; - -static constexpr auto matrix_type = TestMatrixType::UniformDistribution; - -#define PRINT_MATRICES 0 - -template -class TestLoadAndConvert : public ::testing::Test -{ - public: - using XDataType = std::tuple_element_t<0, Tuple>; - using YDataType = std::tuple_element_t<1, Tuple>; - using LoadTranspose = std::tuple_element_t<2, Tuple>; - - protected: - void RunTest() - { - constexpr ck_tile::index_t M = 32; - constexpr ck_tile::index_t N = 32; - constexpr ck_tile::index_t K = 32; - - ck_tile::HostTensor h_a({M, K}); - ck_tile::HostTensor h_c({M, K}); - - if constexpr(matrix_type == TestMatrixType::MonotonicSequence) - { - ck_tile::HostTensor h_a_tmp({M, K}); - ck_tile::FillMonotonicSeq{0.0, 0.1}(h_a_tmp); - ck_tile::reference_unary_elementwise( - h_a_tmp, h_a, [](const auto& x) { return x; }); - } - else if constexpr(matrix_type == TestMatrixType::Identity) - { - ck_tile::FillIdentity{M, K}(h_a); - } - else - { - ck_tile::FillUniformDistributionIntegerValue{-5.0, 5.0}(h_a); - } - - ck_tile::DeviceMem d_a(h_a.get_element_space_size_in_bytes()); - ck_tile::DeviceMem d_c(h_c.get_element_space_size_in_bytes()); - - d_a.ToDevice(h_a.data()); - d_c.ToDevice(h_c.data()); - - using BlockWarps = ck_tile::sequence<1, 1, 1>; - using BlockTile = ck_tile::sequence<32, 32, 16>; - using WarpTile = ck_tile::sequence<32, 32, 16>; - using Vector = ck_tile::sequence<1, 8>; - - using Shape = ck_tile::LoadAndConvertShape; - using Problem = ck_tile::LoadAndConvertProblem; - using Kernel = ck_tile::LoadAndConvertKernel; - - constexpr ck_tile::index_t block_size = Kernel::kBlockSize; - const ck_tile::index_t grid_size = ck_tile::integer_divide_ceil(M, Shape::Block_M) * - ck_tile::integer_divide_ceil(N, Shape::Block_N); - - launch_kernel(ck_tile::stream_config{nullptr, true}, - make_kernel(Kernel{}, - dim3(grid_size), - dim3(block_size), - 0, - static_cast(d_a.GetDeviceBuffer()), - static_cast(d_c.GetDeviceBuffer()), - M, - N, - K)); - - ck_tile::hip_check_error(hipDeviceSynchronize()); - d_c.FromDevice(h_c.data()); - ck_tile::HostTensor h_a_ref({M, K}); - ck_tile::reference_unary_elementwise( - h_a, h_a_ref, [](const auto& x) { return x; }); - bool pass = ck_tile::check_err(h_c, h_a_ref); - -#if PRINT_MATRICES - auto [width, precision] = matrix_type == TestMatrixType::MonotonicSequence - ? std::make_pair(3, 3) - : std::make_pair(2, 6); - print_matrix(h_a, "Matrix A", width, precision); - print_matrix(h_c, "Matrix C", width, precision); -#endif - - EXPECT_TRUE(pass); - } -}; +#include "test_load_and_convert_tile_util.hpp" using TestTypes = ::testing::Types, std::tuple, diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp new file mode 100644 index 000000000000..a9f637e97314 --- /dev/null +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp @@ -0,0 +1,103 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT + +#include + +#include "ck_tile/host.hpp" +#include "ck_tile/ops/common.hpp" +#include "kernel.hpp" + +// Enum struct specifying what kind of test matrix to use +enum struct TestMatrixType +{ + MonotonicSequence = 0, + Identity = 1, + UniformDistribution = 2 +}; + +static constexpr auto matrix_type = TestMatrixType::UniformDistribution; + +#define PRINT_MATRICES 0 + +template +class TestLoadAndConvert : public ::testing::Test +{ + public: + using XDataType = std::tuple_element_t<0, Tuple>; + using YDataType = std::tuple_element_t<1, Tuple>; + using LoadTranspose = std::tuple_element_t<2, Tuple>; + + protected: + void RunTest() + { + constexpr ck_tile::index_t M = 32; + constexpr ck_tile::index_t N = 32; + constexpr ck_tile::index_t K = 32; + + ck_tile::HostTensor h_a({M, K}); + ck_tile::HostTensor h_c({M, K}); + + if constexpr(matrix_type == TestMatrixType::MonotonicSequence) + { + ck_tile::HostTensor h_a_tmp({M, K}); + ck_tile::FillMonotonicSeq{0.0, 0.1}(h_a_tmp); + ck_tile::reference_unary_elementwise( + h_a_tmp, h_a, [](const auto& x) { return x; }); + } + else if constexpr(matrix_type == TestMatrixType::Identity) + { + ck_tile::FillIdentity{M, K}(h_a); + } + else + { + ck_tile::FillUniformDistributionIntegerValue{-5.0, 5.0}(h_a); + } + + ck_tile::DeviceMem d_a(h_a.get_element_space_size_in_bytes()); + ck_tile::DeviceMem d_c(h_c.get_element_space_size_in_bytes()); + + d_a.ToDevice(h_a.data()); + d_c.ToDevice(h_c.data()); + + using BlockWarps = ck_tile::sequence<1, 1, 1>; + using BlockTile = ck_tile::sequence<32, 32, 16>; + using WarpTile = ck_tile::sequence<32, 32, 16>; + using Vector = ck_tile::sequence<1, 8>; + + using Shape = ck_tile::LoadAndConvertShape; + using Problem = ck_tile::LoadAndConvertProblem; + using Kernel = ck_tile::LoadAndConvertKernel; + + constexpr ck_tile::index_t block_size = Kernel::kBlockSize; + const ck_tile::index_t grid_size = ck_tile::integer_divide_ceil(M, Shape::Block_M) * + ck_tile::integer_divide_ceil(N, Shape::Block_N); + + launch_kernel(ck_tile::stream_config{nullptr, true}, + make_kernel(Kernel{}, + dim3(grid_size), + dim3(block_size), + 0, + static_cast(d_a.GetDeviceBuffer()), + static_cast(d_c.GetDeviceBuffer()), + M, + N, + K)); + + ck_tile::hip_check_error(hipDeviceSynchronize()); + d_c.FromDevice(h_c.data()); + ck_tile::HostTensor h_a_ref({M, K}); + ck_tile::reference_unary_elementwise( + h_a, h_a_ref, [](const auto& x) { return x; }); + bool pass = ck_tile::check_err(h_c, h_a_ref); + +#if PRINT_MATRICES + auto [width, precision] = matrix_type == TestMatrixType::MonotonicSequence + ? std::make_pair(3, 3) + : std::make_pair(2, 6); + print_matrix(h_a, "Matrix A", width, precision); + print_matrix(h_c, "Matrix C", width, precision); +#endif + + EXPECT_TRUE(pass); + } +}; From 9ff49703e559f00e14b52c8242d330e9bcbebcfd Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 24 Mar 2026 12:59:24 +0000 Subject: [PATCH 13/23] Split the tests into three separate suites, so that only a subset of the tests can be marked as specific to gfx950 --- .../load_and_convert_tile/CMakeLists.txt | 13 +++++++++--- .../test_load_and_convert_tile.cpp | 20 ------------------- ...est_load_and_convert_tile_no_transpose.cpp | 13 ++++++++++++ ..._load_and_convert_tile_transpose_mixed.cpp | 10 ++++++++++ ...oad_and_convert_tile_transpose_uniform.cpp | 9 +++++++++ 5 files changed, 42 insertions(+), 23 deletions(-) delete mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp create mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_no_transpose.cpp create mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp create mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_uniform.cpp diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt index ce6b1061bf00..0d41c7edbfcf 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt @@ -2,8 +2,15 @@ # SPDX-License-Identifier: MIT set(LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) +list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8 -fverbose-asm --save-temps -Wno-gnu-line-marker) + +add_gtest_executable(test_load_and_convert_tile_no_transpose test_load_and_convert_tile_no_transpose.cpp) +target_compile_options(test_load_and_convert_tile_no_transpose PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) + +add_gtest_executable(test_load_and_convert_tile_transpose_uniform test_load_and_convert_tile_transpose_uniform.cpp) +target_compile_options(test_load_and_convert_tile_transpose_uniform PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) + if(GPU_TARGETS MATCHES "gfx950") - add_gtest_executable(test_load_and_convert_tile test_load_and_convert_tile.cpp) - list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8 -fverbose-asm --save-temps -Wno-gnu-line-marker) - target_compile_options(test_load_and_convert_tile PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) + add_gtest_executable(test_load_and_convert_tile_transpose_mixed test_load_and_convert_tile_transpose_mixed.cpp) + target_compile_options(test_load_and_convert_tile_transpose_mixed PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) endif() diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp deleted file mode 100644 index e636d8f2205e..000000000000 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile.cpp +++ /dev/null @@ -1,20 +0,0 @@ -#include "test_load_and_convert_tile_util.hpp" - -using TestTypes = ::testing::Types, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple, - std::tuple>; - -TYPED_TEST_SUITE(TestLoadAndConvert, TestTypes); - -TYPED_TEST(TestLoadAndConvert, Test) { this->RunTest(); } diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_no_transpose.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_no_transpose.cpp new file mode 100644 index 000000000000..e5eb8e6cc7f3 --- /dev/null +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_no_transpose.cpp @@ -0,0 +1,13 @@ +#include "test_load_and_convert_tile_util.hpp" + +using TestTypes = ::testing::Types, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple, + std::tuple>; + +TYPED_TEST_SUITE(TestLoadAndConvert, TestTypes); + +TYPED_TEST(TestLoadAndConvert, TestNoTranspose) { this->RunTest(); } diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp new file mode 100644 index 000000000000..52057a923012 --- /dev/null +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp @@ -0,0 +1,10 @@ +#include "test_load_and_convert_tile_util.hpp" + +using TestTypes = ::testing::Types, + std::tuple, + std::tuple, + std::tuple>; + +TYPED_TEST_SUITE(TestLoadAndConvert, TestTypes); + +TYPED_TEST(TestLoadAndConvert, Test) { this->RunTest(); } diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_uniform.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_uniform.cpp new file mode 100644 index 000000000000..a7276709b05b --- /dev/null +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_uniform.cpp @@ -0,0 +1,9 @@ +#include "test_load_and_convert_tile_util.hpp" + +using TestTypes = ::testing::Types, + std::tuple, + std::tuple>; + +TYPED_TEST_SUITE(TestLoadAndConvert, TestTypes); + +TYPED_TEST(TestLoadAndConvert, TestTransposeUniform) { this->RunTest(); } From f73169f50a79ac9562aab09f28dec6771b4d019b Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 24 Mar 2026 14:24:24 +0000 Subject: [PATCH 14/23] fixup! Split the tests into three separate suites, so that only a subset of the tests can be marked as specific to gfx950 --- .../test_load_and_convert_tile_transpose_mixed.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp index 52057a923012..58e228f7d84b 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp @@ -7,4 +7,4 @@ using TestTypes = ::testing::TypesRunTest(); } +TYPED_TEST(TestLoadAndConvert, TestTransposeMixed) { this->RunTest(); } From 97a4babf5d8367378001ab2cd47be1e9c17bc93b Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 24 Mar 2026 15:54:52 +0000 Subject: [PATCH 15/23] All transposed loads are a gfx950 feature, so they should all be run only on that architecture --- .../test/ck_tile/load_and_convert_tile/CMakeLists.txt | 7 ++----- .../test_load_and_convert_tile_transpose_uniform.cpp | 9 --------- ...xed.cpp => test_load_and_convert_tile_transposed.cpp} | 9 ++++++--- 3 files changed, 8 insertions(+), 17 deletions(-) delete mode 100644 projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_uniform.cpp rename projects/composablekernel/test/ck_tile/load_and_convert_tile/{test_load_and_convert_tile_transpose_mixed.cpp => test_load_and_convert_tile_transposed.cpp} (53%) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt index 0d41c7edbfcf..694da583b937 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt @@ -7,10 +7,7 @@ list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8 -fverbose- add_gtest_executable(test_load_and_convert_tile_no_transpose test_load_and_convert_tile_no_transpose.cpp) target_compile_options(test_load_and_convert_tile_no_transpose PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) -add_gtest_executable(test_load_and_convert_tile_transpose_uniform test_load_and_convert_tile_transpose_uniform.cpp) -target_compile_options(test_load_and_convert_tile_transpose_uniform PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) - if(GPU_TARGETS MATCHES "gfx950") - add_gtest_executable(test_load_and_convert_tile_transpose_mixed test_load_and_convert_tile_transpose_mixed.cpp) - target_compile_options(test_load_and_convert_tile_transpose_mixed PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) + add_gtest_executable(test_load_and_convert_tile_transposed test_load_and_convert_tile_transposed.cpp) + target_compile_options(test_load_and_convert_tile_transposed PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) endif() diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_uniform.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_uniform.cpp deleted file mode 100644 index a7276709b05b..000000000000 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_uniform.cpp +++ /dev/null @@ -1,9 +0,0 @@ -#include "test_load_and_convert_tile_util.hpp" - -using TestTypes = ::testing::Types, - std::tuple, - std::tuple>; - -TYPED_TEST_SUITE(TestLoadAndConvert, TestTypes); - -TYPED_TEST(TestLoadAndConvert, TestTransposeUniform) { this->RunTest(); } diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transposed.cpp similarity index 53% rename from projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp rename to projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transposed.cpp index 58e228f7d84b..c09299812bfc 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transpose_mixed.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transposed.cpp @@ -1,10 +1,13 @@ #include "test_load_and_convert_tile_util.hpp" -using TestTypes = ::testing::Types, +using TestTypes = ::testing::Types, + std::tuple, std::tuple, + std::tuple, std::tuple, - std::tuple>; + std::tuple, + std::tuple>; TYPED_TEST_SUITE(TestLoadAndConvert, TestTypes); -TYPED_TEST(TestLoadAndConvert, TestTransposeMixed) { this->RunTest(); } +TYPED_TEST(TestLoadAndConvert, TestTransposed) { this->RunTest(); } From 587955e2a4e0dcf6dfce202959e02eef79d9ef70 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Thu, 26 Mar 2026 13:45:12 +0000 Subject: [PATCH 16/23] Adjust compile options --- .../test/ck_tile/load_and_convert_tile/CMakeLists.txt | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt index 694da583b937..a74e369ac671 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt @@ -2,7 +2,10 @@ # SPDX-License-Identifier: MIT set(LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) -list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8 -fverbose-asm --save-temps -Wno-gnu-line-marker) +if(CK_USE_OCP_FP8) + list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8) +endif() +list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) add_gtest_executable(test_load_and_convert_tile_no_transpose test_load_and_convert_tile_no_transpose.cpp) target_compile_options(test_load_and_convert_tile_no_transpose PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) From a807103c0716c6b62d4b25af18c11887cb2b8f2a Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Thu, 26 Mar 2026 14:46:02 +0000 Subject: [PATCH 17/23] Limit tests to gfx9 because of linker problems detected on gfx1201 and gfx1030 --- .../load_and_convert_tile/CMakeLists.txt | 22 ++++++++++--------- 1 file changed, 12 insertions(+), 10 deletions(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt index a74e369ac671..83edc3248ea0 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/CMakeLists.txt @@ -1,16 +1,18 @@ # Copyright (c) Advanced Micro Devices, Inc., or its affiliates. # SPDX-License-Identifier: MIT -set(LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) -if(CK_USE_OCP_FP8) - list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8) -endif() -list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) +if(GPU_TARGETS MATCHES "gfx9") + set(LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) + if(CK_USE_OCP_FP8) + list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS -DCK_TILE_USE_OCP_FP8) + endif() + list(APPEND LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS) -add_gtest_executable(test_load_and_convert_tile_no_transpose test_load_and_convert_tile_no_transpose.cpp) -target_compile_options(test_load_and_convert_tile_no_transpose PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) + add_gtest_executable(test_load_and_convert_tile_no_transpose test_load_and_convert_tile_no_transpose.cpp) + target_compile_options(test_load_and_convert_tile_no_transpose PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) -if(GPU_TARGETS MATCHES "gfx950") - add_gtest_executable(test_load_and_convert_tile_transposed test_load_and_convert_tile_transposed.cpp) - target_compile_options(test_load_and_convert_tile_transposed PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) + if(GPU_TARGETS MATCHES "gfx950") + add_gtest_executable(test_load_and_convert_tile_transposed test_load_and_convert_tile_transposed.cpp) + target_compile_options(test_load_and_convert_tile_transposed PRIVATE ${LOAD_TILE_TRANSPOSE_COMPILE_OPTIONS}) + endif() endif() From 4cc024f81d5f8df50d36381922c178557fd954d7 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Fri, 27 Mar 2026 09:21:47 +0000 Subject: [PATCH 18/23] Modify the M, N, and K parameters to span multiple tiles --- .../test_load_and_convert_tile_util.hpp | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp index a9f637e97314..c6c2d2fbc038 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp @@ -30,9 +30,9 @@ class TestLoadAndConvert : public ::testing::Test protected: void RunTest() { - constexpr ck_tile::index_t M = 32; - constexpr ck_tile::index_t N = 32; - constexpr ck_tile::index_t K = 32; + constexpr ck_tile::index_t M = 256; + constexpr ck_tile::index_t N = 256; + constexpr ck_tile::index_t K = 64; ck_tile::HostTensor h_a({M, K}); ck_tile::HostTensor h_c({M, K}); From 4ced776adfd142f7ed6a67147fa5b77d2420568f Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Fri, 27 Mar 2026 13:59:16 +0000 Subject: [PATCH 19/23] No need to send C tensor to device --- .../load_and_convert_tile/test_load_and_convert_tile_util.hpp | 1 - 1 file changed, 1 deletion(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp index c6c2d2fbc038..b7f7e2363caa 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp @@ -57,7 +57,6 @@ class TestLoadAndConvert : public ::testing::Test ck_tile::DeviceMem d_c(h_c.get_element_space_size_in_bytes()); d_a.ToDevice(h_a.data()); - d_c.ToDevice(h_c.data()); using BlockWarps = ck_tile::sequence<1, 1, 1>; using BlockTile = ck_tile::sequence<32, 32, 16>; From fbae34984a4d70f51ca3523fced576c3dd1c8652 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Tue, 7 Apr 2026 17:38:03 +0000 Subject: [PATCH 20/23] Fix test to work with multiple warps per block, and fix tile distribution encodings so that row permutations are not needed --- .../ck_tile/load_and_convert_tile/kernel.hpp | 206 ++++++++---------- ...est_load_and_convert_tile_no_transpose.cpp | 3 + .../test_load_and_convert_tile_transposed.cpp | 3 + .../test_load_and_convert_tile_util.hpp | 41 ++-- 4 files changed, 119 insertions(+), 134 deletions(-) diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp index f283d44e8ff9..9689038a4d84 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/kernel.hpp @@ -15,21 +15,17 @@ struct LoadAndConvertShape { static constexpr index_t Block_M = BlockTile::at(number<0>{}); static constexpr index_t Block_N = BlockTile::at(number<1>{}); - static constexpr index_t Block_K = BlockTile::at(number<2>{}); static constexpr index_t Warp_M = WarpTile::at(number<0>{}); static constexpr index_t Warp_N = WarpTile::at(number<1>{}); - static constexpr index_t Warp_K = WarpTile::at(number<2>{}); static constexpr index_t Vector_N = Vector::at(number<1>{}); static constexpr index_t WarpPerBlock_M = BlockWarps::at(number<0>{}); static constexpr index_t WarpPerBlock_N = BlockWarps::at(number<1>{}); - static constexpr index_t WarpPerBlock_K = BlockWarps::at(number<2>{}); static constexpr index_t Repeat_M = Block_M / (WarpPerBlock_M * Warp_M); static constexpr index_t Repeat_N = Block_N / (WarpPerBlock_N * Warp_N); - static constexpr index_t Repeat_K = Block_K / (WarpPerBlock_K * Warp_K); static constexpr index_t BlockSize = ck_tile::get_warp_size() * reduce_on_sequence(BlockWarps{}, multiplies<>{}, number<1>{}); @@ -45,35 +41,36 @@ struct LoadAndConvertProblem }; template -struct LoadAndConvertKernel +struct LoadAndConvertPolicy { using Problem = ck_tile::remove_cvref_t; using XDataType = ck_tile::remove_cvref_t; using YDataType = ck_tile::remove_cvref_t; using LoadTranspose = ck_tile::remove_cvref_t; - static constexpr index_t kBlockSize = Problem::BlockShape::BlockSize; - - template - static constexpr auto get_warp_dstr_encoding() + template + CK_TILE_DEVICE static constexpr auto get_warp_dstr_encoding() { using S = typename Problem::BlockShape; if constexpr(NumAccess == 1) - return tile_distribution_encoding, - tuple, sequence<2, S::Vector_N>>, - tuple>, - tuple>, - sequence<2>, - sequence<1>>{}; + return tile_distribution_encoding< + sequence<1>, + tuple, + sequence>, + tuple>, + tuple>, + sequence<2>, + sequence<1>>{}; else return tile_distribution_encoding< - sequence<>, - tuple, sequence>, + sequence<1>, + tuple, + sequence>, tuple>, - tuple>, + tuple>, sequence<2, 2>, - sequence<0, 2>>{}; + sequence<1, 2>>{}; } template @@ -82,150 +79,133 @@ struct LoadAndConvertKernel return DS_READ_TR_SIZE() / sizeof(DataType); } - template + template CK_TILE_DEVICE static constexpr auto MakeDRAMDistribution() { - using S = typename Problem::BlockShape; - constexpr index_t thread_elements = S::Warp_N * S::Warp_K / get_warp_size(); + using S = typename Problem::BlockShape; + + constexpr index_t thread_elements = S::Warp_M * S::Warp_N / get_warp_size(); constexpr index_t NumAccess = LoadTranspose::value ? thread_elements / GetVectorSize() : 1; constexpr auto a_block_outer_dstr_encode = tile_distribution_encoding< - sequence, - tuple, sequence>, - tuple>, - tuple>, - sequence<1, 2>, - sequence<0, 0>>{}; + sequence<>, + tuple, + sequence<>>, + tuple>, + tuple>, + sequence<1>, + sequence<0>>{}; constexpr auto a_block_dstr_encode = detail::make_embed_tile_distribution_encoding( - a_block_outer_dstr_encode, get_warp_dstr_encoding()); + a_block_outer_dstr_encode, get_warp_dstr_encoding()); return make_static_tile_distribution(a_block_dstr_encode); } - template + template CK_TILE_DEVICE static constexpr auto MakeDRAMTransposedDistribution() { return make_static_tile_distribution( typename InputTileDistributionTraits< - typename decltype(MakeDRAMDistribution())::DstrEncode, + typename decltype(MakeDRAMDistribution())::DstrEncode, DataType>::TransposedDstrEncode{}); } +}; + +template +struct LoadAndConvertKernel +{ + using Problem = ck_tile::remove_cvref_t; + using XDataType = ck_tile::remove_cvref_t; + using YDataType = ck_tile::remove_cvref_t; + using Policy = ck_tile::remove_cvref_t; + using LoadTranspose = ck_tile::remove_cvref_t; + + static constexpr index_t kBlockSize = Problem::BlockShape::BlockSize; - CK_TILE_DEVICE void - operator()(const XDataType* a, YDataType* c, index_t M, index_t N, index_t K) const + CK_TILE_HOST static auto BlockSize() { - using S = typename Problem::BlockShape; + if(ck_tile::is_wave32()) + { + return kBlockSize / 2; + } + else + { + return kBlockSize; + } + } - const index_t kMPerBlock = S::WarpPerBlock_M * S::Repeat_M * S::Block_M; - const index_t kNPerBlock = S::WarpPerBlock_N * S::Repeat_N * S::Block_N; + CK_TILE_DEVICE void operator()(const XDataType* a, YDataType* c, index_t M, index_t N) const + { + using S = typename Problem::BlockShape; - constexpr auto block_dims = make_tuple(number{}, number{}); - constexpr auto block_strides = make_tuple(number<1>{}, number{}); - const index_t num_blocks_n = N / kNPerBlock; - const index_t block_m = get_block_id() / num_blocks_n; + constexpr auto block_dims = make_tuple(S::Block_M, S::Block_N); + constexpr auto block_strides = make_tuple(1, S::Block_M); - const index_t m_block_base = block_m * kMPerBlock; + const index_t m_block_base = get_block_id() * S::Block_M; // LDS buffer - __shared__ XDataType a_lds[kMPerBlock * S::Block_K]; + __shared__ XDataType a_lds[S::Block_M * S::Block_N]; - auto a_lds_write_view = make_naive_tensor_view( + auto a_lds_view = make_naive_tensor_view( a_lds, block_dims, block_strides, number<1>{}, number<1>{}); - auto a_block_lds_write_window = make_tile_window(a_lds_write_view, block_dims, {0, 0}); + auto a_block_lds_write_window = make_tile_window(a_lds_view, block_dims, {0, 0}); auto a_block_lds_read_window = [&] { if constexpr(LoadTranspose::value) { - constexpr auto block_dims_t = - make_tuple(number{}, number{}); - constexpr auto block_strides_t = make_tuple(number{}, number<1>{}); + constexpr auto block_dims_t = make_tuple(S::Block_N, S::Block_M); + constexpr auto block_strides_t = make_tuple(S::Block_M, 1); auto view = make_naive_tensor_view( a_lds, block_dims_t, block_strides_t, - number()>{}, + number()>{}, number<1>{}); return make_tile_window( - view, block_dims_t, {0, 0}, MakeDRAMTransposedDistribution()); + view, + block_dims_t, + {0, 0}, + Policy::template MakeDRAMTransposedDistribution()); } else { - auto view = make_naive_tensor_view( - a_lds, block_dims, block_strides, number<1>{}, number<1>{}); - return make_tile_window( - view, block_dims, {0, 0}, MakeDRAMDistribution()); + a_lds_view, + block_dims, + {0, 0}, + Policy::template MakeDRAMDistribution()); } }(); // Input tensor - const auto a_tensor = [&]() { - if constexpr(LoadTranspose::value && !std::is_same_v) - { - // Similar to QuantGemmKernel with PermuteB: reinterpret the output logical layout - // via descriptor transform so YDataType distribution can be used without row-group - // permutation artifacts in mixed-precision transpose loads. - using TransposeGroupType = std:: - conditional_t<(sizeof(XDataType) >= sizeof(YDataType)), XDataType, YDataType>; - constexpr index_t thread_elements = S::Warp_N * S::Warp_K / get_warp_size(); - constexpr index_t n_group_0 = thread_elements / GetVectorSize(); - constexpr index_t n_group_1 = 2; - constexpr index_t n_group_2 = S::Vector_N / n_group_0; - constexpr index_t n_perm_group = n_group_0 * n_group_1 * n_group_2; - - static_assert(n_group_0 > 0, "Invalid derived transpose grouping factor"); - static_assert(S::Vector_N % n_group_0 == 0, - "Vector_N must be divisible by derived grouping factor"); - - const auto a_m_n_desc = make_naive_tensor_descriptor( - make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); - - const auto a_m_n0_b0_b1_n4_desc = transform_tensor_descriptor( - a_m_n_desc, - make_tuple(make_pass_through_transform(M), - make_unmerge_transform(make_tuple(N / n_perm_group, - number{}, - number{}, - number{}))), - make_tuple(sequence<0>{}, sequence<1>{}), - make_tuple(sequence<0>{}, sequence<1, 2, 3, 4>{})); - - const auto a_perm_m_n_desc = transform_tensor_descriptor( - a_m_n0_b0_b1_n4_desc, - make_tuple(make_pass_through_transform(M), - make_merge_transform(make_tuple(N / n_perm_group, - number{}, - number{}, - number{}))), - make_tuple(sequence<0>{}, sequence<1, 3, 2, 4>{}), - make_tuple(sequence<0>{}, sequence<1>{})); - - return make_tensor_view(a, a_perm_m_n_desc); - } - else - { - return make_naive_tensor_view( - a, make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); - } - }(); + const auto a_tensor = make_naive_tensor_view( + a, make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); - auto a_block_window = make_tile_window( - a_tensor, block_dims, {m_block_base, 0}, MakeDRAMDistribution()); + auto a_block_window = + make_tile_window(a_tensor, + block_dims, + {m_block_base, 0}, + Policy::template MakeDRAMDistribution()); // Output tensor - auto c_tensor = make_naive_tensor_view( - c, make_tuple(M, K), make_tuple(1, M), number<1>{}, number<1>{}); + const auto c_tensor = make_naive_tensor_view( + c, make_tuple(M, N), make_tuple(1, M), number<1>{}, number<1>{}); - auto c_block_window = make_tile_window( - c_tensor, block_dims, {m_block_base, 0}, MakeDRAMDistribution()); + auto c_block_window = + make_tile_window(c_tensor, + block_dims, + {m_block_base, 0}, + Policy::template MakeDRAMDistribution()); - const index_t num_k_loops = K / S::Block_K; - for(index_t k_iter = 0; k_iter < num_k_loops; ++k_iter) + const index_t num_n_loops = N / S::Block_N; + for(index_t n_iter = 0; n_iter < num_n_loops; ++n_iter) { auto dram_tile = load_tile(a_block_window); store_tile(a_block_lds_write_window, dram_tile); @@ -235,10 +215,10 @@ struct LoadAndConvertKernel load_and_convert_tile<8, LoadTranspose::value>(c_tile, a_block_lds_read_window); store_tile(c_block_window, c_tile); - if(k_iter < num_k_loops - 1) + if(n_iter < num_n_loops - 1) { - move_tile_window(a_block_window, {0, S::Block_K}); - move_tile_window(c_block_window, {0, S::Block_K}); + move_tile_window(a_block_window, {0, S::Block_N}); + move_tile_window(c_block_window, {0, S::Block_N}); } } } diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_no_transpose.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_no_transpose.cpp index e5eb8e6cc7f3..bffb830a7267 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_no_transpose.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_no_transpose.cpp @@ -1,3 +1,6 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT + #include "test_load_and_convert_tile_util.hpp" using TestTypes = ::testing::Types, diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transposed.cpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transposed.cpp index c09299812bfc..07717902b1e3 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transposed.cpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_transposed.cpp @@ -1,3 +1,6 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT + #include "test_load_and_convert_tile_util.hpp" using TestTypes = ::testing::Types, diff --git a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp index b7f7e2363caa..b2984ece2b2f 100644 --- a/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp +++ b/projects/composablekernel/test/ck_tile/load_and_convert_tile/test_load_and_convert_tile_util.hpp @@ -32,21 +32,20 @@ class TestLoadAndConvert : public ::testing::Test { constexpr ck_tile::index_t M = 256; constexpr ck_tile::index_t N = 256; - constexpr ck_tile::index_t K = 64; - ck_tile::HostTensor h_a({M, K}); - ck_tile::HostTensor h_c({M, K}); + ck_tile::HostTensor h_a({M, N}); + ck_tile::HostTensor h_c({M, N}); if constexpr(matrix_type == TestMatrixType::MonotonicSequence) { - ck_tile::HostTensor h_a_tmp({M, K}); + ck_tile::HostTensor h_a_tmp({M, N}); ck_tile::FillMonotonicSeq{0.0, 0.1}(h_a_tmp); ck_tile::reference_unary_elementwise( h_a_tmp, h_a, [](const auto& x) { return x; }); } else if constexpr(matrix_type == TestMatrixType::Identity) { - ck_tile::FillIdentity{M, K}(h_a); + ck_tile::FillIdentity{M, N}(h_a); } else { @@ -58,33 +57,33 @@ class TestLoadAndConvert : public ::testing::Test d_a.ToDevice(h_a.data()); - using BlockWarps = ck_tile::sequence<1, 1, 1>; - using BlockTile = ck_tile::sequence<32, 32, 16>; - using WarpTile = ck_tile::sequence<32, 32, 16>; + using BlockWarps = ck_tile::sequence<4, 4>; + using BlockTile = ck_tile::sequence<512, 32>; + using WarpTile = ck_tile::sequence<64, 8>; using Vector = ck_tile::sequence<1, 8>; using Shape = ck_tile::LoadAndConvertShape; using Problem = ck_tile::LoadAndConvertProblem; - using Kernel = ck_tile::LoadAndConvertKernel; + using Policy = ck_tile::LoadAndConvertPolicy; + using Kernel = ck_tile::LoadAndConvertKernel; - constexpr ck_tile::index_t block_size = Kernel::kBlockSize; - const ck_tile::index_t grid_size = ck_tile::integer_divide_ceil(M, Shape::Block_M) * + const ck_tile::index_t block_size = Kernel::BlockSize(); + const ck_tile::index_t grid_size = ck_tile::integer_divide_ceil(M, Shape::Block_M) * ck_tile::integer_divide_ceil(N, Shape::Block_N); launch_kernel(ck_tile::stream_config{nullptr, true}, - make_kernel(Kernel{}, - dim3(grid_size), - dim3(block_size), - 0, - static_cast(d_a.GetDeviceBuffer()), - static_cast(d_c.GetDeviceBuffer()), - M, - N, - K)); + make_kernel<1>(Kernel{}, + dim3(grid_size), + dim3(block_size), + 0, + static_cast(d_a.GetDeviceBuffer()), + static_cast(d_c.GetDeviceBuffer()), + M, + N)); ck_tile::hip_check_error(hipDeviceSynchronize()); d_c.FromDevice(h_c.data()); - ck_tile::HostTensor h_a_ref({M, K}); + ck_tile::HostTensor h_a_ref({M, N}); ck_tile::reference_unary_elementwise( h_a, h_a_ref, [](const auto& x) { return x; }); bool pass = ck_tile::check_err(h_c, h_a_ref); From fa2c74d504b6792c67d5d32ea0157573f2884f77 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Thu, 9 Oct 2025 08:07:04 +0000 Subject: [PATCH 21/23] Introduce DetermineWarpPrecType for determining warp GEMM precision types --- .../ck_tile/ops/add_rmsnorm2d_rdquant.hpp | 1 + .../ck_tile/ops/batched_contraction.hpp | 1 + .../include/ck_tile/ops/batched_transpose.hpp | 1 + .../include/ck_tile/ops/common.hpp | 1 + .../ops/common/determine_warp_prec_type.hpp | 134 ++++++++++++++++++ .../include/ck_tile/ops/elementwise.hpp | 1 + .../include/ck_tile/ops/epilogue.hpp | 1 + .../ops/epilogue/cshuffle_epilogue.hpp | 16 +-- .../include/ck_tile/ops/flatmm.hpp | 1 + .../include/ck_tile/ops/fmha.hpp | 1 + .../include/ck_tile/ops/fused_moe.hpp | 1 + .../include/ck_tile/ops/gemm.hpp | 1 + .../block/block_universal_gemm_as_bs_cr.hpp | 9 +- ...emm_universal_pipeline_ag_bg_cr_policy.hpp | 10 +- .../include/ck_tile/ops/gemm_mx.hpp | 1 + .../include/ck_tile/ops/gemm_quant.hpp | 1 + .../ck_tile/ops/grouped_convolution.hpp | 1 + .../include/ck_tile/ops/image_to_column.hpp | 1 + .../include/ck_tile/ops/layernorm2d.hpp | 1 + .../include/ck_tile/ops/norm_reduce.hpp | 1 + .../include/ck_tile/ops/permute.hpp | 1 + .../include/ck_tile/ops/pooling.hpp | 1 + .../include/ck_tile/ops/reduce.hpp | 1 + .../include/ck_tile/ops/rmsnorm2d.hpp | 1 + .../include/ck_tile/ops/smoothquant.hpp | 1 + .../include/ck_tile/ops/softmax.hpp | 1 + .../include/ck_tile/ops/sparse_attn.hpp | 1 + .../include/ck_tile/ops/topk.hpp | 1 + .../include/ck_tile/ops/topk_softmax.hpp | 1 + 29 files changed, 169 insertions(+), 25 deletions(-) create mode 100644 projects/composablekernel/include/ck_tile/ops/common/determine_warp_prec_type.hpp diff --git a/projects/composablekernel/include/ck_tile/ops/add_rmsnorm2d_rdquant.hpp b/projects/composablekernel/include/ck_tile/ops/add_rmsnorm2d_rdquant.hpp index aa0f632c2169..a62bbe981cca 100644 --- a/projects/composablekernel/include/ck_tile/ops/add_rmsnorm2d_rdquant.hpp +++ b/projects/composablekernel/include/ck_tile/ops/add_rmsnorm2d_rdquant.hpp @@ -7,6 +7,7 @@ #include "ck_tile/ops/add_rmsnorm2d_rdquant/pipeline/add_rmsnorm2d_rdquant_fwd_pipeline_one_pass.hpp" #include "ck_tile/ops/add_rmsnorm2d_rdquant/pipeline/add_rmsnorm2d_rdquant_fwd_pipeline_problem.hpp" #include "ck_tile/ops/add_rmsnorm2d_rdquant/pipeline/add_rmsnorm2d_rdquant_fwd_pipeline_three_pass.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/batched_contraction.hpp b/projects/composablekernel/include/ck_tile/ops/batched_contraction.hpp index 9c90db67eddc..71919b61873d 100644 --- a/projects/composablekernel/include/ck_tile/ops/batched_contraction.hpp +++ b/projects/composablekernel/include/ck_tile/ops/batched_contraction.hpp @@ -5,6 +5,7 @@ #include "ck_tile/ops/batched_contraction/kernel/batched_contraction_kernel.hpp" #include "ck_tile/ops/batched_contraction/pipeline/batched_contraction_problem.hpp" #include "ck_tile/ops/batched_contraction/utils/tensor_descriptor_utils.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/batched_transpose.hpp b/projects/composablekernel/include/ck_tile/ops/batched_transpose.hpp index 9cac035c4457..924db5fb60e4 100644 --- a/projects/composablekernel/include/ck_tile/ops/batched_transpose.hpp +++ b/projects/composablekernel/include/ck_tile/ops/batched_transpose.hpp @@ -10,6 +10,7 @@ #include "ck_tile/ops/batched_transpose/pipeline/batched_transpose_pipeline.hpp" #include "ck_tile/ops/batched_transpose/pipeline/batched_transpose_policy.hpp" #include "ck_tile/ops/batched_transpose/pipeline/batched_transpose_problem.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/common.hpp b/projects/composablekernel/include/ck_tile/ops/common.hpp index ad7da5c18339..0113d8c9a280 100644 --- a/projects/composablekernel/include/ck_tile/ops/common.hpp +++ b/projects/composablekernel/include/ck_tile/ops/common.hpp @@ -2,6 +2,7 @@ // SPDX-License-Identifier: MIT #pragma once +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/common/determine_warp_prec_type.hpp b/projects/composablekernel/include/ck_tile/ops/common/determine_warp_prec_type.hpp new file mode 100644 index 000000000000..1fa105937f9d --- /dev/null +++ b/projects/composablekernel/include/ck_tile/ops/common/determine_warp_prec_type.hpp @@ -0,0 +1,134 @@ +// Copyright (c) Advanced Micro Devices, Inc., or its affiliates. +// SPDX-License-Identifier: MIT + +#pragma once + +#include "ck_tile/core.hpp" + +// DetermineWarpPrecType is a set of pattern-matching rules to determine the right precision types +// to use for the warp GEMM, given the precision types defined in the problem, and the compute data +// type. This gives rise to a type conversion: type conversions are sometimes needed to obtain +// a pair of types that are compatible with the hardware matrix operations available. A typical +// use case is mixed precision GEMMs. + +namespace ck_tile { +// For the most general case, default to no conversion. +template +struct DetermineWarpPrecType +{ + using a_prec_type = APrecType; + using b_prec_type = BPrecType; +}; + +// Use tf32_t if compute type if tf32_t +template +struct DetermineWarpPrecType +{ + using a_prec_type = ck_tile::tf32_t; + using b_prec_type = ck_tile::tf32_t; +}; + +// For pk_fp4_t x pk_fp4_t, keep pk_fp4_t +template +struct DetermineWarpPrecType +{ + using a_prec_type = ck_tile::pk_fp4_t; + using b_prec_type = ck_tile::pk_fp4_t; +}; + +// For pk_int4_t x B, use the B type. +template +struct DetermineWarpPrecType +{ + using a_prec_type = BPrecType; + using b_prec_type = BPrecType; +}; + +// For A x pk_int4_t, use the A type. +template +struct DetermineWarpPrecType +{ + using a_prec_type = APrecType; + using b_prec_type = APrecType; +}; + +// For pk_fp4_t x B, use the B type. +template +struct DetermineWarpPrecType +{ + using a_prec_type = BPrecType; + using b_prec_type = BPrecType; +}; + +// For A x pk_fp4_t, use the A type. +template +struct DetermineWarpPrecType +{ + using a_prec_type = APrecType; + using b_prec_type = APrecType; +}; + +// For pk_fp4_raw_t x B, use the B type. +template +struct DetermineWarpPrecType +{ + using a_prec_type = BPrecType; + using b_prec_type = BPrecType; +}; + +// For A x pk_fp4_raw_t, use the A type. +template +struct DetermineWarpPrecType +{ + using a_prec_type = APrecType; + using b_prec_type = APrecType; +}; + +// For fp8 x bf16, use fp8 +template +struct DetermineWarpPrecType +{ + using a_prec_type = ck_tile::fp8_t; + using b_prec_type = ck_tile::fp8_t; +}; + +// For bf16 x fp8, use bf16 +template +struct DetermineWarpPrecType +{ + using a_prec_type = ck_tile::bf16_t; + using b_prec_type = ck_tile::bf16_t; +}; + +// For bf8 x bf16, use bf8 +template +struct DetermineWarpPrecType +{ + using a_prec_type = ck_tile::bf8_t; + using b_prec_type = ck_tile::bf8_t; +}; + +// For bf16 x bf8, use bf16 +template +struct DetermineWarpPrecType +{ + using a_prec_type = ck_tile::bf16_t; + using b_prec_type = ck_tile::bf16_t; +}; + +// For fp8 x fp16, use fp8 +template +struct DetermineWarpPrecType +{ + using a_prec_type = ck_tile::fp8_t; + using b_prec_type = ck_tile::fp8_t; +}; + +// For fp16 x fp8, use fp16 +template +struct DetermineWarpPrecType +{ + using a_prec_type = ck_tile::half_t; + using b_prec_type = ck_tile::half_t; +}; +}; // namespace ck_tile diff --git a/projects/composablekernel/include/ck_tile/ops/elementwise.hpp b/projects/composablekernel/include/ck_tile/ops/elementwise.hpp index bc72f3b0ba1f..2c0ae4ad093f 100644 --- a/projects/composablekernel/include/ck_tile/ops/elementwise.hpp +++ b/projects/composablekernel/include/ck_tile/ops/elementwise.hpp @@ -8,6 +8,7 @@ #include "ck_tile/ops/elementwise/pipeline/elementwise_pipeline_problem.hpp" #include "ck_tile/ops/elementwise/pipeline/elementwise_shape.hpp" #include "ck_tile/ops/elementwise/unary_element_wise_operation.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/epilogue.hpp b/projects/composablekernel/include/ck_tile/ops/epilogue.hpp index d1b38a8bca6f..0eb9e59e723c 100644 --- a/projects/composablekernel/include/ck_tile/ops/epilogue.hpp +++ b/projects/composablekernel/include/ck_tile/ops/epilogue.hpp @@ -10,6 +10,7 @@ #include "ck_tile/ops/epilogue/default_2d_and_dynamic_quant_epilogue.hpp" #include "ck_tile/ops/epilogue/default_2d_epilogue.hpp" #include "ck_tile/ops/epilogue/dynamic_quant_epilogue.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/epilogue/cshuffle_epilogue.hpp b/projects/composablekernel/include/ck_tile/ops/epilogue/cshuffle_epilogue.hpp index fba831e20534..d4fc127a3b76 100644 --- a/projects/composablekernel/include/ck_tile/ops/epilogue/cshuffle_epilogue.hpp +++ b/projects/composablekernel/include/ck_tile/ops/epilogue/cshuffle_epilogue.hpp @@ -7,6 +7,7 @@ #include "ck_tile/core.hpp" #include "ck_tile/ops/common/utils.hpp" #include "ck_tile/ops/gemm/warp/warp_gemm_dispatcher.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/tensor_layout.hpp" #include "ck_tile/ops/elementwise/unary_element_wise_operation.hpp" @@ -100,21 +101,10 @@ struct CShuffleEpilogue // For warp gemm selection: use tf32_t if compute type was tf32_t // For pk_int4/pk_fp4: use the other data type using ATypeToUse = - std::conditional_t, - tf32_t, - std::conditional_t || - std::is_same_v, - BDataTypeBuf, - ADataTypeBuf>>; + typename DetermineWarpPrecType::a_prec_type; // Used for weight-only quantization kernel, B would be dequantized to the same data type as A using BTypeToUse = - std::conditional_t, - tf32_t, - std::conditional_t || - std::is_same_v || - sizeof(BDataTypeBuf) < sizeof(ADataTypeBuf), - ADataTypeBuf, - BDataTypeBuf>>; + typename DetermineWarpPrecType::b_prec_type; using ELayout = remove_cvref_t; using CDElementwise = remove_cvref_t; diff --git a/projects/composablekernel/include/ck_tile/ops/flatmm.hpp b/projects/composablekernel/include/ck_tile/ops/flatmm.hpp index e08fac48c7e9..2e71957ac774 100644 --- a/projects/composablekernel/include/ck_tile/ops/flatmm.hpp +++ b/projects/composablekernel/include/ck_tile/ops/flatmm.hpp @@ -21,6 +21,7 @@ #include "ck_tile/ops/flatmm/pipeline/mx_flatmm_pipeline_agmem_bgmem_creg_v1.hpp" #include "ck_tile/ops/flatmm/pipeline/mx_flatmm_pipeline_agmem_bgmem_creg_v1_policy.hpp" #include "ck_tile/ops/flatmm/pipeline/tile_flatmm_shape.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/fmha.hpp b/projects/composablekernel/include/ck_tile/ops/fmha.hpp index 8a5d77bf462e..633a36b7ff68 100644 --- a/projects/composablekernel/include/ck_tile/ops/fmha.hpp +++ b/projects/composablekernel/include/ck_tile/ops/fmha.hpp @@ -61,6 +61,7 @@ #include "ck_tile/ops/fmha/pipeline/block_fmha_pipeline_qx_ks_vs_custom_policy.hpp" #include "ck_tile/ops/fmha/pipeline/tile_fmha_shape.hpp" #include "ck_tile/ops/fmha/pipeline/tile_fmha_traits.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/fused_moe.hpp b/projects/composablekernel/include/ck_tile/ops/fused_moe.hpp index 60f5bd1c4e35..2eb4abd64117 100644 --- a/projects/composablekernel/include/ck_tile/ops/fused_moe.hpp +++ b/projects/composablekernel/include/ck_tile/ops/fused_moe.hpp @@ -14,6 +14,7 @@ #include "ck_tile/ops/fused_moe/pipeline/fused_moegemm_traits.hpp" #include "ck_tile/ops/fused_moe/pipeline/moe_sorting_pipeline.hpp" #include "ck_tile/ops/fused_moe/pipeline/moe_sorting_policy.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/gemm.hpp b/projects/composablekernel/include/ck_tile/ops/gemm.hpp index 7c087e9186db..8f2125cb6748 100644 --- a/projects/composablekernel/include/ck_tile/ops/gemm.hpp +++ b/projects/composablekernel/include/ck_tile/ops/gemm.hpp @@ -84,6 +84,7 @@ #include "ck_tile/ops/gemm/warp/warp_gemm_smfmac_impl.hpp" #include "ck_tile/ops/gemm/warp/warp_wmma_gemm.hpp" #include "ck_tile/ops/gemm/warp/warp_wmma_gemm_gfx11_utils.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/gemm/block/block_universal_gemm_as_bs_cr.hpp b/projects/composablekernel/include/ck_tile/ops/gemm/block/block_universal_gemm_as_bs_cr.hpp index f7f5cd33dbb9..c5ed8931e7ee 100644 --- a/projects/composablekernel/include/ck_tile/ops/gemm/block/block_universal_gemm_as_bs_cr.hpp +++ b/projects/composablekernel/include/ck_tile/ops/gemm/block/block_universal_gemm_as_bs_cr.hpp @@ -95,12 +95,9 @@ struct BlockUniversalGemmAsBsCr using CDataType = remove_cvref_t; using ATypeToUse = - std::conditional_t, BDataType, ADataType>; - using BTypeToUse = std::conditional_t || - std::is_same_v || - sizeof(BDataType) < sizeof(ADataType), - ADataType, - BDataType>; + typename DetermineWarpPrecType::a_prec_type; + using BTypeToUse = + typename DetermineWarpPrecType::b_prec_type; using WarpGemm = remove_cvref_t; diff --git a/projects/composablekernel/include/ck_tile/ops/gemm/pipeline/gemm_universal_pipeline_ag_bg_cr_policy.hpp b/projects/composablekernel/include/ck_tile/ops/gemm/pipeline/gemm_universal_pipeline_ag_bg_cr_policy.hpp index b4a8e9e8cb48..0f5b12fc65dc 100644 --- a/projects/composablekernel/include/ck_tile/ops/gemm/pipeline/gemm_universal_pipeline_ag_bg_cr_policy.hpp +++ b/projects/composablekernel/include/ck_tile/ops/gemm/pipeline/gemm_universal_pipeline_ag_bg_cr_policy.hpp @@ -910,12 +910,10 @@ struct UniversalGemmPipelineAgBgCrPolicy using BDataType = remove_cvref_t; using ComputeDataType = remove_cvref_t; - using ATypeToUse = if_select_t; - using BTypeToUse = std::conditional_t || - std::is_same_v || - sizeof(BDataType) < sizeof(ADataType), - ADataType, - BDataType>; + using ATypeToUse = + typename DetermineWarpPrecType::a_prec_type; + using BTypeToUse = + typename DetermineWarpPrecType::b_prec_type; using WarpGemm = WarpGemmDispatcher, diff --git a/projects/composablekernel/include/ck_tile/ops/gemm_mx.hpp b/projects/composablekernel/include/ck_tile/ops/gemm_mx.hpp index 29fccf8057b9..fd04c1d0234a 100644 --- a/projects/composablekernel/include/ck_tile/ops/gemm_mx.hpp +++ b/projects/composablekernel/include/ck_tile/ops/gemm_mx.hpp @@ -6,6 +6,7 @@ #include "ck_tile/ops/gemm_mx/kernel/scale_pointer.hpp" #include "ck_tile/ops/gemm_mx/pipeline/gemm_pipeline_ag_bg_cr_comp_async.hpp" #include "ck_tile/ops/gemm_mx/pipeline/gemm_pipeline_ag_bg_cr_comp_async_default_policy.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/gemm_quant.hpp b/projects/composablekernel/include/ck_tile/ops/gemm_quant.hpp index 5b2ce7ff1915..907aabc0f1e7 100644 --- a/projects/composablekernel/include/ck_tile/ops/gemm_quant.hpp +++ b/projects/composablekernel/include/ck_tile/ops/gemm_quant.hpp @@ -33,6 +33,7 @@ #include "ck_tile/ops/gemm_quant/pipeline/gemm_wp_bquant_pipeline_ag_bg_cr_base_policy.hpp" #include "ck_tile/ops/gemm_quant/pipeline/gemm_wp_bquant_pipeline_ag_bg_cr_v2.hpp" #include "ck_tile/ops/gemm_quant/pipeline/tile_gemm_quant_traits.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/grouped_convolution.hpp b/projects/composablekernel/include/ck_tile/ops/grouped_convolution.hpp index 5bc4f0c6a042..1a115f9489e6 100644 --- a/projects/composablekernel/include/ck_tile/ops/grouped_convolution.hpp +++ b/projects/composablekernel/include/ck_tile/ops/grouped_convolution.hpp @@ -12,6 +12,7 @@ #include "ck_tile/ops/grouped_convolution/utils/transform_conv_bwd_data_to_gemm.hpp" #include "ck_tile/ops/grouped_convolution/utils/transform_conv_bwd_weight_to_gemm.hpp" #include "ck_tile/ops/grouped_convolution/utils/transform_conv_fwd_to_gemm.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/image_to_column.hpp b/projects/composablekernel/include/ck_tile/ops/image_to_column.hpp index 07d99890869e..faa165a8b001 100644 --- a/projects/composablekernel/include/ck_tile/ops/image_to_column.hpp +++ b/projects/composablekernel/include/ck_tile/ops/image_to_column.hpp @@ -5,6 +5,7 @@ #include "ck_tile/ops/image_to_column/kernel/image_to_column_kernel.hpp" #include "ck_tile/ops/image_to_column/pipeline/block_image_to_column_problem.hpp" #include "ck_tile/ops/image_to_column/pipeline/tile_image_to_column_shape.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/layernorm2d.hpp b/projects/composablekernel/include/ck_tile/ops/layernorm2d.hpp index 8f9ab205ac45..2266c138729a 100644 --- a/projects/composablekernel/include/ck_tile/ops/layernorm2d.hpp +++ b/projects/composablekernel/include/ck_tile/ops/layernorm2d.hpp @@ -8,6 +8,7 @@ #include "ck_tile/ops/layernorm2d/pipeline/layernorm2d_fwd_pipeline_problem.hpp" #include "ck_tile/ops/layernorm2d/pipeline/layernorm2d_fwd_pipeline_two_pass.hpp" #include "ck_tile/ops/layernorm2d/pipeline/layernorm2d_fwd_traits.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/norm_reduce.hpp b/projects/composablekernel/include/ck_tile/ops/norm_reduce.hpp index eae0ea14a337..9f572ff5cbb2 100644 --- a/projects/composablekernel/include/ck_tile/ops/norm_reduce.hpp +++ b/projects/composablekernel/include/ck_tile/ops/norm_reduce.hpp @@ -5,6 +5,7 @@ #include "ck_tile/ops/norm_reduce/block/block_norm_reduce.hpp" #include "ck_tile/ops/norm_reduce/block/block_norm_reduce_problem.hpp" #include "ck_tile/ops/norm_reduce/thread/thread_welford.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/permute.hpp b/projects/composablekernel/include/ck_tile/ops/permute.hpp index 4d37f4fbc12a..c7747a67e70c 100644 --- a/projects/composablekernel/include/ck_tile/ops/permute.hpp +++ b/projects/composablekernel/include/ck_tile/ops/permute.hpp @@ -4,6 +4,7 @@ #include "ck_tile/ops/permute/kernel/generic_permute_kernel.hpp" #include "ck_tile/ops/permute/pipeline/generic_petmute_problem.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/pooling.hpp b/projects/composablekernel/include/ck_tile/ops/pooling.hpp index faa77d53273e..43b24c7f8caf 100644 --- a/projects/composablekernel/include/ck_tile/ops/pooling.hpp +++ b/projects/composablekernel/include/ck_tile/ops/pooling.hpp @@ -6,6 +6,7 @@ #include "ck_tile/ops/pooling/pipeline/pool_default_policy.hpp" #include "ck_tile/ops/pooling/pipeline/pool_problem.hpp" #include "ck_tile/ops/pooling/pipeline/pool_shape.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/reduce.hpp b/projects/composablekernel/include/ck_tile/ops/reduce.hpp index b5e53283e485..e680d2574538 100644 --- a/projects/composablekernel/include/ck_tile/ops/reduce.hpp +++ b/projects/composablekernel/include/ck_tile/ops/reduce.hpp @@ -13,6 +13,7 @@ #include "ck_tile/ops/reduce/pipeline/reduce2d_default_policy.hpp" #include "ck_tile/ops/reduce/pipeline/reduce2d_problem.hpp" #include "ck_tile/ops/reduce/pipeline/reduce2d_shape.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/rmsnorm2d.hpp b/projects/composablekernel/include/ck_tile/ops/rmsnorm2d.hpp index f271be50068c..7ee67334d6b9 100644 --- a/projects/composablekernel/include/ck_tile/ops/rmsnorm2d.hpp +++ b/projects/composablekernel/include/ck_tile/ops/rmsnorm2d.hpp @@ -9,6 +9,7 @@ #include "ck_tile/ops/rmsnorm2d/pipeline/rmsnorm2d_fwd_pipeline_problem.hpp" #include "ck_tile/ops/rmsnorm2d/pipeline/rmsnorm2d_fwd_pipeline_two_pass.hpp" #include "ck_tile/ops/rmsnorm2d/pipeline/rmsnorm2d_fwd_traits.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/smoothquant.hpp b/projects/composablekernel/include/ck_tile/ops/smoothquant.hpp index 4c2fe9bee434..ad984c033f02 100644 --- a/projects/composablekernel/include/ck_tile/ops/smoothquant.hpp +++ b/projects/composablekernel/include/ck_tile/ops/smoothquant.hpp @@ -8,6 +8,7 @@ #include "ck_tile/ops/smoothquant/pipeline/smoothquant_pipeline_one_pass.hpp" #include "ck_tile/ops/smoothquant/pipeline/smoothquant_pipeline_problem.hpp" #include "ck_tile/ops/smoothquant/pipeline/smoothquant_pipeline_two_pass.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/softmax.hpp b/projects/composablekernel/include/ck_tile/ops/softmax.hpp index c79ba06abfea..b810a57dda0d 100644 --- a/projects/composablekernel/include/ck_tile/ops/softmax.hpp +++ b/projects/composablekernel/include/ck_tile/ops/softmax.hpp @@ -4,6 +4,7 @@ #include "ck_tile/ops/softmax/block/block_softmax_2d.hpp" #include "ck_tile/ops/softmax/block/block_softmax_2d_problem.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/sparse_attn.hpp b/projects/composablekernel/include/ck_tile/ops/sparse_attn.hpp index c7c4171874aa..56e074794ddf 100644 --- a/projects/composablekernel/include/ck_tile/ops/sparse_attn.hpp +++ b/projects/composablekernel/include/ck_tile/ops/sparse_attn.hpp @@ -6,6 +6,7 @@ #include "ck_tile/ops/sparse_attn/kernel/fmha_fwd_vsa_kernel.hpp" #include "ck_tile/ops/sparse_attn/pipeline/block_fmha_pipeline_qr_ks_vs_async_jenga.hpp" #include "ck_tile/ops/sparse_attn/pipeline/block_fmha_pipeline_qr_ks_vs_async_vsa.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/topk.hpp b/projects/composablekernel/include/ck_tile/ops/topk.hpp index 474ba932270c..13d818174e62 100644 --- a/projects/composablekernel/include/ck_tile/ops/topk.hpp +++ b/projects/composablekernel/include/ck_tile/ops/topk.hpp @@ -4,6 +4,7 @@ #include "ck_tile/ops/topk/block/block_topk_stream_2d.hpp" #include "ck_tile/ops/topk/block/block_topk_stream_2d_problem.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" diff --git a/projects/composablekernel/include/ck_tile/ops/topk_softmax.hpp b/projects/composablekernel/include/ck_tile/ops/topk_softmax.hpp index 066fbf5feea2..b7219511faa3 100644 --- a/projects/composablekernel/include/ck_tile/ops/topk_softmax.hpp +++ b/projects/composablekernel/include/ck_tile/ops/topk_softmax.hpp @@ -6,6 +6,7 @@ #include "ck_tile/ops/topk_softmax/pipeline/topk_softmax_warp_per_row_pipeline.hpp" #include "ck_tile/ops/topk_softmax/pipeline/topk_softmax_warp_per_row_policy.hpp" #include "ck_tile/ops/topk_softmax/pipeline/topk_softmax_warp_per_row_problem.hpp" +#include "ck_tile/ops/common/determine_warp_prec_type.hpp" #include "ck_tile/ops/common/generic_2d_block_shape.hpp" #include "ck_tile/ops/common/load_and_convert_tile.hpp" #include "ck_tile/ops/common/streamk_common.hpp" From b6a16fab00686eee4ecd3243a90ed9363a70fb8c Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Thu, 9 Oct 2025 09:04:13 +0000 Subject: [PATCH 22/23] Add functionality and tests for bf16 x fp8 and fp8 x bf16 --- .../test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/projects/composablekernel/test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp b/projects/composablekernel/test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp index 47a0267020e7..f9a61d276940 100644 --- a/projects/composablekernel/test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp +++ b/projects/composablekernel/test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp @@ -95,9 +95,11 @@ using KernelTypesCompV3 = ::testing::Types< std::tuple< Row, Col, Row, F16, F16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F16, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, BF16, BF16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, + std::tuple< Row, Col, Row, BF16, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, BF16, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, INT8, INT8, INT32, INT32, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F8, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, + std::tuple< Row, Col, Row, F8, BF16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F8, BF8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F8, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, BF8, BF8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, @@ -115,9 +117,11 @@ using KernelTypesCompV3 = ::testing::Types< std::tuple< Col, Col, Row, F16, F16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F16, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, BF16, BF16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, + std::tuple< Col, Col, Row, BF16, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, BF16, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, INT8, INT8, INT32, INT32, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F8, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, + std::tuple< Col, Col, Row, F8, BF16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F8, BF8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F8, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, BF8, BF8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, From 26ce440487d3354ab80093873cbabc1be50015c4 Mon Sep 17 00:00:00 2001 From: Sami Aario Date: Wed, 12 Nov 2025 15:09:01 +0000 Subject: [PATCH 23/23] Add functionality and tests for fp16 x fp8 and fp8 x fp16 --- .../test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/projects/composablekernel/test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp b/projects/composablekernel/test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp index f9a61d276940..98d365f552ac 100644 --- a/projects/composablekernel/test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp +++ b/projects/composablekernel/test/ck_tile/gemm/test_gemm_pipeline_kernel_types.hpp @@ -93,12 +93,14 @@ using KernelTypesCompV3 = ::testing::Types< std::tuple< Row, Row, Row, BF8, BF8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Row, Row, BF8, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F16, F16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, + std::tuple< Row, Col, Row, F16, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F16, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, BF16, BF16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, BF16, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, BF16, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, INT8, INT8, INT32, INT32, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F8, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, + std::tuple< Row, Col, Row, F8, F16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F8, BF16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F8, BF8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Row, Col, Row, F8, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, @@ -115,12 +117,14 @@ using KernelTypesCompV3 = ::testing::Types< std::tuple< Col, Row, Row, BF8, BF8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Row, Row, BF8, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F16, F16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, + std::tuple< Col, Col, Row, F16, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F16, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, BF16, BF16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, BF16, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, BF16, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, INT8, INT8, INT32, INT32, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F8, F8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, + std::tuple< Col, Col, Row, F8, F16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F8, BF16, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F8, BF8, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>, std::tuple< Col, Col, Row, F8, I4, F32, F16, I256, I256, I64, I32, I32, I16, Intrawave, CompV3>,