CUTLASS 2.4 (Implicit GEMM convolution) (#147)
CUTLASS 2.4 (Implicit GEMM Convolution) Co-authored-by: Manish Gupta <manigupta@nvidia.com>, Haicheng Wu <haichengw@nvidia.com>, Dustyn Blasig <dblasig@nvidia.com>, Andrew Kerr <akerr@nvidia.com>
This commit is contained in:
co-authored by
Manish Gupta <manigupta@nvidia.com>, Haicheng Wu <haichengw@nvidia.com>, Dustyn Blasig <dblasig@nvidia.com>, Andrew Kerr <akerr@nvidia.com>
parent
c2b80ad4e4
commit
6615010cd0
@@ -161,6 +161,42 @@ struct negate {
|
||||
}
|
||||
};
|
||||
|
||||
/// Greater equal
|
||||
template <typename T>
|
||||
struct greater_equal {
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool operator()(T const &lhs, T const &rhs) const {
|
||||
return (lhs >= rhs);
|
||||
}
|
||||
};
|
||||
|
||||
/// Greater
|
||||
template <typename T>
|
||||
struct greater {
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool operator()(T const &lhs, T const &rhs) const {
|
||||
return (lhs > rhs);
|
||||
}
|
||||
};
|
||||
|
||||
/// Less equal
|
||||
template <typename T>
|
||||
struct less_equal {
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool operator()(T const &lhs, T const &rhs) const {
|
||||
return (lhs <= rhs);
|
||||
}
|
||||
};
|
||||
|
||||
/// Less
|
||||
template <typename T>
|
||||
struct less {
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool operator()(T const &lhs, T const &rhs) const {
|
||||
return (lhs < rhs);
|
||||
}
|
||||
};
|
||||
|
||||
/// Fused multiply-add
|
||||
template <typename A, typename B = A, typename C = A>
|
||||
struct multiply_add {
|
||||
@@ -189,6 +225,40 @@ struct xor_add {
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
struct conjugate {
|
||||
CUTLASS_HOST_DEVICE
|
||||
T operator()(T const &a) const {
|
||||
return a;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <typename T>
|
||||
struct conjugate<complex<T>> {
|
||||
CUTLASS_HOST_DEVICE
|
||||
complex<T> operator()(complex<T> const &a) const {
|
||||
return conj(a);
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, int N>
|
||||
struct conjugate<Array<T, N> > {
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator()(Array<T, N> const &a) const {
|
||||
|
||||
conjugate<T> conj_op;
|
||||
|
||||
Array<T, N> ca;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < N; ++i) {
|
||||
ca[i] = conj_op(a[i]);
|
||||
}
|
||||
return ca;
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// Partial specialization for complex<T> to target four scalar fused multiply-adds.
|
||||
@@ -1499,6 +1569,86 @@ struct multiply_add<Array<bfloat16_t, N>, Array<bfloat16_t, N>, Array<bfloat16_t
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator+(Array<T, N> const &lhs, Array<T, N> const &rhs) {
|
||||
plus<Array<T, N>> op;
|
||||
return op(lhs, rhs);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator-(Array<T, N> const &lhs, Array<T, N> const &rhs) {
|
||||
minus<Array<T, N>> op;
|
||||
return op(lhs, rhs);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator-(Array<T, N> const &lhs) {
|
||||
negate<Array<T, N>> op;
|
||||
return op(lhs);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator*(Array<T, N> const &lhs, Array<T, N> const &rhs) {
|
||||
multiplies<Array<T, N>> op;
|
||||
return op(lhs, rhs);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator*(T lhs, Array<T, N> const &rhs) {
|
||||
multiplies<Array<T, N>> op;
|
||||
return op(lhs, rhs);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator*(Array<T, N> const &lhs, T rhs) {
|
||||
multiplies<Array<T, N>> op;
|
||||
return op(lhs, rhs);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> operator/(Array<T, N> const &lhs, Array<T, N> const &rhs) {
|
||||
divides<Array<T, N>> op;
|
||||
return op(lhs, rhs);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> fma(Array<T, N> const &a, Array<T, N> const &b, Array<T, N> const &c) {
|
||||
multiply_add<Array<T, N>> op;
|
||||
return op(a, b, c);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> fma(T a, Array<T, N> const &b, Array<T, N> const &c) {
|
||||
multiply_add<Array<T, N>> op;
|
||||
return op(a, b, c);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> fma(Array<T, N> const &a, T b, Array<T, N> const &c) {
|
||||
multiply_add<Array<T, N>> op;
|
||||
return op(a, b, c);
|
||||
}
|
||||
|
||||
template <typename T, int N>
|
||||
CUTLASS_HOST_DEVICE
|
||||
Array<T, N> fma(Array<T, N> const &a, Array<T, N> const &b, T c) {
|
||||
multiply_add<Array<T, N>> op;
|
||||
return op(a, b, c);
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
Reference in New Issue
Block a user