Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Next Next commit
Refactor common::Node and remove cpu::jit::Node
* Remove backend specific cpu::jit::Node and use common::Node as a
  base class for all backends.
* Remove std::string for types in the node class in favor of
  the enum af::dtype.
* Remove UnaryOp's default implementation
  • Loading branch information
umar456 committed May 20, 2020
commit df44610fd36c33307c478c109aa4dea8186de1af
7 changes: 3 additions & 4 deletions src/backend/common/jit/BinaryNode.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,10 +14,9 @@
namespace common {
class BinaryNode : public NaryNode {
public:
BinaryNode(const char *out_type_str, const char *name_str,
const char *op_str, common::Node_ptr lhs, common::Node_ptr rhs,
int op)
: NaryNode(out_type_str, name_str, op_str, 2, {{lhs, rhs}}, op,
BinaryNode(const af::dtype type, const char *op_str, common::Node_ptr lhs,
common::Node_ptr rhs, int op)
: NaryNode(type, op_str, 2, {{lhs, rhs}}, op,
std::max(lhs->getHeight(), rhs->getHeight()) + 1) {}
};
} // namespace common
12 changes: 6 additions & 6 deletions src/backend/common/jit/BufferNodeBase.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -28,8 +28,7 @@ class BufferNodeBase : public common::Node {
bool m_linear_buffer;

public:
BufferNodeBase(const char *type_str, const char *name_str)
: Node(type_str, name_str, 0, {}) {}
BufferNodeBase(af::dtype type) : Node(type, 0, {}) {}

bool isBuffer() const final { return true; }

Expand All @@ -54,14 +53,15 @@ class BufferNodeBase : public common::Node {

void genKerName(std::stringstream &kerStream,
const common::Node_ids &ids) const final {
kerStream << "_" << m_name_str;
kerStream << "_" << getNameStr();
kerStream << std::setw(3) << std::setfill('0') << std::dec << ids.id
<< std::dec;
}

void genParams(std::stringstream &kerStream, int id,
bool is_linear) const final {
detail::generateParamDeclaration(kerStream, id, is_linear, m_type_str);
detail::generateParamDeclaration(kerStream, id, is_linear,
getTypeStr());
}

int setArgs(int start_id, bool is_linear,
Expand All @@ -73,12 +73,12 @@ class BufferNodeBase : public common::Node {

void genOffsets(std::stringstream &kerStream, int id,
bool is_linear) const final {
detail::generateBufferOffsets(kerStream, id, is_linear, m_type_str);
detail::generateBufferOffsets(kerStream, id, is_linear, getTypeStr());
}

void genFuncs(std::stringstream &kerStream,
const common::Node_ids &ids) const final {
detail::generateBufferRead(kerStream, ids.id, m_type_str);
detail::generateBufferRead(kerStream, ids.id, getTypeStr());
}

void getInfo(unsigned &len, unsigned &buf_count,
Expand Down
8 changes: 4 additions & 4 deletions src/backend/common/jit/NaryNode.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -29,12 +29,11 @@ class NaryNode : public Node {
const std::string m_op_str;

public:
NaryNode(const char *out_type_str, const char *name_str, const char *op_str,
const int num_children,
NaryNode(const af::dtype type, const char *op_str, const int num_children,
const std::array<common::Node_ptr, Node::kMaxChildren> &&children,
const int op, const int height)
: common::Node(
out_type_str, name_str, height,
type, height,
std::forward<
const std::array<common::Node_ptr, Node::kMaxChildren>>(
children))
Expand All @@ -57,7 +56,8 @@ class NaryNode : public Node {

void genFuncs(std::stringstream &kerStream,
const common::Node_ids &ids) const final {
kerStream << m_type_str << " val" << ids.id << " = " << m_op_str << "(";
kerStream << getTypeStr() << " val" << ids.id << " = " << m_op_str
<< "(";
for (int i = 0; i < m_num_children; i++) {
if (i > 0) kerStream << ", ";
kerStream << "val" << ids.child_ids[i];
Expand Down
31 changes: 29 additions & 2 deletions src/backend/common/jit/Node.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9,16 +9,18 @@

#include <common/defines.hpp>
#include <common/jit/Node.hpp>
#include <common/util.hpp>

#include <sstream>
#include <string>
#include <vector>

using std::vector;

namespace common {

int Node::getNodesMap(Node_map_t &node_map, vector<const Node *> &full_nodes,
vector<Node_ids> &full_ids) const {
int Node::getNodesMap(Node_map_t &node_map, vector<Node *> &full_nodes,
vector<Node_ids> &full_ids) {
auto iter = node_map.find(this);
if (iter == node_map.end()) {
Node_ids ids{};
Expand All @@ -36,4 +38,29 @@ int Node::getNodesMap(Node_map_t &node_map, vector<const Node *> &full_nodes,
return iter->second;
}

std::string getFuncName(const vector<Node *> &output_nodes,
const vector<Node *> &full_nodes,
const vector<Node_ids> &full_ids, bool is_linear) {
std::stringstream funcName;
std::stringstream hashName;

if (is_linear) {
funcName << "L_"; // Kernel Linear
} else {
funcName << "G_"; // Kernel General
}

for (const auto &node : output_nodes) {
funcName << node->getNameStr() << "_";
}

for (int i = 0; i < static_cast<int>(full_nodes.size()); i++) {
full_nodes[i]->genKerName(funcName, full_ids[i]);
}

hashName << "KER";
hashName << deterministicHash(funcName.str());
return hashName.str();
}

} // namespace common
104 changes: 91 additions & 13 deletions src/backend/common/jit/Node.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,8 +8,11 @@
********************************************************/

#pragma once
#include <backend.hpp>
#include <common/defines.hpp>
#include <optypes.hpp>
#include <platform.hpp>
#include <types.hpp>
#include <af/defines.h>

#include <array>
Expand All @@ -31,31 +34,78 @@ class Node;
struct Node_ids;

using Node_ptr = std::shared_ptr<Node>;
using Node_map_t = std::unordered_map<const Node *, int>;
using Node_map_t = std::unordered_map<Node *, int>;
using Node_map_iter = Node_map_t::iterator;

static const char *getFullName(af::dtype type) {
switch (type) {
case f32: return detail::getFullName<float>();
case f64: return detail::getFullName<double>();
case c32: return detail::getFullName<detail::cfloat>();
case c64: return detail::getFullName<detail::cdouble>();
case u32: return detail::getFullName<unsigned>();
case s32: return detail::getFullName<int>();
case u64: return detail::getFullName<unsigned long long>();
case s64: return detail::getFullName<long long>();
case u16: return detail::getFullName<unsigned short>();
case s16: return detail::getFullName<short>();
case b8: return detail::getFullName<char>();
case u8: return detail::getFullName<unsigned char>();
case f16: return "half";
}
return "";
}

static const char *getShortName(af::dtype type) {
switch (type) {
case f32: return detail::shortname<float>();
case f64: return detail::shortname<double>();
case c32: return detail::shortname<detail::cfloat>();
case c64: return detail::shortname<detail::cdouble>();
case u32: return detail::shortname<unsigned>();
case s32: return detail::shortname<int>();
case u64: return detail::shortname<unsigned long long>();
case s64: return detail::shortname<long long>();
case u16: return detail::shortname<unsigned short>();
case s16: return detail::shortname<short>();
case b8: return detail::shortname<char>();
case u8: return detail::shortname<unsigned char>();
case f16: return "h";
}
return "";
}

class Node {
public:
static const int kMaxChildren = 3;

protected:
const std::array<Node_ptr, kMaxChildren> m_children;
const std::string m_type_str;
const std::string m_name_str;
const af::dtype m_type;
const int m_height;

template<typename T>
friend class NodeIterator;

public:
Node(const char *type_str, const char *name_str, const int height,
Node(const af::dtype type, const int height,
const std::array<Node_ptr, kMaxChildren> children)
: m_children(children)
, m_type_str(type_str)
, m_name_str(name_str)
, m_height(height) {}
: m_children(children), m_type(type), m_height(height) {}

/// Default copy constructor
Node(Node &node) = default;

int getNodesMap(Node_map_t &node_map, std::vector<const Node *> &full_nodes,
std::vector<Node_ids> &full_ids) const;
/// Default move constructor
Node(Node &&node) = default;

/// Default copy assignment operator
Node &operator=(const Node &node) = default;

/// Default move assignment operator
Node &operator=(Node &&node) = default;

int getNodesMap(Node_map_t &node_map, std::vector<Node *> &full_nodes,
std::vector<Node_ids> &full_ids);

/// Generates the string that will be used to hash the kernel
virtual void genKerName(std::stringstream &kerStream,
Expand All @@ -73,6 +123,18 @@ class Node {
UNUSED(is_linear);
}

virtual void calc(int x, int y, int z, int w, int lim) {
UNUSED(x);
UNUSED(y);
UNUSED(z);
UNUSED(w);
}

virtual void calc(int idx, int lim) {
UNUSED(idx);
UNUSED(lim);
}

/// Generates the variable that stores the thread's/work-item's offset into
/// the memory.
///
Expand Down Expand Up @@ -132,19 +194,35 @@ class Node {

// Returns true if this node is a Buffer
virtual bool isBuffer() const { return false; }

/// Returns true if the buffer is linear
virtual bool isLinear(dim_t dims[4]) const {
UNUSED(dims);
return true;
}
std::string getTypeStr() const { return m_type_str; }

/// Returns the string representation of the type
std::string getTypeStr() const { return getFullName(m_type); }

/// Returns the height of the JIT tree from this node
int getHeight() const { return m_height; }
std::string getNameStr() const { return m_name_str; }

virtual ~Node() {}
/// Returns the short name for this type
/// \note For the shift node this is "Sh" appended by the short name of the
/// type
virtual std::string getNameStr() const { return getShortName(m_type); }

/// Default destructor
virtual ~Node() = default;
};

struct Node_ids {
std::array<int, Node::kMaxChildren> child_ids;
int id;
};

std::string getFuncName(const std::vector<Node *> &output_nodes,
const std::vector<Node *> &full_nodes,
const std::vector<Node_ids> &full_ids, bool is_linear);

} // namespace common
12 changes: 8 additions & 4 deletions src/backend/common/jit/ScalarNode.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,9 @@
********************************************************/

#pragma once
#include <backend.hpp>
#include <common/jit/Node.hpp>
#include <af/traits.hpp>

#include <math.hpp>
#include <types.hpp>
Expand All @@ -23,20 +25,20 @@ class ScalarNode : public common::Node {

public:
ScalarNode(T val)
: Node(detail::getFullName<T>(), detail::shortname<T>(false), 0, {})
: Node(static_cast<af::dtype>(af::dtype_traits<T>::af_type), 0, {})
, m_val(val) {}

void genKerName(std::stringstream& kerStream,
const common::Node_ids& ids) const final {
kerStream << "_" << m_name_str;
kerStream << "_" << getTypeStr();
kerStream << std::setw(3) << std::setfill('0') << std::dec << ids.id
<< std::dec;
}

void genParams(std::stringstream& kerStream, int id,
bool is_linear) const final {
UNUSED(is_linear);
kerStream << m_type_str << " scalar" << id << ", \n";
kerStream << getTypeStr() << " scalar" << id << ", \n";
}

int setArgs(int start_id, bool is_linear,
Expand All @@ -49,10 +51,12 @@ class ScalarNode : public common::Node {

void genFuncs(std::stringstream& kerStream,
const common::Node_ids& ids) const final {
kerStream << m_type_str << " val" << ids.id << " = scalar" << ids.id
kerStream << getTypeStr() << " val" << ids.id << " = scalar" << ids.id
<< ";\n";
}

std::string getNameStr() const final { return detail::shortname<T>(false); }

// Return the info for the params and the size of the buffers
virtual size_t getParamBytes() const final { return sizeof(T); }
};
Expand Down
18 changes: 10 additions & 8 deletions src/backend/common/jit/ShiftNodeBase.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -29,12 +29,9 @@ class ShiftNodeBase : public Node {
const std::array<int, 4> m_shifts;

public:
ShiftNodeBase(const char *type_str, const char *name_str,
std::shared_ptr<BufferNode> buffer_node,
ShiftNodeBase(const af::dtype type, std::shared_ptr<BufferNode> buffer_node,
const std::array<int, 4> shifts)
: Node(type_str, name_str, 0, {})
, m_buffer_node(buffer_node)
, m_shifts(shifts) {}
: Node(type, 0, {}), m_buffer_node(buffer_node), m_shifts(shifts) {}

bool isLinear(dim_t dims[4]) const final {
UNUSED(dims);
Expand All @@ -43,7 +40,7 @@ class ShiftNodeBase : public Node {

void genKerName(std::stringstream &kerStream,
const common::Node_ids &ids) const final {
kerStream << "_" << m_name_str;
kerStream << "_" << getNameStr();
kerStream << std::setw(3) << std::setfill('0') << std::dec << ids.id
<< std::dec;
}
Expand All @@ -69,17 +66,22 @@ class ShiftNodeBase : public Node {

void genOffsets(std::stringstream &kerStream, int id,
bool is_linear) const final {
detail::generateShiftNodeOffsets(kerStream, id, is_linear, m_type_str);
detail::generateShiftNodeOffsets(kerStream, id, is_linear,
getTypeStr());
}

void genFuncs(std::stringstream &kerStream,
const common::Node_ids &ids) const final {
detail::generateShiftNodeRead(kerStream, ids.id, m_type_str);
detail::generateShiftNodeRead(kerStream, ids.id, getTypeStr());
}

void getInfo(unsigned &len, unsigned &buf_count,
unsigned &bytes) const final {
m_buffer_node->getInfo(len, buf_count, bytes);
}

std::string getNameStr() const final {
return std::string("Sh") + getShortName(m_type);
}
};
} // namespace common
Loading