InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
tensordot.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_TENSORDOT_H_
2#define INFINI_OPS_BASE_TENSORDOT_H_
3
4#include <cassert>
5#include <cstddef>
6#include <cstdint>
7#include <vector>
8
9#include "operator.h"
10
11namespace infini::ops {
12
13class Tensordot : public Operator<Tensordot> {
14 public:
15 Tensordot(const Tensor a, const Tensor b, const int64_t dims, Tensor out)
16 : Tensordot{a, b, ExpandDims(a, b, dims), out} {}
17
18 Tensordot(const Tensor a, const Tensor b,
19 const std::vector<std::vector<int64_t>> dims, Tensor out)
20 : input_shape_{a.shape()},
21 input_strides_{a.strides()},
22 input_type_{a.dtype()},
23 other_shape_{b.shape()},
24 other_strides_{b.strides()},
25 other_type_{b.dtype()},
26 out_shape_{out.shape()},
27 out_strides_{out.strides()},
28 out_type_{out.dtype()},
29 dims_self_{DimGroup(dims, 0)},
30 dims_other_{DimGroup(dims, 1)},
31 device_index_{out.device().index()} {}
32
34 [[deprecated("Use an overload taking a single `dims` argument instead.")]]
35 Tensordot(const Tensor input, const Tensor other,
36 const std::vector<int64_t> dims_self,
37 const std::vector<int64_t> dims_other, Tensor out)
38 : input_shape_{input.shape()},
39 input_strides_{input.strides()},
40 input_type_{input.dtype()},
41 other_shape_{other.shape()},
42 other_strides_{other.strides()},
43 other_type_{other.dtype()},
44 out_shape_{out.shape()},
45 out_strides_{out.strides()},
46 out_type_{out.dtype()},
47 dims_self_{dims_self},
48 dims_other_{dims_other},
49 device_index_{out.device().index()} {}
50
51 void operator()(const Tensor a, const Tensor b, const int64_t dims,
52 Tensor out) const {
53 const auto expanded_dims = ExpandDims(a, b, dims);
54 (*this)(a, b, expanded_dims[0], expanded_dims[1], out);
55 }
56
57 void operator()(const Tensor a, const Tensor b,
58 const std::vector<std::vector<int64_t>> dims,
59 Tensor out) const {
60 (*this)(a, b, DimGroup(dims, 0), DimGroup(dims, 1), out);
61 }
62
64 [[deprecated("Use an overload taking a single `dims` argument instead.")]]
65 virtual void operator()(const Tensor input, const Tensor other,
66 const std::vector<int64_t> dims_self,
67 const std::vector<int64_t> dims_other,
68 Tensor out) const = 0;
69
70 protected:
71 Tensor::Shape input_shape_;
72
73 Tensor::Strides input_strides_;
74
75 DataType input_type_;
76
77 Tensor::Shape other_shape_;
78
79 Tensor::Strides other_strides_;
80
81 DataType other_type_;
82
83 Tensor::Shape out_shape_;
84
85 Tensor::Strides out_strides_;
86
87 DataType out_type_;
88
89 std::vector<int64_t> dims_self_{};
90
91 std::vector<int64_t> dims_other_{};
92
94
95 private:
96 static std::vector<std::vector<int64_t>> ExpandDims(const Tensor a,
97 const Tensor b,
98 const int64_t dims) {
99 assert(dims >= 0 && dims <= static_cast<int64_t>(a.ndim()) &&
100 dims <= static_cast<int64_t>(b.ndim()) &&
101 "`Tensordot` expects non-negative `dims` no greater than either "
102 "input rank");
103
104 std::vector<std::vector<int64_t>> expanded_dims(2);
105 for (int64_t dim = 0; dim < dims; ++dim) {
106 expanded_dims[0].push_back(dim - dims);
107 expanded_dims[1].push_back(dim);
108 }
109
110 return expanded_dims;
111 }
112
113 static const std::vector<int64_t>& DimGroup(
114 const std::vector<std::vector<int64_t>>& dims, const std::size_t index) {
115 assert(dims.size() == 2 &&
116 "`Tensordot` expects `dims` to contain exactly two dimension "
117 "groups");
118 return dims[index];
119 }
120};
121
122} // namespace infini::ops
123
124#endif
Definition generated/include/operator.h:282
Definition tensordot.h:13
DataType out_type_
Definition tensordot.h:87
Tensor::Strides input_strides_
Definition tensordot.h:73
Tensor::Shape out_shape_
Definition tensordot.h:83
DataType other_type_
Definition tensordot.h:81
Tensordot(const Tensor input, const Tensor other, const std::vector< int64_t > dims_self, const std::vector< int64_t > dims_other, Tensor out)
Definition tensordot.h:35
std::vector< int64_t > dims_other_
Definition tensordot.h:91
Tensor::Strides other_strides_
Definition tensordot.h:79
Tensor::Strides out_strides_
Definition tensordot.h:85
int device_index_
Definition tensordot.h:93
DataType input_type_
Definition tensordot.h:75
Tensor::Shape input_shape_
Definition tensordot.h:71
void operator()(const Tensor a, const Tensor b, const std::vector< std::vector< int64_t > > dims, Tensor out) const
Definition tensordot.h:57
Tensor::Shape other_shape_
Definition tensordot.h:77
void operator()(const Tensor a, const Tensor b, const int64_t dims, Tensor out) const
Definition tensordot.h:51
std::vector< int64_t > dims_self_
Definition tensordot.h:89
Tensordot(const Tensor a, const Tensor b, const int64_t dims, Tensor out)
Definition tensordot.h:15
virtual void operator()(const Tensor input, const Tensor other, const std::vector< int64_t > dims_self, const std::vector< int64_t > dims_other, Tensor out) const =0
Tensordot(const Tensor a, const Tensor b, const std::vector< std::vector< int64_t > > dims, Tensor out)
Definition tensordot.h:18
Definition generated/include/operator.h:28
infini::rt::TensorView Tensor
Definition tensor.h:8