From 1a2f610d654a64b25b864b0832a1b0f01f16f8e2 Mon Sep 17 00:00:00 2001 From: willyborn Date: Fri, 27 Sep 2024 12:12:48 +0200 Subject: [PATCH] Enforced argument checking for direct backend calls on join --- src/backend/cpu/join.cpp | 68 ++++++++++++++++++++--------- src/backend/cuda/join.cpp | 66 ++++++++++++++++++++++++---- src/backend/oneapi/join.cpp | 76 ++++++++++++++++++++++++++------- src/backend/opencl/join.cpp | 85 +++++++++++++++++++++++++++++-------- 4 files changed, 236 insertions(+), 59 deletions(-) diff --git a/src/backend/cpu/join.cpp b/src/backend/cpu/join.cpp index e9fed65df1..15c6ae46f0 100644 --- a/src/backend/cpu/join.cpp +++ b/src/backend/cpu/join.cpp @@ -15,47 +15,77 @@ #include #include +#include +#include +using af::dim4; using arrayfire::common::half; namespace arrayfire { namespace cpu { template -Array join(const int dim, const Array &first, const Array &second) { - // All dimensions except join dimension must be equal +Array join(const int jdim, const Array &first, const Array &second) { // Compute output dims - af::dim4 odims; - af::dim4 fdims = first.dims(); - af::dim4 sdims = second.dims(); - - for (int i = 0; i < 4; i++) { - if (i == dim) { - odims[i] = fdims[i] + sdims[i]; - } else { - odims[i] = fdims[i]; - } - } - + const dim4 &fdims = first.dims(); + const dim4 &sdims = second.dims(); + // All dimensions except join dimension must be equal + assert((jdim == 0 ? true : fdims.dims[0] == sdims.dims[0]) && + (jdim == 1 ? true : fdims.dims[1] == sdims.dims[1]) && + (jdim == 2 ? true : fdims.dims[2] == sdims.dims[2]) && + (jdim == 3 ? true : fdims.dims[3] == sdims.dims[3])); + + // compute output dms + dim4 odims(fdims); + odims.dims[jdim] += sdims.dims[jdim]; Array out = createEmptyArray(odims); std::vector> v{first, second}; - getQueue().enqueue(kernel::join, dim, out, v, 2); + getQueue().enqueue(kernel::join, jdim, out, v, 2); return out; } template -void join(Array &out, const int dim, const std::vector> &inputs) { +void join(Array &out, const int jdim, const std::vector> &inputs) { const dim_t n_arrays = inputs.size(); - - std::vector *> input_ptrs(inputs.size()); + if (n_arrays == 0) return; + + // avoid buffer overflow + const dim4 &odims{out.dims()}; + const dim4 &fdims{inputs[0].dims()}; + // All dimensions of inputs needs to be equal except for the join + // dimension + assert(std::all_of(inputs.begin(), inputs.end(), + [jdim, &fdims](const Array &in) { + bool eq{true}; + for (int i = 0; i < 4; ++i) { + if (i != jdim) { + eq &= fdims.dims[i] == in.dims().dims[i]; + }; + }; + return eq; + })); + // All dimensions of out needs to cover all input dimensions + assert( + (odims.dims[0] >= fdims.dims[0]) && (odims.dims[1] >= fdims.dims[1]) && + (odims.dims[2] >= fdims.dims[2]) && (odims.dims[3] >= fdims.dims[3])); + // The join dimension of out needs to be larger than the + // sum of all input join dimensions + assert(odims.dims[jdim] >= + std::accumulate(inputs.begin(), inputs.end(), 0, + [jdim](dim_t dim, const Array &in) { + return dim += in.dims()[jdim]; + })); + assert(out.strides().dims[0] == 1); + + std::vector *> input_ptrs(n_arrays); std::transform( begin(inputs), end(inputs), begin(input_ptrs), [](const Array &input) { return const_cast *>(&input); }); evalMultiple(input_ptrs); std::vector> inputParams(inputs.begin(), inputs.end()); - getQueue().enqueue(kernel::join, dim, out, inputParams, n_arrays); + getQueue().enqueue(kernel::join, jdim, out, inputParams, n_arrays); } #define INSTANTIATE(T) \ diff --git a/src/backend/cuda/join.cpp b/src/backend/cuda/join.cpp index 3eed6f7fb5..89cd2fdaf9 100644 --- a/src/backend/cuda/join.cpp +++ b/src/backend/cuda/join.cpp @@ -14,7 +14,9 @@ #include #include +#include #include +#include #include #include @@ -29,9 +31,14 @@ namespace cuda { template Array join(const int jdim, const Array &first, const Array &second) { - // All dimensions except join dimension must be equal const dim4 &fdims{first.dims()}; const dim4 &sdims{second.dims()}; + // All dimensions except join dimension must be equal + assert((jdim == 0 ? true : fdims.dims[0] == sdims.dims[0]) && + (jdim == 1 ? true : fdims.dims[1] == sdims.dims[1]) && + (jdim == 2 ? true : fdims.dims[2] == sdims.dims[2]) && + (jdim == 3 ? true : fdims.dims[3] == sdims.dims[3])); + // Compute output dims dim4 odims(fdims); odims.dims[jdim] += sdims.dims[jdim]; @@ -117,6 +124,42 @@ Array join(const int jdim, const Array &first, const Array &second) { template void join(Array &out, const int jdim, const vector> &inputs) { + const dim_t n_arrays = inputs.size(); + if (n_arrays == 0) return; + + // avoid buffer overflow + const dim4 &odims{out.dims()}; + const dim4 &fdims{inputs[0].dims()}; + // All dimensions of inputs needs to be equal except for the join + // dimension + assert(std::all_of(inputs.begin(), inputs.end(), + [jdim, &fdims](const Array &in) { + bool eq{true}; + for (int i = 0; i < 4; ++i) { + if (i != jdim) { + eq &= fdims.dims[i] == in.dims().dims[i]; + }; + }; + return eq; + })); + // All dimensions of out needs to cover all input dimensions + assert( + (odims.dims[0] >= fdims.dims[0]) && (odims.dims[1] >= fdims.dims[1]) && + (odims.dims[2] >= fdims.dims[2]) && (odims.dims[3] >= fdims.dims[3])); + // The join dimension of out needs to be larger than the + // sum of all input join dimensions + assert(odims.dims[jdim] >= + std::accumulate(inputs.begin(), inputs.end(), 0, + [jdim](dim_t dim, const Array &in) { + return dim += in.dims().dims[jdim]; + })); + assert(out.strides().dims[0] == 1); + + // out is an external defined array: + // - with the only restriction that the dims have to be larger than the + // joined inputs. + // - no restrictions on the strides. + // The part of out, that is not overwritten by the join remains as is!! class eval { public: vector> outputs; @@ -126,6 +169,7 @@ void join(Array &out, const int jdim, const vector> &inputs) { }; std::map evals; const cudaStream_t activeStream{getActiveStream()}; + const dim4 &ostrides{out.strides()}; const size_t L2CacheSize{getL2CacheSize(getActiveDeviceId())}; // topspeed is achieved when byte size(in+out) ~= L2CacheSize @@ -148,18 +192,19 @@ void join(Array &out, const int jdim, const vector> &inputs) { // has to be called multiple times // Group all arrays according to size - dim_t outOffset{0}; + dim_t odim{0}, outOffset{0}; for (const Array &iArray : inputs) { - const dim_t *idims{iArray.dims().dims}; - eval &e{evals[idims[jdim]]}; - e.outputs.emplace_back(out.get() + outOffset, idims, - out.strides().dims); + const dim4 &idims{iArray.dims()}; + eval &e{evals[idims.dims[jdim]]}; + e.outputs.emplace_back(out.get() + outOffset, idims.dims, + ostrides.dims); // Extend life of the returned node by saving the corresponding // shared_ptr e.nodePtrs.emplace_back(iArray.getNode()); e.nodes.push_back(e.nodePtrs.back().get()); e.ins.push_back(&iArray); - outOffset += idims[jdim] * out.strides().dims[jdim]; + odim += idims.dims[jdim]; + outOffset = odim * ostrides.dims[jdim]; } for (auto &eval : evals) { @@ -173,7 +218,12 @@ void join(Array &out, const int jdim, const vector> &inputs) { auto outputIt{begin(s.outputs)}; for (const Array *in : s.ins) { if (in->isReady()) { - if (1LL + jdim >= in->ndims() && in->isLinear()) { + const dim4 &istrides{in->strides()}; + bool lin = in->isLinear() & (ostrides.dims[0] == 1); + for (int i{1}; i < in->ndims(); ++i) { + lin &= (ostrides.dims[i] == istrides.dims[i]); + } + if (lin) { CUDA_CHECK(cudaMemcpyAsync(outputIt->ptr, in->get(), in->elements() * sizeof(T), cudaMemcpyHostToDevice, diff --git a/src/backend/oneapi/join.cpp b/src/backend/oneapi/join.cpp index 37c7c14fc9..aaaf196aa9 100644 --- a/src/backend/oneapi/join.cpp +++ b/src/backend/oneapi/join.cpp @@ -42,6 +42,11 @@ Array join(const int jdim, const Array &first, const Array &second) { // All dimensions except join dimension must be equal const dim4 &fdims{first.dims()}; const dim4 &sdims{second.dims()}; + // All dimensions except join dimension must be equal + assert((jdim == 0 ? true : fdims.dims[0] == sdims.dims[0]) && + (jdim == 1 ? true : fdims.dims[1] == sdims.dims[1]) && + (jdim == 2 ? true : fdims.dims[2] == sdims.dims[2]) && + (jdim == 3 ? true : fdims.dims[3] == sdims.dims[3])); // Compute output dims dim4 odims(fdims); @@ -69,16 +74,18 @@ Array join(const int jdim, const Array &first, const Array &second) { // Both arrays have same size & everything fits into the cache, // so thread in 1 JIT kernel, iso individual copies which is // always slower - const dim_t *outStrides{out.strides().dims}; + const dim4 &outStrides{out.strides()}; vector> outputs{ {out.get(), {{fdims.dims[0], fdims.dims[1], fdims.dims[2], fdims.dims[3]}, - {outStrides[0], outStrides[1], outStrides[2], outStrides[3]}, + {outStrides.dims[0], outStrides.dims[1], outStrides.dims[2], + outStrides.dims[3]}, 0}}, {out.get(), {{sdims.dims[0], sdims.dims[1], sdims.dims[2], sdims.dims[3]}, - {outStrides[0], outStrides[1], outStrides[2], outStrides[3]}, - fdims.dims[jdim] * outStrides[jdim]}}}; + {outStrides.dims[0], outStrides.dims[1], outStrides.dims[2], + outStrides.dims[3]}, + fdims.dims[jdim] * outStrides.dims[jdim]}}}; // Extend the life of the returned node, bij saving the // corresponding shared_ptr const Node_ptr fNode{first.getNode()}; @@ -113,11 +120,12 @@ Array join(const int jdim, const Array &first, const Array &second) { } } else { // Write the result directly in the out array - const dim_t *outStrides{out.strides().dims}; + const dim4 &outStrides{out.strides()}; Param output{ out.get(), {{fdims.dims[0], fdims.dims[1], fdims.dims[2], fdims.dims[3]}, - {outStrides[0], outStrides[1], outStrides[2], outStrides[3]}, + {outStrides.dims[0], outStrides.dims[1], outStrides.dims[2], + outStrides.dims[3]}, 0}}; evalNodes(output, first.getNode().get()); } @@ -146,12 +154,13 @@ Array join(const int jdim, const Array &first, const Array &second) { } } else { // Write the result directly in the out array - const dim_t *outStrides{out.strides().dims}; + const dim4 &outStrides{out.strides()}; Param output{ out.get(), {{sdims.dims[0], sdims.dims[1], sdims.dims[2], sdims.dims[3]}, - {outStrides[0], outStrides[1], outStrides[2], outStrides[3]}, - fdims.dims[jdim] * outStrides[jdim]}}; + {outStrides.dims[0], outStrides.dims[1], outStrides.dims[2], + outStrides.dims[3]}, + fdims.dims[jdim] * outStrides.dims[jdim]}}; evalNodes(output, second.getNode().get()); } return out; @@ -159,6 +168,37 @@ Array join(const int jdim, const Array &first, const Array &second) { template void join(Array &out, const int jdim, const vector> &inputs) { + const dim_t n_arrays = inputs.size(); + if (n_arrays == 0) return; + + // avoid buffer overflow + const dim4 &odims{out.dims()}; + const dim4 &fdims{inputs[0].dims()}; + // All dimensions of inputs needs to be equal except for the join + // dimension + assert(std::all_of(inputs.begin(), inputs.end(), + [jdim, &fdims](const Array &in) { + bool eq{true}; + for (int i = 0; i < 4; ++i) { + if (i != jdim) { + eq &= fdims.dims[i] == in.dims().dims[i]; + }; + }; + return eq; + })); + // All dimensions of out needs to cover all input dimensions + assert( + (odims.dims[0] >= fdims.dims[0]) && (odims.dims[1] >= fdims.dims[1]) && + (odims.dims[2] >= fdims.dims[2]) && (odims.dims[3] >= fdims.dims[3])); + // The join dimension of out needs to be larger than the + // sum of all input join dimensions + assert(odims.dims[jdim] >= + std::accumulate(inputs.begin(), inputs.end(), 0, + [jdim](dim_t dim, const Array &in) { + return dim += in.dims().dims[jdim]; + })); + assert(out.strides().dims[0] == 1); + class eval { public: vector> outputs; @@ -167,7 +207,7 @@ void join(Array &out, const int jdim, const vector> &inputs) { vector *> ins; }; std::map evals; - const dim_t *ostrides{out.strides().dims}; + const dim4 &ostrides{out.strides()}; const size_t L2CacheSize{getL2CacheSize(oneapi::getDevice())}; // topspeed is achieved when byte size(in+out) ~= L2CacheSize @@ -188,12 +228,13 @@ void join(Array &out, const int jdim, const vector> &inputs) { // Group all arrays according to size dim_t outOffset{0}; for (const Array &iArray : inputs) { - const dim_t *idims{iArray.dims().dims}; + const dim4 &idims{iArray.dims()}; eval &e{evals[idims[jdim]]}; const Param output{ out.get(), - {{idims[0], idims[1], idims[2], idims[3]}, - {ostrides[0], ostrides[1], ostrides[2], ostrides[3]}, + {{idims.dims[0], idims.dims[1], idims.dims[2], idims.dims[3]}, + {ostrides.dims[0], ostrides.dims[1], ostrides.dims[2], + ostrides.dims[3]}, outOffset}}; e.outputs.push_back(output); // Extend life of the returned node by saving the corresponding @@ -201,7 +242,7 @@ void join(Array &out, const int jdim, const vector> &inputs) { e.nodePtrs.emplace_back(iArray.getNode()); e.nodes.push_back(e.nodePtrs.back().get()); e.ins.push_back(&iArray); - outOffset += idims[jdim] * ostrides[jdim]; + outOffset += idims.dims[jdim] * ostrides.dims[jdim]; } for (auto &eval : evals) { @@ -215,7 +256,12 @@ void join(Array &out, const int jdim, const vector> &inputs) { auto outputIt{begin(s.outputs)}; for (const Array *in : s.ins) { if (in->isReady()) { - if (1LL + jdim >= in->ndims() && in->isLinear()) { + const dim_t *istrides{in->strides().dims}; + bool lin = in->isLinear() & (ostrides.dims[0] == 1); + for (int i{1}; i < in->ndims(); ++i) { + lin &= (ostrides.dims[i] == istrides[i]); + } + if (lin) { getQueue().submit([&](sycl::handler &h) { sycl::range sz(in->elements()); sycl::id src_offset(in->getOffset()); diff --git a/src/backend/opencl/join.cpp b/src/backend/opencl/join.cpp index 22875d0e61..b65f639465 100644 --- a/src/backend/opencl/join.cpp +++ b/src/backend/opencl/join.cpp @@ -14,7 +14,9 @@ #include #include +#include #include +#include #include #include @@ -28,9 +30,14 @@ namespace arrayfire { namespace opencl { template Array join(const int jdim, const Array &first, const Array &second) { - // All dimensions except join dimension must be equal const dim4 &fdims{first.dims()}; const dim4 &sdims{second.dims()}; + // All dimensions except join dimension must be equal + assert((jdim == 0 ? true : fdims.dims[0] == sdims.dims[0]) && + (jdim == 1 ? true : fdims.dims[1] == sdims.dims[1]) && + (jdim == 2 ? true : fdims.dims[2] == sdims.dims[2]) && + (jdim == 3 ? true : fdims.dims[3] == sdims.dims[3])); + // Compute output dims dim4 odims(fdims); odims.dims[jdim] += sdims.dims[jdim]; @@ -57,16 +64,18 @@ Array join(const int jdim, const Array &first, const Array &second) { // Both arrays have same size & everything fits into the cache, // so thread in 1 JIT kernel, iso individual copies which is // always slower - const dim_t *outStrides{out.strides().dims}; + const dim4 &outStrides{out.strides()}; vector outputs{ {out.get(), {{fdims.dims[0], fdims.dims[1], fdims.dims[2], fdims.dims[3]}, - {outStrides[0], outStrides[1], outStrides[2], outStrides[3]}, + {outStrides.dims[0], outStrides.dims[1], outStrides.dims[2], + outStrides.dims[3]}, 0}}, {out.get(), {{sdims.dims[0], sdims.dims[1], sdims.dims[2], sdims.dims[3]}, - {outStrides[0], outStrides[1], outStrides[2], outStrides[3]}, - fdims.dims[jdim] * outStrides[jdim]}}}; + {outStrides.dims[0], outStrides.dims[1], outStrides.dims[2], + outStrides.dims[3]}, + fdims.dims[jdim] * outStrides.dims[jdim]}}}; // Extend the life of the returned node, bij saving the // corresponding shared_ptr const Node_ptr fNode{first.getNode()}; @@ -92,11 +101,12 @@ Array join(const int jdim, const Array &first, const Array &second) { } } else { // Write the result directly in the out array - const dim_t *outStrides{out.strides().dims}; + const dim4 &outStrides{out.strides()}; Param output{ out.get(), {{fdims.dims[0], fdims.dims[1], fdims.dims[2], fdims.dims[3]}, - {outStrides[0], outStrides[1], outStrides[2], outStrides[3]}, + {outStrides.dims[0], outStrides.dims[1], outStrides.dims[2], + outStrides.dims[3]}, 0}}; evalNodes(output, first.getNode().get()); } @@ -116,12 +126,13 @@ Array join(const int jdim, const Array &first, const Array &second) { } } else { // Write the result directly in the out array - const dim_t *outStrides{out.strides().dims}; + const dim4 &outStrides{out.strides()}; Param output{ out.get(), {{sdims.dims[0], sdims.dims[1], sdims.dims[2], sdims.dims[3]}, - {outStrides[0], outStrides[1], outStrides[2], outStrides[3]}, - fdims.dims[jdim] * outStrides[jdim]}}; + {outStrides.dims[0], outStrides.dims[1], outStrides.dims[2], + outStrides.dims[3]}, + fdims.dims[jdim] * outStrides.dims[jdim]}}; evalNodes(output, second.getNode().get()); } @@ -130,6 +141,37 @@ Array join(const int jdim, const Array &first, const Array &second) { template void join(Array &out, const int jdim, const vector> &inputs) { + const dim_t n_arrays = inputs.size(); + if (n_arrays == 0) return; + + // avoid buffer overflow + const dim4 &odims{out.dims()}; + const dim4 &fdims{inputs[0].dims()}; + // All dimensions of inputs needs to be equal except for the join + // dimension + assert(std::all_of(inputs.begin(), inputs.end(), + [jdim, &fdims](const Array &in) { + bool eq{true}; + for (int i = 0; i < 4; ++i) { + if (i != jdim) { + eq &= fdims.dims[i] == in.dims().dims[i]; + }; + }; + return eq; + })); + // All dimensions of out needs to cover all input dimensions + assert( + (odims.dims[0] >= fdims.dims[0]) && (odims.dims[1] >= fdims.dims[1]) && + (odims.dims[2] >= fdims.dims[2]) && (odims.dims[3] >= fdims.dims[3])); + // The join dimension of out needs to be larger than the + // sum of all input join dimensions + assert(odims.dims[jdim] >= + std::accumulate(inputs.begin(), inputs.end(), 0, + [jdim](dim_t dim, const Array &in) { + return dim += in.dims().dims[jdim]; + })); + assert(out.strides().dims[0] == 1); + class eval { public: vector outputs; @@ -138,7 +180,7 @@ void join(Array &out, const int jdim, const vector> &inputs) { vector *> ins; }; std::map evals; - const dim_t *ostrides{out.strides().dims}; + const dim4 &ostrides{out.strides()}; const size_t L2CacheSize{getL2CacheSize(opencl::getDevice())}; // topspeed is achieved when byte size(in+out) ~= L2CacheSize @@ -157,14 +199,17 @@ void join(Array &out, const int jdim, const vector> &inputs) { // will be called twice // Group all arrays according to size - dim_t outOffset{0}; + dim_t odim{0}, outOffset{0}; for (const Array &iArray : inputs) { - const dim_t *idims{iArray.dims().dims}; + const dim4 &idims{iArray.dims()}; + for (int i = 0; i < AF_MAX_DIMS; ++i) + ARG_ASSERT(1, odims.dims[i] >= idims.dims[i]); eval &e{evals[idims[jdim]]}; const Param output{ out.get(), - {{idims[0], idims[1], idims[2], idims[3]}, - {ostrides[0], ostrides[1], ostrides[2], ostrides[3]}, + {{idims.dims[0], idims.dims[1], idims.dims[2], idims.dims[3]}, + {ostrides.dims[0], ostrides.dims[1], ostrides.dims[2], + ostrides.dims[3]}, outOffset}}; e.outputs.push_back(output); // Extend life of the returned node by saving the corresponding @@ -172,7 +217,8 @@ void join(Array &out, const int jdim, const vector> &inputs) { e.nodePtrs.emplace_back(iArray.getNode()); e.nodes.push_back(e.nodePtrs.back().get()); e.ins.push_back(&iArray); - outOffset += idims[jdim] * ostrides[jdim]; + odim += idims.dims[jdim]; + outOffset = odim * ostrides.dims[jdim]; } for (auto &eval : evals) { @@ -186,7 +232,12 @@ void join(Array &out, const int jdim, const vector> &inputs) { auto outputIt{begin(s.outputs)}; for (const Array *in : s.ins) { if (in->isReady()) { - if (1LL + jdim >= in->ndims() && in->isLinear()) { + const dim4 &istrides{in->strides()}; + bool lin = in->isLinear() & (ostrides.dims[0] == 1); + for (int i{1}; i < in->ndims(); ++i) { + lin &= (ostrides.dims[i] == istrides.dims[i]); + } + if (lin) { getQueue().enqueueCopyBuffer( *in->get(), *outputIt->data, in->getOffset() * sizeof(T),