CUTLASS 3.2.1 (#1113)

* Updates for 3.2.1 release.

* Minor fix in gemm op profiler for raster order.

* Add scheduler mapping for raster order in the kernels.
This commit is contained in:
ANIKET SHIVAM
2023-09-26 17:24:26 -04:00
committed by GitHub
parent e0aaa3c3b3
commit 90d3b0fb18
428 changed files with 22252 additions and 21761 deletions
+44 -60
View File
@@ -130,7 +130,7 @@ struct Mma<
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -196,7 +196,7 @@ struct Mma<
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -257,13 +257,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k16.row.col.s32.s8.s8.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -318,13 +317,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k16.row.col.s32.u8.s8.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -379,14 +377,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k16.row.col.s8.u8 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -441,13 +437,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k16.row.col.s32.u8.u8.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -461,7 +456,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = S8 * S8 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,16>,
gemm::GemmShape<8, 8, 16>,
32,
int8_t,
layout::RowMajor,
@@ -471,7 +466,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAddSaturate> {
using Shape = gemm::GemmShape<8,8,16>;
using Shape = gemm::GemmShape<8, 8, 16>;
using ElementA = int8_t;
using LayoutA = layout::RowMajor;
@@ -508,13 +503,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k16.row.col.satfinite.s32.s8.s8.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -522,7 +516,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = U8 * S8 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,16>,
gemm::GemmShape<8, 8, 16>,
32,
uint8_t,
layout::RowMajor,
@@ -532,7 +526,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAddSaturate> {
using Shape = gemm::GemmShape<8,8,16>;
using Shape = gemm::GemmShape<8, 8, 16>;
using ElementA = uint8_t;
using LayoutA = layout::RowMajor;
@@ -569,13 +563,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k16.row.col.satfinite.s32.u8.s8.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -583,7 +576,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = S8 * U8 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,16>,
gemm::GemmShape<8, 8, 16>,
32,
int8_t,
layout::RowMajor,
@@ -593,7 +586,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAddSaturate> {
using Shape = gemm::GemmShape<8,8,16>;
using Shape = gemm::GemmShape<8, 8, 16>;
using ElementA = int8_t;
using LayoutA = layout::RowMajor;
@@ -630,13 +623,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k16.row.col.satfinite.s32.s8.u8.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -644,7 +636,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = U8 * U8 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,16>,
gemm::GemmShape<8, 8, 16>,
32,
uint8_t,
layout::RowMajor,
@@ -654,7 +646,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAddSaturate> {
using Shape = gemm::GemmShape<8,8,16>;
using Shape = gemm::GemmShape<8, 8, 16>;
using ElementA = uint8_t;
using LayoutA = layout::RowMajor;
@@ -691,13 +683,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k16.row.col.satfinite.s32.u8.u8.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -711,7 +702,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = S4 * S4 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,32>,
gemm::GemmShape<8, 8, 32>,
32,
int4b_t,
layout::RowMajor,
@@ -721,7 +712,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAdd> {
using Shape = gemm::GemmShape<8,8,32>;
using Shape = gemm::GemmShape<8, 8, 32>;
using ElementA = int4b_t;
using LayoutA = layout::RowMajor;
@@ -751,19 +742,19 @@ struct Mma<
unsigned const & A = reinterpret_cast<unsigned const &>(a);
unsigned const & B = reinterpret_cast<unsigned const &>(b);
int const *C = reinterpret_cast<int const *>(&c);
int *D = reinterpret_cast<int *>(&d);
asm volatile("mma.sync.aligned.m8n8k32.row.col.s32.s4.s4.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -771,7 +762,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = U4 * S4 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,32>,
gemm::GemmShape<8, 8, 32>,
32,
uint4b_t,
layout::RowMajor,
@@ -781,7 +772,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAdd> {
using Shape = gemm::GemmShape<8,8,32>;
using Shape = gemm::GemmShape<8, 8, 32>;
using ElementA = uint4b_t;
using LayoutA = layout::RowMajor;
@@ -818,13 +809,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k32.row.col.s32.u4.s4.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -832,7 +822,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = S4 * U4 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,32>,
gemm::GemmShape<8, 8, 32>,
32,
int4b_t,
layout::RowMajor,
@@ -842,7 +832,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAdd> {
using Shape = gemm::GemmShape<8,8,32>;
using Shape = gemm::GemmShape<8, 8, 32>;
using ElementA = int4b_t;
using LayoutA = layout::RowMajor;
@@ -879,13 +869,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k32.row.col.s32.s4.u4.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -893,7 +882,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = U4 * U4 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,32>,
gemm::GemmShape<8, 8, 32>,
32,
uint4b_t,
layout::RowMajor,
@@ -903,7 +892,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAdd> {
using Shape = gemm::GemmShape<8,8,32>;
using Shape = gemm::GemmShape<8, 8, 32>;
using ElementA = uint4b_t;
using LayoutA = layout::RowMajor;
@@ -940,13 +929,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k32.row.col.s32.u4.u4.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -960,7 +948,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = S4 * S4 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,32>,
gemm::GemmShape<8, 8, 32>,
32,
int4b_t,
layout::RowMajor,
@@ -970,7 +958,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAddSaturate> {
using Shape = gemm::GemmShape<8,8,32>;
using Shape = gemm::GemmShape<8, 8, 32>;
using ElementA = int4b_t;
using LayoutA = layout::RowMajor;
@@ -1007,13 +995,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k32.row.col.satfinite.s32.s4.s4.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -1021,7 +1008,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = U4 * S4 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,32>,
gemm::GemmShape<8, 8, 32>,
32,
uint4b_t,
layout::RowMajor,
@@ -1031,7 +1018,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAddSaturate> {
using Shape = gemm::GemmShape<8,8,32>;
using Shape = gemm::GemmShape<8, 8, 32>;
using ElementA = uint4b_t;
using LayoutA = layout::RowMajor;
@@ -1068,13 +1055,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k32.row.col.satfinite.s32.u4.s4.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -1082,7 +1068,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = S4 * U4 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,32>,
gemm::GemmShape<8, 8, 32>,
32,
int4b_t,
layout::RowMajor,
@@ -1092,7 +1078,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAddSaturate> {
using Shape = gemm::GemmShape<8,8,32>;
using Shape = gemm::GemmShape<8, 8, 32>;
using ElementA = int4b_t;
using LayoutA = layout::RowMajor;
@@ -1129,13 +1115,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k32.row.col.satfinite.s32.s4.u4.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -1143,7 +1128,7 @@ struct Mma<
/// Matrix multiply-add operation: S32 = U4 * U4 + S32
template <>
struct Mma<
gemm::GemmShape<8,8,32>,
gemm::GemmShape<8, 8, 32>,
32,
uint4b_t,
layout::RowMajor,
@@ -1153,7 +1138,7 @@ struct Mma<
layout::RowMajor,
OpMultiplyAddSaturate> {
using Shape = gemm::GemmShape<8,8,32>;
using Shape = gemm::GemmShape<8, 8, 32>;
using ElementA = uint4b_t;
using LayoutA = layout::RowMajor;
@@ -1190,13 +1175,12 @@ struct Mma<
asm volatile("mma.sync.aligned.m8n8k32.row.col.satfinite.s32.u4.u4.s32 {%0,%1}, {%2}, {%3}, {%4,%5};\n"
: "=r"(D[0]), "=r"(D[1])
: "r"(A), "r"(B), "r"(C[0]), "r"(C[1]));
#else
CUTLASS_UNUSED(a);
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0);
CUTLASS_NOT_IMPLEMENTED();
#endif
}
};
@@ -1287,7 +1271,7 @@ struct Mma<
CUTLASS_UNUSED(b);
CUTLASS_UNUSED(c);
CUTLASS_UNUSED(d);
assert(0); // WMMA must be supported to issue binary matrix multiply-accumulate instructions.
CUTLASS_NOT_IMPLEMENTED(); // WMMA must be supported to issue binary matrix multiply-accumulate instructions.
#endif // defined(CUTLASS_ARCH_WMMA_ENABLED)
+13 -4
View File
@@ -53,7 +53,16 @@
#if (defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800))
#define CUTLASS_ARCH_MMA_SM80_ENABLED
#if (__CUDA_ARCH__ <= 900)
#define CUTLASS_ARCH_MMA_B1_AND_SM80_ENABLED
#endif
#if (__CUDA_ARCH__ <= 890)
#define CUTLASS_ARCH_MMA_B1_XOR_SM80_ENABLED
#endif
#endif
#endif
////////////////////////////////////////////////////////////////////////////////
@@ -2084,7 +2093,7 @@ struct Mma<
FragmentC const &c
) const {
#if defined(CUTLASS_ARCH_MMA_SM80_ENABLED)
#if defined(CUTLASS_ARCH_MMA_B1_AND_SM80_ENABLED)
uint32_t const *A = reinterpret_cast<uint32_t const *>(&a);
uint32_t const *B = reinterpret_cast<uint32_t const *>(&b);
@@ -2149,7 +2158,7 @@ struct Mma<
FragmentC const &c
) const {
#if defined(CUTLASS_ARCH_MMA_SM80_ENABLED)
#if defined(CUTLASS_ARCH_MMA_B1_AND_SM80_ENABLED)
uint32_t const *A = reinterpret_cast<uint32_t const *>(&a);
uint32_t const *B = reinterpret_cast<uint32_t const *>(&b);
@@ -2220,7 +2229,7 @@ struct Mma<
FragmentC const &c
) const {
#if defined(CUTLASS_ARCH_MMA_SM80_ENABLED)
#if defined(CUTLASS_ARCH_MMA_B1_XOR_SM80_ENABLED)
uint32_t const *A = reinterpret_cast<uint32_t const *>(&a);
uint32_t const *B = reinterpret_cast<uint32_t const *>(&b);
@@ -2244,7 +2253,7 @@ struct Mma<
CUTLASS_UNUSED(d);
assert(0);
#endif // defined(CUTLASS_ARCH_MMA_SM80_ENABLED)
#endif // defined(CUTLASS_ARCH_MMA_B1_XOR_SM80_ENABLED)
}
};