InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
copy.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_COPY_H_
2#define INFINI_OPS_BASE_COPY_H_
3
4#include <cassert>
5
6#include "operator.h"
7
8namespace infini::ops {
9
10class Copy : public Operator<Copy> {
11 public:
12 Copy(const Tensor src, const bool non_blocking, Tensor out)
13 : input_shape_{out.shape()},
15 input_type_{src.dtype()},
16 non_blocking_{non_blocking},
17 out_shape_{out.shape()},
18 out_strides_{out.strides()},
19 out_type_{out.dtype()},
20 output_size_{out.numel()},
21 ndim_{out.ndim()},
22 is_input_contiguous_{src.shape() == out.shape() && src.IsContiguous()},
23 is_out_contiguous_{out.IsContiguous()} {
24 assert(input_type_ == out_type_ &&
25 "`Copy` currently requires input and output dtypes to match");
26 assert(src.device() == out.device() &&
27 "`Copy` currently requires input and output on the same device");
28 assert(!out.HasBroadcastDim() &&
29 "`Copy` output must not have broadcasted dimensions");
30 }
31
32 virtual void operator()(const Tensor src, const bool non_blocking,
33 Tensor out) const = 0;
34
35 protected:
36 static Tensor::Strides BroadcastStrides(const Tensor src, const Tensor out) {
37 assert(src.ndim() <= out.ndim() &&
38 "`Copy` input rank must not exceed output rank");
39 Tensor::Strides strides(out.ndim(), 0);
40 auto offset = out.ndim() - src.ndim();
41
42 for (Tensor::Size i = 0; i < src.ndim(); ++i) {
43 auto out_dim = i + offset;
44 assert((src.size(i) == 1 || src.size(i) == out.size(out_dim)) &&
45 "`Copy` input shape must be broadcastable to output shape");
46 strides[out_dim] = src.size(i) == 1 ? 0 : src.stride(i);
47 }
48
49 return strides;
50 }
51
52 Tensor::Shape input_shape_;
53
54 Tensor::Strides input_strides_;
55
56 DataType input_type_;
57
58 bool non_blocking_{false};
59
60 Tensor::Shape out_shape_;
61
62 Tensor::Strides out_strides_;
63
64 DataType out_type_;
65
66 Tensor::Size output_size_{0};
67
68 Tensor::Size ndim_{0};
69
71
72 bool is_out_contiguous_{false};
73};
74
75} // namespace infini::ops
76
77#endif
Definition copy.h:10
Tensor::Size output_size_
Definition copy.h:66
bool non_blocking_
Definition copy.h:58
Tensor::Shape input_shape_
Definition copy.h:52
Tensor::Shape out_shape_
Definition copy.h:60
Copy(const Tensor src, const bool non_blocking, Tensor out)
Definition copy.h:12
DataType out_type_
Definition copy.h:64
DataType input_type_
Definition copy.h:56
Tensor::Strides input_strides_
Definition copy.h:54
virtual void operator()(const Tensor src, const bool non_blocking, Tensor out) const =0
Tensor::Strides out_strides_
Definition copy.h:62
bool is_input_contiguous_
Definition copy.h:70
bool is_out_contiguous_
Definition copy.h:72
Tensor::Size ndim_
Definition copy.h:68
static Tensor::Strides BroadcastStrides(const Tensor src, const Tensor out)
Definition copy.h:36
Definition generated/include/operator.h:282
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8