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

Join does not always respect the order of provided parameters (#3511) by willyborn · Pull Request #3513 · 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  (5) .hpp  (1) All 2 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
6 changes: 6 additions & 0 deletions src/backend/common/jit/Node.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 @@ -42,6 +42,7 @@ int Node::getNodesMap(Node_map_t &node_map, vector<Node *> &full_nodes,
}

std::string getFuncName(const vector<Node *> &output_nodes,
const vector<int> &output_ids,
const vector<Node *> &full_nodes,
const vector<Node_ids> &full_ids, const bool is_linear,
const bool loop0, const bool loop1, const bool loop2,
Expand All @@ -59,6 +60,11 @@ std::string getFuncName(const vector<Node *> &output_nodes,
funcName += node->getNameStr();
}

for (const int id : output_ids) {
funcName += '-';
funcName += std::to_string(id);
}

for (int i = 0; i < static_cast<int>(full_nodes.size()); i++) {
full_nodes[i]->genKerName(funcName, full_ids[i]);
}
Expand Down
1 change: 1 addition & 0 deletions src/backend/common/jit/Node.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 @@ -326,6 +326,7 @@ struct Node_ids {
};

std::string getFuncName(const std::vector<Node *> &output_nodes,
const std::vector<int> &output_ids,
const std::vector<Node *> &full_nodes,
const std::vector<Node_ids> &full_ids,
const bool is_linear, const bool loop0,
Expand Down
27 changes: 15 additions & 12 deletions src/backend/cuda/jit.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 @@ -244,16 +244,18 @@ struct Param {
node->genOffsets(inOffsetsStream, ids_curr.id, is_linear);
// Generate the core function body, needs children ids as well
node->genFuncs(opsStream, ids_curr);
for (auto outIt{begin(output_ids)}, endIt{end(output_ids)};
(outIt = find(outIt, endIt, ids_curr.id)) != endIt; ++outIt) {
// Generate also output parameters
outParamStream << (oid == 0 ? "" : ",\n") << "Param<"
<< full_nodes[ids_curr.id]->getTypeStr()
<< "> out" << oid;
// Generate code to write the output (offset already in ptr)
opsStream << "out" << oid << ".ptr[idx] = val" << ids_curr.id
<< ";\n";
++oid;
for (size_t output_idx{0}; output_idx < output_ids.size();
++output_idx) {
if (output_ids[output_idx] == ids_curr.id) {
// Generate also output parameters
outParamStream << (oid == 0 ? "" : ",\n") << "Param<"
<< full_nodes[ids_curr.id]->getTypeStr()
<< "> out" << oid;
// Generate code to write the output (offset already in ptr)
opsStream << "out" << output_idx << ".ptr[idx] = val"
<< ids_curr.id << ";\n";
++oid;
}
}
}

Expand Down Expand Up @@ -322,8 +324,9 @@ static CUfunction getKernel(const vector<Node*>& output_nodes,
const bool is_linear, const bool loop0,
const bool loop1, const bool loop2,
const bool loop3) {
const string funcName{getFuncName(output_nodes, full_nodes, full_ids,
is_linear, loop0, loop1, loop2, loop3)};
const string funcName{getFuncName(output_nodes, output_ids, full_nodes,
full_ids, is_linear, loop0, loop1, loop2,
loop3)};
// A forward lookup in module cache helps avoid recompiling
// the JIT source generated from identical JIT-trees.
const auto entry{
Expand Down
4 changes: 2 additions & 2 deletions src/backend/oneapi/jit.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 @@ -478,8 +478,8 @@ void evalNodes(vector<Param<T>>& outputs, const vector<Node*>& output_nodes) {
full_nodes.clear();
for (Node_ptr& node : node_clones) { full_nodes.push_back(node.get()); }

const string funcName{getFuncName(output_nodes, full_nodes, full_ids,
is_linear, false, false, false,
const string funcName{getFuncName(output_nodes, output_ids, full_nodes,
full_ids, is_linear, false, false, false,
outputs[0].info.dims[2] > 1)};

getQueue().submit([&](sycl::handler& h) {
Expand Down
118 changes: 67 additions & 51 deletions src/backend/opencl/jit.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 @@ -188,62 +188,77 @@ __kernel void )JIT";
thread_local stringstream outOffsetStream;
thread_local stringstream inOffsetsStream;
thread_local stringstream opsStream;
thread_local stringstream kerStream;

int oid{0};
for (size_t i{0}; i < full_nodes.size(); i++) {
const auto& node{full_nodes[i]};
const auto& ids_curr{full_ids[i]};
// Generate input parameters, only needs current id
node->genParams(inParamStream, ids_curr.id, is_linear);
// Generate input offsets, only needs current id
node->genOffsets(inOffsetsStream, ids_curr.id, is_linear);
// Generate the core function body, needs children ids as well
node->genFuncs(opsStream, ids_curr);
for (auto outIt{begin(output_ids)}, endIt{end(output_ids)};
(outIt = find(outIt, endIt, ids_curr.id)) != endIt; ++outIt) {
// Generate also output parameters
outParamStream << "__global "
<< full_nodes[ids_curr.id]->getTypeStr() << " *out"
<< oid << ", int offset" << oid << ",\n";
// Apply output offset
outOffsetStream << "\nout" << oid << " += offset" << oid << ';';
// Generate code to write the output
opsStream << "out" << oid << "[idx] = val" << ids_curr.id << ";\n";
++oid;
string ret;
try {
int oid{0};
for (size_t i{0}; i < full_nodes.size(); i++) {
const auto& node{full_nodes[i]};
const auto& ids_curr{full_ids[i]};
// Generate input parameters, only needs current id
node->genParams(inParamStream, ids_curr.id, is_linear);
// Generate input offsets, only needs current id
node->genOffsets(inOffsetsStream, ids_curr.id, is_linear);
// Generate the core function body, needs children ids as well
node->genFuncs(opsStream, ids_curr);
for (size_t output_idx{0}; output_idx < output_ids.size();
++output_idx) {
if (output_ids[output_idx] == ids_curr.id) {
outParamStream
<< "__global " << full_nodes[ids_curr.id]->getTypeStr()
<< " *out" << oid << ", int offset" << oid << ",\n";
// Apply output offset
outOffsetStream << "\nout" << oid << " += offset" << oid
<< ';';
// Generate code to write the output
opsStream << "out" << output_idx << "[idx] = val"
<< ids_curr.id << ";\n";
++oid;
}
}
}
}

thread_local stringstream kerStream;
kerStream << kernelVoid << funcName << "(\n"
<< inParamStream.str() << outParamStream.str() << dimParams << ")"
<< blockStart;
if (is_linear) {
kerStream << linearInit << inOffsetsStream.str()
<< outOffsetStream.str() << '\n';
if (loop0) kerStream << linearLoop0Start;
kerStream << "\n\n" << opsStream.str();
if (loop0) kerStream << linearLoop0End;
kerStream << linearEnd;
} else {
if (loop0) {
kerStream << stridedLoop0Init << outOffsetStream.str() << '\n'
<< stridedLoop0Start;
kerStream << kernelVoid << funcName << "(\n"
<< inParamStream.str() << outParamStream.str() << dimParams
<< ")" << blockStart;
if (is_linear) {
kerStream << linearInit << inOffsetsStream.str()
<< outOffsetStream.str() << '\n';
if (loop0) kerStream << linearLoop0Start;
kerStream << "\n\n" << opsStream.str();
if (loop0) kerStream << linearLoop0End;
kerStream << linearEnd;
} else {
kerStream << stridedLoopNInit << outOffsetStream.str() << '\n';
if (loop3) kerStream << stridedLoop3Init;
if (loop1) kerStream << stridedLoop1Init << stridedLoop1Start;
if (loop3) kerStream << stridedLoop3Start;
if (loop0) {
kerStream << stridedLoop0Init << outOffsetStream.str() << '\n'
<< stridedLoop0Start;
} else {
kerStream << stridedLoopNInit << outOffsetStream.str() << '\n';
if (loop3) kerStream << stridedLoop3Init;
if (loop1) kerStream << stridedLoop1Init << stridedLoop1Start;
if (loop3) kerStream << stridedLoop3Start;
}
kerStream << "\n\n" << inOffsetsStream.str() << opsStream.str();
if (loop3) kerStream << stridedLoop3End;
if (loop1) kerStream << stridedLoop1End;
if (loop0) kerStream << stridedLoop0End;
kerStream << stridedEnd;
}
kerStream << "\n\n" << inOffsetsStream.str() << opsStream.str();
if (loop3) kerStream << stridedLoop3End;
if (loop1) kerStream << stridedLoop1End;
if (loop0) kerStream << stridedLoop0End;
kerStream << stridedEnd;
kerStream << blockEnd;
ret = kerStream.str();
} catch (...) {
// Prepare for next round
inParamStream.str("");
outParamStream.str("");
inOffsetsStream.str("");
outOffsetStream.str("");
opsStream.str("");
kerStream.str("");
throw;
}
kerStream << blockEnd;
const string ret{kerStream.str()};

// Prepare for next round, limit memory
// Prepare for next round
inParamStream.str("");
outParamStream.str("");
inOffsetsStream.str("");
Expand All @@ -259,8 +274,9 @@ cl::Kernel getKernel(const vector<Node*>& output_nodes,
const vector<Node*>& full_nodes,
const vector<Node_ids>& full_ids, const bool is_linear,
const bool loop0, const bool loop1, const bool loop3) {
const string funcName{getFuncName(output_nodes, full_nodes, full_ids,
is_linear, loop0, loop1, false, loop3)};
const string funcName{getFuncName(output_nodes, output_ids, full_nodes,
full_ids, is_linear, loop0, loop1, false,
loop3)};
// A forward lookup in module cache helps avoid recompiling the JIT
// source generated from identical JIT-trees.
const auto entry{
Expand Down
46 changes: 46 additions & 0 deletions test/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 @@ -15,6 +15,7 @@
#include <af/index.h>
#include <af/traits.hpp>

#include <array>
#include <complex>
#include <iostream>
#include <numeric>
Expand Down Expand Up @@ -266,3 +267,48 @@ TEST(Join, ManyEmpty) {
ASSERT_ARRAYS_EQ(gold, eace);
ASSERT_ARRAYS_EQ(gold, acee);
}

TEST(Join, respect_parameters_order_ISSUE3511) {
const float column_host1[] = {1., 2., 3.};
const float column_host2[] = {4., 5., 6.};
const af::array buf1(3, 1, column_host1);
const af::array buf2(3, 1, column_host2);

// We need to avoid that JIT arrays are evaluated during whatever call,
// so we will have to work with copies for single use
const af::array jit1{buf1 + 1.0};
const af::array jit2{buf2 + 2.0};
const std::array<af::array, 8> cases{jit1, -jit1, jit1 + 1.0, jit2,
-jit2, jit1 + jit2, buf1, buf2};
const std::array<char*, 8> cases_name{"JIT1", "-JIT1", "JIT1+1.0",
"JIT2", "-JIT2", "JIT1+JIT2",
"BUF1", "BUF2"};
assert(cases.size() == cases_name.size());
for (size_t cl0{0}; cl0 < cases.size(); ++cl0) {
for (size_t cl1{0}; cl1 < cases.size(); ++cl1) {
printf("Testing: af::join(1,%s,%s)\n", cases_name[cl0],
cases_name[cl1]);
const array col0{cases[cl0]};
const array col1{cases[cl1]};
const array result{af::join(1, col0, col1)};
ASSERT_ARRAYS_EQ(result(af::span, 0), col0);
ASSERT_ARRAYS_EQ(result(af::span, 1), col1);
}
}
// Join of 3 arrays
for (size_t cl0{0}; cl0 < cases.size(); ++cl0) {
for (size_t cl1{0}; cl1 < cases.size(); ++cl1) {
for (size_t cl2{0}; cl2 < cases.size(); ++cl2) {
printf("Testing: af::join(1,%s,%s,%s)\n", cases_name[cl0],
cases_name[cl1], cases_name[cl2]);
const array col0{cases[cl0]};
const array col1{cases[cl1]};
const array col2{cases[cl2]};
const array result{af::join(1, col0, col1, col2)};
ASSERT_ARRAYS_EQ(result(af::span, 0), col0);
ASSERT_ARRAYS_EQ(result(af::span, 1), col1);
ASSERT_ARRAYS_EQ(result(af::span, 2), col2);
}
}
}
}

Back | FazBrowse Home | New Git URL