CUTLASS 2.1 (#83)
CUTLASS 2.1 contributes: - BLAS-style host-side API added to CUTLASS Library - Planar Complex GEMM kernels targeting Volta and Turing Tensor Cores - Minor enhancements and bug fixes
This commit is contained in:
@@ -34,12 +34,12 @@
|
||||
|
||||
#include "cutlass/util/reference/device/tensor_compare.h"
|
||||
#include "cutlass/util/reference/device/tensor_fill.h"
|
||||
|
||||
#include "cutlass/util/reference/host/tensor_fill.h"
|
||||
|
||||
#include "cutlass/util/host_tensor.h"
|
||||
#include "cutlass/util/tensor_view_io.h"
|
||||
|
||||
#include "cutlass/library/util.h"
|
||||
|
||||
#include "device_allocation.h"
|
||||
|
||||
namespace cutlass {
|
||||
@@ -106,6 +106,18 @@ std::vector<int> DeviceAllocation::get_packed_layout(
|
||||
case library::LayoutTypeID::kRowMajorInterleavedK16:
|
||||
stride = get_packed_layout_stride<cutlass::layout::RowMajorInterleaved<16>>(extent);
|
||||
break;
|
||||
case library::LayoutTypeID::kColumnMajorInterleavedK32:
|
||||
stride = get_packed_layout_stride<cutlass::layout::ColumnMajorInterleaved<32>>(extent);
|
||||
break;
|
||||
case library::LayoutTypeID::kRowMajorInterleavedK32:
|
||||
stride = get_packed_layout_stride<cutlass::layout::RowMajorInterleaved<32>>(extent);
|
||||
break;
|
||||
case library::LayoutTypeID::kColumnMajorInterleavedK64:
|
||||
stride = get_packed_layout_stride<cutlass::layout::ColumnMajorInterleaved<64>>(extent);
|
||||
break;
|
||||
case library::LayoutTypeID::kRowMajorInterleavedK64:
|
||||
stride = get_packed_layout_stride<cutlass::layout::RowMajorInterleaved<64>>(extent);
|
||||
break;
|
||||
case library::LayoutTypeID::kTensorNCHW:
|
||||
stride = get_packed_layout_stride<cutlass::layout::TensorNCHW>(extent);
|
||||
break;
|
||||
@@ -200,6 +212,18 @@ size_t DeviceAllocation::construct_layout(
|
||||
case library::LayoutTypeID::kRowMajorInterleavedK16:
|
||||
return construct_layout_<cutlass::layout::RowMajorInterleaved<16>>(bytes, layout_id, extent, stride);
|
||||
|
||||
case library::LayoutTypeID::kColumnMajorInterleavedK32:
|
||||
return construct_layout_<cutlass::layout::ColumnMajorInterleaved<32>>(bytes, layout_id, extent, stride);
|
||||
|
||||
case library::LayoutTypeID::kRowMajorInterleavedK32:
|
||||
return construct_layout_<cutlass::layout::RowMajorInterleaved<32>>(bytes, layout_id, extent, stride);
|
||||
|
||||
case library::LayoutTypeID::kColumnMajorInterleavedK64:
|
||||
return construct_layout_<cutlass::layout::ColumnMajorInterleaved<64>>(bytes, layout_id, extent, stride);
|
||||
|
||||
case library::LayoutTypeID::kRowMajorInterleavedK64:
|
||||
return construct_layout_<cutlass::layout::RowMajorInterleaved<64>>(bytes, layout_id, extent, stride);
|
||||
|
||||
case library::LayoutTypeID::kTensorNCHW:
|
||||
return construct_layout_<cutlass::layout::TensorNHWC>(bytes, layout_id, extent, stride);
|
||||
|
||||
@@ -415,6 +439,14 @@ void DeviceAllocation::initialize_random_device(int seed, Distribution dist) {
|
||||
dist
|
||||
);
|
||||
break;
|
||||
case library::NumericTypeID::kCF64:
|
||||
cutlass::reference::device::BlockFillRandom<complex<double>>(
|
||||
reinterpret_cast<complex<double> *>(pointer_),
|
||||
capacity_,
|
||||
seed,
|
||||
dist
|
||||
);
|
||||
break;
|
||||
case library::NumericTypeID::kS8:
|
||||
cutlass::reference::device::BlockFillRandom<int8_t>(
|
||||
reinterpret_cast<int8_t *>(pointer_),
|
||||
@@ -508,6 +540,14 @@ void DeviceAllocation::initialize_random_host(int seed, Distribution dist) {
|
||||
dist
|
||||
);
|
||||
break;
|
||||
case library::NumericTypeID::kCF16:
|
||||
cutlass::reference::host::BlockFillRandom<cutlass::complex<cutlass::half_t>>(
|
||||
reinterpret_cast<cutlass::complex<cutlass::half_t> *>(host_data.data()),
|
||||
capacity_,
|
||||
seed,
|
||||
dist
|
||||
);
|
||||
break;
|
||||
case library::NumericTypeID::kF64:
|
||||
cutlass::reference::host::BlockFillRandom<double>(
|
||||
reinterpret_cast<double *>(host_data.data()),
|
||||
@@ -516,6 +556,14 @@ void DeviceAllocation::initialize_random_host(int seed, Distribution dist) {
|
||||
dist
|
||||
);
|
||||
break;
|
||||
case library::NumericTypeID::kCF64:
|
||||
cutlass::reference::host::BlockFillRandom<cutlass::complex<double>>(
|
||||
reinterpret_cast<cutlass::complex<double> *>(host_data.data()),
|
||||
capacity_,
|
||||
seed,
|
||||
dist
|
||||
);
|
||||
break;
|
||||
case library::NumericTypeID::kS8:
|
||||
cutlass::reference::host::BlockFillRandom<int8_t>(
|
||||
reinterpret_cast<int8_t *>(host_data.data()),
|
||||
@@ -607,13 +655,25 @@ bool DeviceAllocation::block_compare_equal(
|
||||
reinterpret_cast<float const *>(ptr_A),
|
||||
reinterpret_cast<float const *>(ptr_B),
|
||||
capacity);
|
||||
|
||||
|
||||
case library::NumericTypeID::kCF16:
|
||||
return reference::device::BlockCompareEqual<complex<half_t>>(
|
||||
reinterpret_cast<complex<half_t> const *>(ptr_A),
|
||||
reinterpret_cast<complex<half_t> const *>(ptr_B),
|
||||
capacity);
|
||||
|
||||
case library::NumericTypeID::kF64:
|
||||
return reference::device::BlockCompareEqual<double>(
|
||||
reinterpret_cast<double const *>(ptr_A),
|
||||
reinterpret_cast<double const *>(ptr_B),
|
||||
capacity);
|
||||
|
||||
case library::NumericTypeID::kCF64:
|
||||
return reference::device::BlockCompareEqual<complex<double>>(
|
||||
reinterpret_cast<complex<double> const *>(ptr_A),
|
||||
reinterpret_cast<complex<double> const *>(ptr_B),
|
||||
capacity);
|
||||
|
||||
case library::NumericTypeID::kS8:
|
||||
return reference::device::BlockCompareEqual<int8_t>(
|
||||
reinterpret_cast<int8_t const *>(ptr_A),
|
||||
|
||||
Reference in New Issue
Block a user