Skip to content

Commit f9e33b1

Browse files
umar4569prady9
authored andcommitted
Add static asserts and move constructors for several classes
1 parent a8e86cd commit f9e33b1

13 files changed

Lines changed: 169 additions & 28 deletions

File tree

‎.github/pull_request_template.md‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11

2-
<!--
2+
<!--
33
Short description of change
44
55
This should be one or two sentences that describe the overall
@@ -16,7 +16,7 @@ Additional information about the PR answering following questions:
1616
* More detail if necessary to describe all commits in pull request.
1717
* Why these changes are necessary.
1818
* Potential impact on specific hardware, software or backends.
19-
* New functions and their functionallity.
19+
* New functions and their functionality.
2020
* Can this PR be backported to older versions?
2121
* Future changes not implemented in this PR.
2222
-->

‎src/backend/common/ArrayInfo.hpp‎

Lines changed: 20 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,10 @@ class ArrayInfo {
5656
, dim_strides(stride)
5757
, is_sparse(false) {
5858
setId(id);
59+
static_assert(std::is_move_assignable<ArrayInfo>::value,
60+
"ArrayInfo is not move assignable");
61+
static_assert(std::is_move_constructible<ArrayInfo>::value,
62+
"ArrayInfo is not move constructible");
5963
static_assert(
6064
offsetof(ArrayInfo, devId) == 0,
6165
"ArrayInfo::devId must be the first member variable of ArrayInfo. \
@@ -79,10 +83,24 @@ class ArrayInfo {
7983
This is then used in the unified backend to check mismatched arrays.");
8084
}
8185

82-
// Copy constructors are deprecated if there is a
83-
// user-defined destructor in c++11
8486
ArrayInfo() = default;
8587
ArrayInfo(const ArrayInfo& other) = default;
88+
ArrayInfo(ArrayInfo&& other) = default;
89+
90+
ArrayInfo& operator=(ArrayInfo other) noexcept {
91+
swap(other);
92+
return *this;
93+
}
94+
95+
void swap(ArrayInfo& other) noexcept {
96+
using std::swap;
97+
swap(devId, other.devId);
98+
swap(type, other.type);
99+
swap(dim_size, other.dim_size);
100+
swap(offset, other.offset);
101+
swap(dim_strides, other.dim_strides);
102+
swap(is_sparse, other.is_sparse);
103+
}
86104

87105
const af_dtype& getType() const { return type; }
88106

‎src/backend/common/half.hpp‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -844,6 +844,14 @@ class alignas(2) half {
844844
data_(bits)
845845
#endif
846846
{
847+
#ifndef __CUDACC_RTC__
848+
static_assert(std::is_standard_layout<half>::value,
849+
"half must be a standard layout type");
850+
static_assert(std::is_nothrow_move_assignable<half>::value,
851+
"half is not move assignable");
852+
static_assert(std::is_nothrow_move_constructible<half>::value,
853+
"half is not move constructible");
854+
#endif
847855
}
848856

849857
#if defined(__CUDA_ARCH__)

‎src/backend/common/jit/BufferNodeBase.hpp‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,9 @@ class BufferNodeBase : public common::Node {
2828
bool m_linear_buffer;
2929

3030
public:
31-
BufferNodeBase(af::dtype type) : Node(type, 0, {}) {}
31+
BufferNodeBase(af::dtype type) : Node(type, 0, {}) {
32+
// This class is not movable because of std::once_flag
33+
}
3234

3335
bool isBuffer() const final { return true; }
3436

‎src/backend/common/jit/NaryNode.hpp‎

Lines changed: 27 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -24,9 +24,9 @@ namespace common {
2424

2525
class NaryNode : public Node {
2626
private:
27-
const int m_num_children;
28-
const int m_op;
29-
const std::string m_op_str;
27+
int m_num_children;
28+
int m_op;
29+
std::string m_op_str;
3030

3131
public:
3232
NaryNode(const af::dtype type, const char *op_str, const int num_children,
@@ -39,7 +39,30 @@ class NaryNode : public Node {
3939
children))
4040
, m_num_children(num_children)
4141
, m_op(op)
42-
, m_op_str(op_str) {}
42+
, m_op_str(op_str) {
43+
static_assert(std::is_nothrow_move_assignable<NaryNode>::value,
44+
"NaryNode is not move assignable");
45+
static_assert(std::is_nothrow_move_constructible<NaryNode>::value,
46+
"NaryNode is not move constructible");
47+
}
48+
49+
NaryNode(NaryNode &&other) = default;
50+
51+
NaryNode(const NaryNode &other) = default;
52+
53+
/// Default copy assignment operator
54+
NaryNode &operator=(const NaryNode &node) = default;
55+
56+
/// Default move assignment operator
57+
NaryNode &operator=(NaryNode &&node) noexcept = default;
58+
59+
void swap(NaryNode &other) noexcept {
60+
using std::swap;
61+
Node::swap(other);
62+
swap(m_num_children, other.m_num_children);
63+
swap(m_op, other.m_op);
64+
swap(m_op_str, other.m_op_str);
65+
}
4366

4467
void genKerName(std::stringstream &kerStream,
4568
const common::Node_ids &ids) const final {

‎src/backend/common/jit/Node.hpp‎

Lines changed: 21 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
#include <memory>
2121
#include <string>
2222
#include <unordered_map>
23+
#include <utility>
2324
#include <vector>
2425

2526
enum class kJITHeuristics {
@@ -80,30 +81,37 @@ class Node {
8081
static const int kMaxChildren = 3;
8182

8283
protected:
83-
const std::array<Node_ptr, kMaxChildren> m_children;
84-
const af::dtype m_type;
85-
const int m_height;
84+
std::array<Node_ptr, kMaxChildren> m_children;
85+
af::dtype m_type;
86+
int m_height;
8687

8788
template<typename T>
8889
friend class NodeIterator;
8990

91+
void swap(Node &other) noexcept {
92+
using std::swap;
93+
for (int i = 0; i < kMaxChildren; i++) {
94+
swap(m_children[i], other.m_children[i]);
95+
}
96+
swap(m_type, other.m_type);
97+
swap(m_height, other.m_height);
98+
}
99+
90100
public:
101+
Node() = default;
91102
Node(const af::dtype type, const int height,
92103
const std::array<Node_ptr, kMaxChildren> children)
93-
: m_children(children), m_type(type), m_height(height) {}
94-
95-
/// Default copy constructor
96-
Node(Node &node) = default;
104+
: m_children(children), m_type(type), m_height(height) {
105+
static_assert(std::is_nothrow_move_assignable<Node>::value,
106+
"Node is not move assignable");
107+
}
97108

98-
/// Default move constructor
99-
Node(Node &&node) = default;
109+
/// Default copy constructor operator
110+
Node(const Node &node) = default;
100111

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

104-
/// Default move assignment operator
105-
Node &operator=(Node &&node) = default;
106-
107115
int getNodesMap(Node_map_t &node_map, std::vector<Node *> &full_nodes,
108116
std::vector<Node_ids> &full_ids);
109117

@@ -213,7 +221,7 @@ class Node {
213221
virtual std::string getNameStr() const { return getShortName(m_type); }
214222

215223
/// Default destructor
216-
virtual ~Node() = default;
224+
virtual ~Node() noexcept = default;
217225
};
218226

219227
struct Node_ids {

‎src/backend/common/jit/ScalarNode.hpp‎

Lines changed: 25 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,31 @@ class ScalarNode : public common::Node {
2626
public:
2727
ScalarNode(T val)
2828
: Node(static_cast<af::dtype>(af::dtype_traits<T>::af_type), 0, {})
29-
, m_val(val) {}
29+
, m_val(val) {
30+
static_assert(std::is_nothrow_move_assignable<ScalarNode>::value,
31+
"ScalarNode is not move assignable");
32+
static_assert(std::is_nothrow_move_constructible<ScalarNode>::value,
33+
"ScalarNode is not move constructible");
34+
}
35+
36+
/// Default move copy constructor
37+
ScalarNode(const ScalarNode& other) = default;
38+
39+
/// Default move constructor
40+
ScalarNode(ScalarNode&& other) = default;
41+
42+
/// Default move/copy assignment operator(Rule of 4)
43+
ScalarNode& operator=(ScalarNode node) noexcept {
44+
swap(node);
45+
return *this;
46+
}
47+
48+
// Swap specilization
49+
void swap(ScalarNode& other) noexcept {
50+
using std::swap;
51+
Node::swap(other);
52+
swap(m_val, other.m_val);
53+
}
3054

3155
void genKerName(std::stringstream& kerStream,
3256
const common::Node_ids& ids) const final {

‎src/backend/common/jit/ShiftNodeBase.hpp‎

Lines changed: 27 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,12 +26,37 @@ template<typename BufferNode>
2626
class ShiftNodeBase : public Node {
2727
private:
2828
std::shared_ptr<BufferNode> m_buffer_node;
29-
const std::array<int, 4> m_shifts;
29+
std::array<int, 4> m_shifts;
3030

3131
public:
3232
ShiftNodeBase(const af::dtype type, std::shared_ptr<BufferNode> buffer_node,
3333
const std::array<int, 4> shifts)
34-
: Node(type, 0, {}), m_buffer_node(buffer_node), m_shifts(shifts) {}
34+
: Node(type, 0, {}), m_buffer_node(buffer_node), m_shifts(shifts) {
35+
static_assert(std::is_nothrow_move_assignable<ShiftNodeBase>::value,
36+
"ShiftNode is not move assignable");
37+
static_assert(std::is_nothrow_move_constructible<ShiftNodeBase>::value,
38+
"ShiftNode is not move constructible");
39+
}
40+
41+
/// Default move copy constructor
42+
ShiftNodeBase(const ShiftNodeBase &other) = default;
43+
44+
/// Default move constructor
45+
ShiftNodeBase(ShiftNodeBase &&other) = default;
46+
47+
/// Default move/copy assignment operator(Rule of 4)
48+
ShiftNodeBase &operator=(ShiftNodeBase node) noexcept {
49+
swap(node);
50+
return *this;
51+
}
52+
53+
// Swap specilization
54+
void swap(ShiftNodeBase &other) noexcept {
55+
using std::swap;
56+
Node::swap(other);
57+
swap(m_buffer_node, other.m_buffer_node);
58+
swap(m_shifts, other.m_shifts);
59+
}
3560

3661
bool isLinear(dim_t dims[4]) const final {
3762
UNUSED(dims);

‎src/backend/common/jit/UnaryNode.hpp‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,11 @@ namespace common {
1515
class UnaryNode : public NaryNode {
1616
public:
1717
UnaryNode(const af::dtype type, const char *op_str, Node_ptr child, int op)
18-
: NaryNode(type, op_str, 1, {{child}}, op, child->getHeight() + 1) {}
18+
: NaryNode(type, op_str, 1, {{child}}, op, child->getHeight() + 1) {
19+
static_assert(std::is_nothrow_move_assignable<UnaryNode>::value,
20+
"UnaryNode is not move assignable");
21+
static_assert(std::is_nothrow_move_constructible<UnaryNode>::value,
22+
"UnaryNode is not move constructible");
23+
}
1924
};
2025
} // namespace common

‎src/backend/cpu/Array.cpp‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -80,6 +80,10 @@ Array<T>::Array(const dim4 &dims, T *const in_data, bool is_device,
8080
, owner(true) {
8181
static_assert(is_standard_layout<Array<T>>::value,
8282
"Array<T> must be a standard layout type");
83+
static_assert(std::is_move_assignable<Array<T>>::value,
84+
"Array<T> is not move assignable");
85+
static_assert(std::is_move_constructible<Array<T>>::value,
86+
"Array<T> is not move constructible");
8387
static_assert(
8488
offsetof(Array<T>, info) == 0,
8589
"Array<T>::info must be the first member variable of Array<T>");

0 commit comments

Comments
 (0)