CUTLASS 2.6 (#298)

CUTLASS 2.6
This commit is contained in:
Manish Gupta
2021-07-23 00:40:53 -04:00
committed by GitHub
parent 6c29fe20ba
commit e5d51840e8
308 changed files with 32408 additions and 4722 deletions
@@ -22,6 +22,7 @@
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*! \file
\brief Unit tests for threadblock-level GEMM
*/
@@ -3824,4 +3825,5 @@ TEST(SM80_gemm_threadblock_crosswise_f64,
}
////////////////////////////////////////////////////////////////////////////////
#endif
@@ -59,7 +59,8 @@ __global__ void kernel_multistage_mma_sparse(cutlass::gemm::GemmCoord problem_si
typename Mma::IteratorA::TensorRef ref_A,
typename Mma::IteratorB::Params params_B,
typename Mma::IteratorB::TensorRef ref_B,
typename Mma::ElementC *ptr_C, int ldc,
typename Mma::ElementC *ptr_C,
typename Mma::LayoutC::Stride::Index ldc,
typename Mma::IteratorE::Params params_E,
typename Mma::IteratorE::TensorRef ref_E) {
// Shared storage needed by threadblock-scoped matrix multiply-
@@ -57,7 +57,8 @@ __global__ void kernel_multistage_mma(cutlass::gemm::GemmCoord problem_size,
typename Mma::IteratorA::TensorRef ref_A,
typename Mma::IteratorB::Params params_B,
typename Mma::IteratorB::TensorRef ref_B,
typename Mma::ElementC *ptr_C, int ldc) {
typename Mma::ElementC *ptr_C,
typename Mma::LayoutC::Stride::Index ldc) {
// Shared storage needed by threadblock-scoped matrix multiply-accumulate
// Dynamic shared memory base pointer
@@ -67,7 +67,8 @@ __global__ void kernel_mma(cutlass::gemm::GemmCoord problem_size,
typename Mma::IteratorA::TensorRef ref_A,
typename Mma::IteratorB::Params params_B,
typename Mma::IteratorB::TensorRef ref_B,
typename Mma::ElementC *ptr_C, int ldc) {
typename Mma::ElementC *ptr_C,
typename Mma::LayoutC::Stride::Index ldc) {
// Shared storage needed by threadblock-scoped matrix multiply-accumulate
__shared__ typename Mma::SharedStorage shared_storage;
@@ -67,7 +67,8 @@ __global__ void kernel_mma_planar_complex(
typename Mma::IteratorB::Params params_B,
typename Mma::IteratorB::Element *ptr_B,
int64_t imaginary_stride_B,
typename Mma::ElementC *ptr_C, int ldc, int64_t imaginary_stride_C) {
typename Mma::ElementC *ptr_C,
typename Mma::LayoutC::Stride::Index ldc, int64_t imaginary_stride_C) {
// Shared storage needed by threadblock-scoped matrix multiply-accumulate
__shared__ typename Mma::SharedStorage shared_storage;