@@ -63,11 +63,17 @@ template <index_t NDimSpatial,
6363 typename DsDataTypes = Tuple<>,
6464 typename OutElementOp = PassThrough>
6565using device_grouped_conv_fwd_xdl_bf16_comp_instances_2x = std::tuple<
66- // clang-format off
66+ // clang-format off
6767 // ########################################| NumDim| A| B| Ds| E| AData| BData| AccData| CShuffle| Ds| EData| A| B| CDE| ConvForward| GEMM| NumGemmK| Block| MPer| NPer| KPer| AK1| BK1| MPer| NPer| MXdl| NXdl| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockLds| BBlockTransfer| BBlockTransfer| BBlockTransfer| BlockTransfer| BBlockTransfer| BBlockTransfer| BBlockLds| CShuffle| CShuffle| CBlockTransferClusterLengths| CBlockTransfer|
6868 // ########################################| Spatial| Layout| Layout| Layout| Layout| Type| Type| Type| DataType| DataType| Type| Elementwise| Elementwise| Elementwise| Specialization| Specialization| Prefetch| Size| Block| Block| Block| | | XDL| XDL| Per| Per| ThreadCluster| ThreadCluster| SrcAccessOrder| SrcVectorDim| SrcScalar| DstScalar| AddExtraM| ThreadCluster| ThreadCluster| SrcAccessOrder| SrcVectorDim| SrcScalar| DstScalar| AddExtraN| MXdlPerWave| NXdlPerWave| _MBlock_MWaveMPerXdl| ScalarPerVector|
6969 // ########################################| | | | | | | | | | | | Operation| Operation| Operation| | | Stage| | | | | | | | | Wave| Wave| Lengths_K0_M_K1| ArrangeOrder| | | PerVector| PerVector_K1| | Lengths_K0_N_K1| ArrangeOrder| | | PerVector| PerVector_K1| | PerShuffle| PerShuffle| _NBlock_NWaveNPerXdl| _NWaveNPerXdl|
7070 // ########################################| | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | |
71+ #if defined(CK_USE_GFX950)
72+ // Instances optimized for G=1
73+ // The use more shared memory than what is available for non gfx950 architectures.
74+ DeviceGroupedConvFwdMultipleABD_Xdl_CShuffle_V3<NDimSpatial,ALayout,BLayout, DsLayout,ELayout, BF16 , BF16 , F32 , BF16 , DsDataTypes, BF16 , PassThrough, PassThrough, OutElementOp, ConvSpec, GemmMNKPadding, 256 , 256 , 256 , 64 , 8 , 8 , 32 , 32 , 4 , 4 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , 1 , 1 , S<1 , 32 , 1 , 8 >, 8 , BlockGemmPipelineScheduler::Intrawave, BlockGemmPipelineVersion::v3>,
75+ DeviceGroupedConvFwdMultipleABD_Xdl_CShuffle_V3<NDimSpatial,ALayout,BLayout, DsLayout,ELayout, BF16 , BF16 , F32 , BF16 , DsDataTypes, BF16 , PassThrough, PassThrough, OutElementOp, ConvSpec, GemmMNKPadding, 256 , 512 , 128 , 32 , 8 , 8 , 32 , 32 , 8 , 2 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , 1 , 1 , S<1 , 32 , 1 , 8 >, 8 , BlockGemmPipelineScheduler::Intrawave, BlockGemmPipelineVersion::v3>,
76+ #endif
7177 DeviceGroupedConvFwdMultipleABD_Xdl_CShuffle_V3<NDimSpatial,ALayout,BLayout, DsLayout,ELayout, BF16 , BF16 , F32 , BF16 , DsDataTypes, BF16 , PassThrough, PassThrough, OutElementOp, ConvSpec, GemmMNKPadding, 256 , 128 , 128 , 64 , 16 , 16 , 32 , 32 , 2 , 2 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , 1 , 1 , S<1 , 32 , 1 , 8 >, 8 , BlockGemmPipelineScheduler::Interwave, BlockGemmPipelineVersion::v1>
7278 // clang-format on
7379 >;
@@ -131,11 +137,17 @@ template <index_t NDimSpatial,
131137 typename DsDataTypes = Tuple<>,
132138 typename OutElementOp = PassThrough>
133139using device_grouped_conv_fwd_xdl_f16_comp_instances_2x = std::tuple<
134- // clang-format off
140+ // clang-format off
135141 // ########################################| NumDim| A| B| Ds| E| AData| BData| AccData| CShuffle| Ds| EData| A| B| CDE| ConvForward| GEMM| NumGemmK| Block| MPer| NPer| KPer| AK1| BK1| MPer| NPer| MXdl| NXdl| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockTransfer| ABlockLds| BBlockTransfer| BBlockTransfer| BBlockTransfer| BlockTransfer| BBlockTransfer| BBlockTransfer| BBlockLds| CShuffle| CShuffle| CBlockTransferClusterLengths| CBlockTransfer|
136142 // ########################################| Spatial| Layout| Layout| Layout| Layout| Type| Type| Type| DataType| DataType| Type| Elementwise| Elementwise| Elementwise| Specialization| Specialization| Prefetch| Size| Block| Block| Block| | | XDL| XDL| Per| Per| ThreadCluster| ThreadCluster| SrcAccessOrder| SrcVectorDim| SrcScalar| DstScalar| AddExtraM| ThreadCluster| ThreadCluster| SrcAccessOrder| SrcVectorDim| SrcScalar| DstScalar| AddExtraN| MXdlPerWave| NXdlPerWave| _MBlock_MWaveMPerXdl| ScalarPerVector|
137143 // ########################################| | | | | | | | | | | | Operation| Operation| Operation| | | Stage| | | | | | | | | Wave| Wave| Lengths_K0_M_K1| ArrangeOrder| | | PerVector| PerVector_K1| | Lengths_K0_N_K1| ArrangeOrder| | | PerVector| PerVector_K1| | PerShuffle| PerShuffle| _NBlock_NWaveNPerXdl| _NWaveNPerXdl|
138144 // ########################################| | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | |
145+ #if defined(CK_USE_GFX950)
146+ // Instances optimized for G=1
147+ // The use more shared memory than what is available for non gfx950 architectures.
148+ DeviceGroupedConvFwdMultipleABD_Xdl_CShuffle_V3<NDimSpatial,ALayout,BLayout, DsLayout,ELayout, F16 , F16 , F32 , F16 , DsDataTypes, F16 , PassThrough, PassThrough, OutElementOp, ConvSpec, GemmMNKPadding, 256 , 256 , 256 , 64 , 8 , 8 , 32 , 32 , 4 , 4 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , 1 , 1 , S<1 , 32 , 1 , 8 >, 8 , BlockGemmPipelineScheduler::Intrawave, BlockGemmPipelineVersion::v3>,
149+ DeviceGroupedConvFwdMultipleABD_Xdl_CShuffle_V3<NDimSpatial,ALayout,BLayout, DsLayout,ELayout, F16 , F16 , F32 , F16 , DsDataTypes, F16 , PassThrough, PassThrough, OutElementOp, ConvSpec, GemmMNKPadding, 256 , 512 , 128 , 32 , 8 , 8 , 32 , 32 , 8 , 2 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , 1 , 1 , S<1 , 32 , 1 , 8 >, 8 , BlockGemmPipelineScheduler::Intrawave, BlockGemmPipelineVersion::v3>,
150+ #endif
139151 DeviceGroupedConvFwdMultipleABD_Xdl_CShuffle_V3<NDimSpatial,ALayout,BLayout, DsLayout,ELayout, F16 , F16 , F32 , F16 , DsDataTypes, F16 , PassThrough, PassThrough, OutElementOp, ConvSpec, GemmMNKPadding, 256 , 128 , 128 , 64 , 16 , 16 , 32 , 32 , 2 , 2 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , S<4 , 64 , 1 >, S<1 , 0 , 2 >, S<1 , 0 , 2 >, 2 , 8 , 8 , 1 , 1 , 1 , S<1 , 32 , 1 , 8 >, 8 , BlockGemmPipelineScheduler::Interwave, BlockGemmPipelineVersion::v1>
140152 // clang-format on
141153 >;
0 commit comments