summaryrefslogtreecommitdiff
path: root/src/routines
diff options
context:
space:
mode:
authorCedric Nugteren <web@cedricnugteren.nl>2017-08-12 20:50:00 +0200
committerCedric Nugteren <web@cedricnugteren.nl>2017-08-12 20:50:00 +0200
commit777681dcbdf18493320dd7b94fccd5c6faee9455 (patch)
treeb8597f5d79f8ef33bffbf33f3de2548cc51d4c5c /src/routines
parent97bcf77d4bc9b31e32a8785787e0497ac5440e44 (diff)
parentd67fd6604b4a6584c4f9e856057fcc8076ce377d (diff)
Merge branch 'master' into im_to_col
Diffstat (limited to 'src/routines')
-rw-r--r--src/routines/common.cpp75
-rw-r--r--src/routines/common.hpp25
-rw-r--r--src/routines/level3/xgemm.cpp6
-rw-r--r--src/routines/levelx/xgemmbatched.cpp4
4 files changed, 83 insertions, 27 deletions
diff --git a/src/routines/common.cpp b/src/routines/common.cpp
index c995dc12..5b178e53 100644
--- a/src/routines/common.cpp
+++ b/src/routines/common.cpp
@@ -73,4 +73,79 @@ void RunKernel(Kernel &kernel, Queue &queue, const Device &device,
}
// =================================================================================================
+
+// Sets all elements of a matrix to a constant value
+template <typename T>
+void FillMatrix(Queue &queue, const Device &device,
+ const Program &program, const Databases &,
+ EventPointer event, const std::vector<Event> &waitForEvents,
+ const size_t m, const size_t n, const size_t ld, const size_t offset,
+ const Buffer<T> &dest,
+ const T constant_value) {
+ auto kernel = Kernel(program, "FillMatrix");
+ kernel.SetArgument(0, static_cast<int>(m));
+ kernel.SetArgument(1, static_cast<int>(n));
+ kernel.SetArgument(2, static_cast<int>(ld));
+ kernel.SetArgument(3, static_cast<int>(offset));
+ kernel.SetArgument(4, dest());
+ kernel.SetArgument(5, GetRealArg(constant_value));
+ auto local = std::vector<size_t>{8, 8};
+ auto global = std::vector<size_t>{Ceil(m, 8), Ceil(n, 8)};
+ RunKernel(kernel, queue, device, global, local, event, waitForEvents);
+}
+
+// Compiles the above function
+template void FillMatrix<half>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const size_t, const Buffer<half>&, const half);
+template void FillMatrix<float>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const size_t, const Buffer<float>&, const float);
+template void FillMatrix<double>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const size_t, const Buffer<double>&, const double);
+template void FillMatrix<float2>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const size_t, const Buffer<float2>&, const float2);
+template void FillMatrix<double2>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const size_t, const Buffer<double2>&, const double2);
+
+// Sets all elements of a vector to a constant value
+template <typename T>
+void FillVector(Queue &queue, const Device &device,
+ const Program &program, const Databases &,
+ EventPointer event, const std::vector<Event> &waitForEvents,
+ const size_t n, const size_t inc, const size_t offset,
+ const Buffer<T> &dest,
+ const T constant_value) {
+ auto kernel = Kernel(program, "FillVector");
+ kernel.SetArgument(0, static_cast<int>(n));
+ kernel.SetArgument(1, static_cast<int>(inc));
+ kernel.SetArgument(2, static_cast<int>(offset));
+ kernel.SetArgument(3, dest());
+ kernel.SetArgument(4, GetRealArg(constant_value));
+ auto local = std::vector<size_t>{64};
+ auto global = std::vector<size_t>{Ceil(n, 64)};
+ RunKernel(kernel, queue, device, global, local, event, waitForEvents);
+}
+
+// Compiles the above function
+template void FillVector<half>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const Buffer<half>&, const half);
+template void FillVector<float>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const Buffer<float>&, const float);
+template void FillVector<double>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const Buffer<double>&, const double);
+template void FillVector<float2>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const Buffer<float2>&, const float2);
+template void FillVector<double2>(Queue&, const Device&, const Program&, const Databases&,
+ EventPointer, const std::vector<Event>&, const size_t, const size_t,
+ const size_t, const Buffer<double2>&, const double2);
+
+// =================================================================================================
} // namespace clblast
diff --git a/src/routines/common.hpp b/src/routines/common.hpp
index 28a43da5..84ccd9d2 100644
--- a/src/routines/common.hpp
+++ b/src/routines/common.hpp
@@ -40,18 +40,7 @@ void FillMatrix(Queue &queue, const Device &device,
EventPointer event, const std::vector<Event> &waitForEvents,
const size_t m, const size_t n, const size_t ld, const size_t offset,
const Buffer<T> &dest,
- const T constant_value) {
- auto kernel = Kernel(program, "FillMatrix");
- kernel.SetArgument(0, static_cast<int>(m));
- kernel.SetArgument(1, static_cast<int>(n));
- kernel.SetArgument(2, static_cast<int>(ld));
- kernel.SetArgument(3, static_cast<int>(offset));
- kernel.SetArgument(4, dest());
- kernel.SetArgument(5, GetRealArg(constant_value));
- auto local = std::vector<size_t>{8, 8};
- auto global = std::vector<size_t>{Ceil(m, 8), Ceil(n, 8)};
- RunKernel(kernel, queue, device, global, local, event, waitForEvents);
-}
+ const T constant_value);
// Sets all elements of a vector to a constant value
template <typename T>
@@ -60,17 +49,7 @@ void FillVector(Queue &queue, const Device &device,
EventPointer event, const std::vector<Event> &waitForEvents,
const size_t n, const size_t inc, const size_t offset,
const Buffer<T> &dest,
- const T constant_value) {
- auto kernel = Kernel(program, "FillVector");
- kernel.SetArgument(0, static_cast<int>(n));
- kernel.SetArgument(1, static_cast<int>(inc));
- kernel.SetArgument(2, static_cast<int>(offset));
- kernel.SetArgument(3, dest());
- kernel.SetArgument(4, GetRealArg(constant_value));
- auto local = std::vector<size_t>{64};
- auto global = std::vector<size_t>{Ceil(n, 64)};
- RunKernel(kernel, queue, device, global, local, event, waitForEvents);
-}
+ const T constant_value);
// =================================================================================================
diff --git a/src/routines/level3/xgemm.cpp b/src/routines/level3/xgemm.cpp
index 30e5999c..136eec43 100644
--- a/src/routines/level3/xgemm.cpp
+++ b/src/routines/level3/xgemm.cpp
@@ -283,8 +283,10 @@ void Xgemm<T>::GemmDirect(const size_t m, const size_t n, const size_t k,
const auto m_ceiled = Ceil(m, db_["WGD"]);
const auto n_ceiled = Ceil(n, db_["WGD"]);
const auto global = std::vector<size_t>{
- (m_ceiled * db_["MDIMCD"]) / db_["WGD"],
- (n_ceiled * db_["NDIMCD"]) / db_["WGD"]
+ // CeilDiv(m * db_["MDIMCD"], db_["WGD"]),
+ // CeilDiv(n * db_["NDIMCD"], db_["WGD"])
+ (m_ceiled * db_["MDIMCD"]) / db_["WGD"],
+ (n_ceiled * db_["NDIMCD"]) / db_["WGD"]
};
const auto local = std::vector<size_t>{db_["MDIMCD"], db_["NDIMCD"]};
diff --git a/src/routines/levelx/xgemmbatched.cpp b/src/routines/levelx/xgemmbatched.cpp
index 0fea1922..ee8448d2 100644
--- a/src/routines/levelx/xgemmbatched.cpp
+++ b/src/routines/levelx/xgemmbatched.cpp
@@ -94,8 +94,8 @@ void XgemmBatched<T>::DoGemmBatched(const Layout layout, const Transpose a_trans
// Tests the matrices for validity
for (auto batch = size_t{0}; batch < batch_count; ++batch) {
- TestMatrixA(a_one, a_two, a_buffer, a_offsets[batch], a_ld);
- TestMatrixB(b_one, b_two, b_buffer, b_offsets[batch], b_ld);
+ TestMatrixA(a_one, a_two, a_buffer, a_offsets[batch], a_ld, false); // don't test for invalid LD
+ TestMatrixB(b_one, b_two, b_buffer, b_offsets[batch], b_ld, false); // don't test for invalid LD
TestMatrixC(c_one, c_two, c_buffer, c_offsets[batch], c_ld);
}