CUTLASS 3.4.0 (#1286)

* CUTLASS 3.4.0

* Update CHANGELOG.md

---------

Co-authored-by: Pradeep Ramani <prramani@nvidia.com>
This commit is contained in:
Pradeep Ramani
2023-12-29 15:21:31 -05:00
committed by GitHub
co-authored by Pradeep Ramani
parent b7508e3379
commit 8236f30675
211 changed files with 11409 additions and 2763 deletions
@@ -49,7 +49,6 @@ namespace cutlass {
namespace epilogue {
namespace thread {
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace detail {
@@ -66,13 +65,65 @@ struct ArrayMaximum {
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < ElementsPerAccess; ++i) {
result[i] = fmax(lhs[i], rhs[i]);
result[i] = platform::max(lhs[i].get(), rhs[i]);
}
return result;
}
CUTLASS_HOST_DEVICE
Array<Element, ElementsPerAccess> operator()(
Array<Element, ElementsPerAccess> const &lhs,
Element rhs) const {
Array<Element, ElementsPerAccess> result;
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < ElementsPerAccess; ++i) {
result[i] = platform::max(lhs[i].get(), rhs);
}
return result;
}
};
/// Partial specialization: Element=float
template <int ElementsPerAccess>
struct ArrayMaximum<float, ElementsPerAccess> {
CUTLASS_HOST_DEVICE
Array<float, ElementsPerAccess> operator()(
Array<float, ElementsPerAccess> const &lhs,
Array<float, ElementsPerAccess> const &rhs) const {
Array<float, ElementsPerAccess> result;
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < ElementsPerAccess; ++i) {
result[i] = fmax(lhs[i], rhs[i]);
}
return result;
}
CUTLASS_HOST_DEVICE
Array<float, ElementsPerAccess> operator()(
Array<float, ElementsPerAccess> const &lhs,
float rhs) const {
Array<float, ElementsPerAccess> result;
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < ElementsPerAccess; ++i) {
result[i] = fmax(lhs[i], rhs);
}
return result;
}
};
/// Partial specialization: Element=half
template <int ElementsPerAccess>
struct ArrayMaximum<half_t, ElementsPerAccess> {
@@ -96,6 +147,8 @@ struct ArrayMaximum<half_t, ElementsPerAccess> {
res_ptr[i] = __hmax2(lhs_ptr[i], rhs_ptr[i]);
}
static_assert(!(ElementsPerAccess % 2), "Output array must be divisible by vector length.");
#else
__half const *lhs_ptr = reinterpret_cast<__half const *>(lhs.raw_data());
__half const *rhs_ptr = reinterpret_cast<__half const *>(rhs.raw_data());
@@ -133,6 +186,8 @@ struct ArrayMaximum<half_t, ElementsPerAccess> {
res_ptr[i] = __hmax2(lhs_ptr[i], rhs_pair);
}
static_assert(!(ElementsPerAccess % 2), "Output array must be divisible by vector length.");
#else
__half const *lhs_ptr = reinterpret_cast<__half const *>(lhs.raw_data());
@@ -150,6 +205,90 @@ struct ArrayMaximum<half_t, ElementsPerAccess> {
}
};
/// Partial specialization: Element=bfloat16_t
template <int ElementsPerAccess>
struct ArrayMaximum<bfloat16_t, ElementsPerAccess> {
using NvType = __nv_bfloat16;
using NvTypeV2 = __nv_bfloat162;
CUTLASS_DEVICE
Array<bfloat16_t, ElementsPerAccess> operator()(
Array<bfloat16_t, ElementsPerAccess> const &lhs,
Array<bfloat16_t, ElementsPerAccess> const &rhs) const {
Array<bfloat16_t, ElementsPerAccess> result;
#if __CUDA_ARCH__ >= 800
int const kVectorCount = ElementsPerAccess / 2;
NvTypeV2 const *lhs_ptr = reinterpret_cast<NvTypeV2 const *>(lhs.raw_data());
NvTypeV2 const *rhs_ptr = reinterpret_cast<NvTypeV2 const *>(rhs.raw_data());
NvTypeV2 *res_ptr = reinterpret_cast<NvTypeV2 *>(result.raw_data());
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < kVectorCount; ++i) {
res_ptr[i] = __hmax2(lhs_ptr[i], rhs_ptr[i]);
}
#else
NvType const *lhs_ptr = reinterpret_cast<NvType const *>(lhs.raw_data());
NvType const *rhs_ptr = reinterpret_cast<NvType const *>(rhs.raw_data());
NvType *res_ptr = reinterpret_cast<NvType *>(result.raw_data());
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < ElementsPerAccess; ++i) {
res_ptr[i] = ((lhs_ptr[i] < rhs_ptr[i]) ? rhs_ptr[i] : lhs_ptr[i]);
}
#endif
return result;
}
CUTLASS_DEVICE
Array<bfloat16_t, ElementsPerAccess> operator()(
Array<bfloat16_t, ElementsPerAccess> const &lhs,
bfloat16_t rhs) const {
Array<bfloat16_t, ElementsPerAccess> result;
#if __CUDA_ARCH__ >= 800
int const kVectorCount = ElementsPerAccess / 2;
NvType rhs_raw = reinterpret_cast<NvType const &>(rhs);
NvTypeV2 rhs_pair = __bfloat162bfloat162(rhs_raw);
NvTypeV2 const *lhs_ptr = reinterpret_cast<NvTypeV2 const *>(lhs.raw_data());
NvTypeV2 *res_ptr = reinterpret_cast<NvTypeV2 *>(result.raw_data());
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < kVectorCount; ++i) {
res_ptr[i] = __hmax2(lhs_ptr[i], rhs_pair);
}
static_assert(!(ElementsPerAccess % 2), "Output array must be divisible by vector length.");
#else
NvType const *lhs_ptr = reinterpret_cast<NvType const *>(lhs.raw_data());
NvType const rhs_raw = reinterpret_cast<NvType const &>(rhs);
NvType *res_ptr = reinterpret_cast<NvType *>(result.raw_data());
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < ElementsPerAccess; ++i) {
res_ptr[i] = ((lhs_ptr[i] < rhs_raw) ? rhs_raw : lhs_ptr[i]);
}
#endif
return result;
}
};
/////////////////////////////////////////////////////////////////////////////////////////////////
template <typename Element, int ElementsPerAccess>
@@ -187,6 +326,25 @@ struct ReluConditional<half_t, ElementsPerAccess> {
}
};
template <int ElementsPerAccess>
struct ReluConditional<bfloat16_t, ElementsPerAccess> {
CUTLASS_DEVICE
void operator()(
bool conditional[],
Array<bfloat16_t, ElementsPerAccess> const &fragment,
bfloat16_t threshold) const {
__nv_bfloat16 y = reinterpret_cast<__nv_bfloat16 const &>(threshold);
__nv_bfloat16 const *x = reinterpret_cast<__nv_bfloat16 const *>(fragment.raw_data());
CUTLASS_PRAGMA_UNROLL
for (int i = 0; i < ElementsPerAccess; ++i) {
conditional[i] = !__hlt(x[i], y);
}
}
};
} // namespace detail
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -92,9 +92,9 @@ public:
ElementCompute alpha; ///< scales accumulators
ElementCompute beta; ///< scales source tensor
ElementCompute threshold; ///< minimum value that is output
ElementCompute const *alpha_ptr; ///< pointer to accumulator scalar - if not null, loads it from memory
ElementCompute const *beta_ptr; ///< pointer to source scalar - if not null, loads it from memory
ElementCompute threshold; ///< minimum value that is output
//
// Methods
//
@@ -87,9 +87,9 @@ public:
ElementCompute alpha; ///< scales accumulators
ElementCompute beta; ///< scales source tensor
ElementCompute threshold; ///< minimum value that is output
ElementCompute const *alpha_ptr; ///< pointer to accumulator scalar - if not null, loads it from memory
ElementCompute const *beta_ptr; ///< pointer to source scalar - if not null, loads it from memory
ElementCompute threshold; ///< minimum value that is output
//
// Methods
//