FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

handles empty arrays in join_many by syurkevi · Pull Request #3211 · arrayfire/arrayfire · GitHub

Repository navigation

Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension .cpp  (9) .h  (1) .hpp  (3) All 3 file types selected
Viewed files
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Unified
Split
Hide whitespace
Diff view
Unified
Split
Hide whitespace
10 changes: 10 additions & 0 deletions include/af/data.h
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -200,6 +200,8 @@ namespace af
\param[in] second is the second input array
\return the array that joins input arrays along the given dimension

\note empty arrays will be ignored

\ingroup manip_func_join
*/
AFAPI array join(const int dim, const array &first, const array &second);
Expand All @@ -213,6 +215,8 @@ namespace af
\param[in] third is the third input array
\return the array that joins input arrays along the given dimension

\note empty arrays will be ignored

\ingroup manip_func_join
*/
AFAPI array join(const int dim, const array &first, const array &second, const array &third);
Expand All @@ -227,6 +231,8 @@ namespace af
\param[in] fourth is the fourth input array
\return the array that joins input arrays along the given dimension

\note empty arrays will be ignored

\ingroup manip_func_join
*/
AFAPI array join(const int dim, const array &first, const array &second,
Expand Down Expand Up @@ -547,6 +553,8 @@ extern "C" {
\param[in] first is the first input array
\param[in] second is the second input array

\note empty arrays will be ignored

\ingroup manip_func_join
*/
AFAPI af_err af_join(af_array *out, const int dim, const af_array first, const af_array second);
Expand All @@ -561,6 +569,8 @@ extern "C" {
\param[in] n_arrays number of arrays to join
\param[in] inputs is an array of af_arrays containing handles to the arrays to be joined

\note empty arrays will be ignored

\ingroup manip_func_join
*/
AFAPI af_err af_join_many(af_array *out, const int dim, const unsigned n_arrays, const af_array *inputs);
Expand Down
56 changes: 49 additions & 7 deletions src/api/c/join.cpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -14,13 +14,15 @@
#include <handle.hpp>
#include <join.hpp>
#include <af/data.h>
#include <algorithm>
#include <vector>

using af::dim4;
using common::half;
using detail::Array;
using detail::cdouble;
using detail::cfloat;
using detail::createEmptyArray;
using detail::intl;
using detail::uchar;
using detail::uint;
Expand All @@ -43,8 +45,30 @@ static inline af_array join_many(const int dim, const unsigned n_arrays,

for (unsigned i = 0; i < n_arrays; i++) {
inputs_.push_back(getArray<T>(inputs[i]));
if (inputs_.back().isEmpty()) { inputs_.pop_back(); }
}
return getHandle(join<T>(dim, inputs_));

// All dimensions except join dimension must be equal
// calculate odims size
std::vector<af::dim4> idims(inputs_.size());
dim_t dim_size = 0;
for (unsigned i = 0; i < idims.size(); i++) {
idims[i] = inputs_[i].dims();
dim_size += idims[i][dim];
}

af::dim4 odims;
for (int i = 0; i < 4; i++) {
if (i == dim) {
odims[i] = dim_size;
} else {
odims[i] = idims[0][i];
}
}

Array<T> out = createEmptyArray<T>(odims);
join<T>(out, dim, inputs_);
return getHandle(out);
}

af_err af_join(af_array *out, const int dim, const af_array first,
Expand Down Expand Up @@ -117,24 +141,42 @@ af_err af_join_many(af_array *out, const int dim, const unsigned n_arrays,

ARG_ASSERT(1, dim >= 0 && dim < 4);

bool allEmpty = std::all_of(
info.begin(), info.end(),
[](const ArrayInfo &i) -> bool { return i.elements() <= 0; });
if (allEmpty) {
af_array ret = nullptr;
AF_CHECK(af_retain_array(&ret, inputs[0]));
std::swap(*out, ret);
return AF_SUCCESS;
}

auto first_valid_afinfo = std::find_if(
info.begin(), info.end(),
[](const ArrayInfo &i) -> bool { return i.elements() > 0; });

af_dtype assertType = first_valid_afinfo->getType();
for (unsigned i = 1; i < n_arrays; i++) {
ARG_ASSERT(3, info[0].getType() == info[i].getType());
DIM_ASSERT(3, info[i].elements() > 0);
if (info[i].elements() > 0) {
ARG_ASSERT(3, assertType == info[i].getType());
}
}

// All dimensions except join dimension must be equal
// Compute output dims
af::dim4 assertDims = first_valid_afinfo->dims();
for (int i = 0; i < 4; i++) {
if (i != dim) {
for (unsigned j = 1; j < n_arrays; j++) {
DIM_ASSERT(3, dims[0][i] == dims[j][i]);
for (unsigned j = 0; j < n_arrays; j++) {
if (info[j].elements() > 0) {
DIM_ASSERT(3, assertDims[i] == dims[j][i]);
}
}
}
}

af_array output;

switch (info[0].getType()) {
switch (assertType) {
case f32: output = join_many<float>(dim, n_arrays, inputs); break;
case c32: output = join_many<cfloat>(dim, n_arrays, inputs); break;
case f64: output = join_many<double>(dim, n_arrays, inputs); break;
Expand Down
6 changes: 5 additions & 1 deletion src/api/c/rgb_gray.cpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ using af::dim4;
using common::cast;
using detail::arithOp;
using detail::Array;
using detail::createEmptyArray;
using detail::createValueArray;
using detail::join;
using detail::scalar;
Expand Down Expand Up @@ -96,7 +97,10 @@ static af_array gray2rgb(const af_array& in, const float r, const float g,
AF_CHECK(af_release_array(mod_input));

// join channels
return getHandle(join<cType>(2, {expr3, expr1, expr2}));
dim4 odims(expr1.dims()[0], expr1.dims()[1], 3);
Array<cType> out = createEmptyArray<cType>(odims);
join<cType>(out, 2, {expr3, expr1, expr2});
return getHandle(out);
}

template<typename T, typename cType, bool isRGB2GRAY>
Expand Down
7 changes: 6 additions & 1 deletion src/api/c/surface.cpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ using af::dim4;
using common::modDims;
using detail::Array;
using detail::copy_surface;
using detail::createEmptyArray;
using detail::forgeManager;
using detail::reduce_all;
using detail::uchar;
Expand Down Expand Up @@ -72,7 +73,11 @@ fg_chart setup_surface(fg_window window, const af_array xVals,

// Now join along first dimension, skip reorder
std::vector<Array<T>> inputs{xIn, yIn, zIn};
Array<T> Z = join(0, inputs);

dim4 odims(3, rowDims[1]);
Array<T> out = createEmptyArray<T>(odims);
join(out, 0, inputs);
Array<T> Z = out;

ForgeManager& fgMngr = forgeManager();

Expand Down
10 changes: 8 additions & 2 deletions src/api/c/vector_field.cpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
using af::dim4;
using detail::Array;
using detail::copy_vector_field;
using detail::createEmptyArray;
using detail::forgeManager;
using detail::reduce;
using detail::transpose;
Expand All @@ -50,8 +51,13 @@ fg_chart setup_vector_field(fg_window window, const vector<af_array>& points,
}

// Join for set up vector
Array<T> pIn = detail::join(1, pnts);
Array<T> dIn = detail::join(1, dirs);
dim4 odims(3, points.size());
Array<T> out_pnts = createEmptyArray<T>(odims);
Array<T> out_dirs = createEmptyArray<T>(odims);
detail::join(out_pnts, 1, pnts);
detail::join(out_dirs, 1, dirs);
Array<T> pIn = out_pnts;
Array<T> dIn = out_dirs;

// do transpose if required
if (transpose_) {
Expand Down
11 changes: 9 additions & 2 deletions src/api/c/ycbcr_rgb.cpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
using af::dim4;
using detail::arithOp;
using detail::Array;
using detail::createEmptyArray;
using detail::createValueArray;
using detail::join;
using detail::scalar;
Expand Down Expand Up @@ -108,7 +109,10 @@ static af_array convert(const af_array& in, const af_ycc_std standard) {
INV_112 * (kb - 1) * kb * invKl);
Array<T> B = mix<T>(Y_, Cb_, INV_219, INV_112 * (1 - kb));
// join channels
return getHandle(join<T>(2, {R, G, B}));
dim4 odims(R.dims()[0], R.dims()[1], 3);
Array<T> rgbout = createEmptyArray<T>(odims);
join<T>(rgbout, 2, {R, G, B});
return getHandle(rgbout);
}
Array<T> Ey = mix<T>(X, Y, Z, kr, kl, kb);
Array<T> Ecr =
Expand All @@ -119,7 +123,10 @@ static af_array convert(const af_array& in, const af_ycc_std standard) {
Array<T> Cr = digitize<T>(Ecr, 224.0, 128.0);
Array<T> Cb = digitize<T>(Ecb, 224.0, 128.0);
// join channels
return getHandle(join<T>(2, {Y_, Cb, Cr}));
dim4 odims(Y_.dims()[0], Y_.dims()[1], 3);
Array<T> ycbcrout = createEmptyArray<T>(odims);
join<T>(ycbcrout, 2, {Y_, Cb, Cr});
return getHandle(ycbcrout);
}

template<bool isYCbCr2RGB>
Expand Down
29 changes: 4 additions & 25 deletions src/backend/cpu/join.cpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -44,38 +44,17 @@ Array<T> join(const int dim, const Array<T> &first, const Array<T> &second) {
}

template<typename T>
Array<T> join(const int dim, const std::vector<Array<T>> &inputs) {
// All dimensions except join dimension must be equal
// Compute output dims
af::dim4 odims;
void join(Array<T> &out, const int dim, const std::vector<Array<T>> &inputs) {
const dim_t n_arrays = inputs.size();
std::vector<af::dim4> idims(n_arrays);

dim_t dim_size = 0;
for (unsigned i = 0; i < idims.size(); i++) {
idims[i] = inputs[i].dims();
dim_size += idims[i][dim];
}

for (int i = 0; i < 4; i++) {
if (i == dim) {
odims[i] = dim_size;
} else {
odims[i] = idims[0][i];
}
}

std::vector<Array<T> *> input_ptrs(inputs.size());
std::transform(
begin(inputs), end(inputs), begin(input_ptrs),
[](const Array<T> &input) { return const_cast<Array<T> *>(&input); });
evalMultiple(input_ptrs);
std::vector<CParam<T>> inputParams(inputs.begin(), inputs.end());
Array<T> out = createEmptyArray<T>(odims);

getQueue().enqueue(kernel::join<T>, dim, out, inputParams, n_arrays);

return out;
}

#define INSTANTIATE(T) \
Expand All @@ -98,9 +77,9 @@ INSTANTIATE(half)

#undef INSTANTIATE

#define INSTANTIATE(T) \
template Array<T> join<T>(const int dim, \
const std::vector<Array<T>> &inputs);
#define INSTANTIATE(T) \
template void join<T>(Array<T> & out, const int dim, \
const std::vector<Array<T>> &inputs);

INSTANTIATE(float)
INSTANTIATE(double)
Expand Down
2 changes: 1 addition & 1 deletion src/backend/cpu/join.hpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -15,5 +15,5 @@ template<typename T>
Array<T> join(const int dim, const Array<T> &first, const Array<T> &second);

template<typename T>
Array<T> join(const int dim, const std::vector<Array<T>> &inputs);
void join(Array<T> &output, const int dim, const std::vector<Array<T>> &inputs);
} // namespace cpu
30 changes: 4 additions & 26 deletions src/backend/cuda/join.cpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -69,36 +69,14 @@ void join_wrapper(const int dim, Array<T> &out,
}

template<typename T>
Array<T> join(const int dim, const std::vector<Array<T>> &inputs) {
// All dimensions except join dimension must be equal
// Compute output dims
af::dim4 odims;
const dim_t n_arrays = inputs.size();
std::vector<af::dim4> idims(n_arrays);

dim_t dim_size = 0;
for (size_t i = 0; i < idims.size(); i++) {
idims[i] = inputs[i].dims();
dim_size += idims[i][dim];
}

for (int i = 0; i < 4; i++) {
if (i == dim) {
odims[i] = dim_size;
} else {
odims[i] = idims[0][i];
}
}

void join(Array<T> &out, const int dim, const std::vector<Array<T>> &inputs) {
std::vector<Array<T> *> input_ptrs(inputs.size());
std::transform(
begin(inputs), end(inputs), begin(input_ptrs),
[](const Array<T> &input) { return const_cast<Array<T> *>(&input); });
evalMultiple(input_ptrs);
Array<T> out = createEmptyArray<T>(odims);

join_wrapper<T>(dim, out, inputs);
return out;
}

#define INSTANTIATE(T) \
Expand All @@ -121,9 +99,9 @@ INSTANTIATE(half)

#undef INSTANTIATE

#define INSTANTIATE(T) \
template Array<T> join<T>(const int dim, \
const std::vector<Array<T>> &inputs);
#define INSTANTIATE(T) \
template void join<T>(Array<T> & out, const int dim, \
const std::vector<Array<T>> &inputs);

INSTANTIATE(float)
INSTANTIATE(double)
Expand Down
2 changes: 1 addition & 1 deletion src/backend/cuda/join.hpp
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters. Learn more about bidirectional Unicode characters
Original file line number Diff line number Diff line change
Expand Up @@ -14,5 +14,5 @@ template<typename T>
Array<T> join(const int dim, const Array<T> &first, const Array<T> &second);

template<typename T>
Array<T> join(const int dim, const std::vector<Array<T>> &inputs);
void join(Array<T> &out, const int dim, const std::vector<Array<T>> &inputs);
} // namespace cuda
Loading

Back | FazBrowse Home | New Git URL