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:
co-authored by
Pradeep Ramani
parent
b7508e3379
commit
8236f30675
@@ -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
|
||||
//
|
||||
|
||||
Reference in New Issue
Block a user