1#ifndef INFINI_OPS_BASE_TENSORDOT_H_
2#define INFINI_OPS_BASE_TENSORDOT_H_
16 :
Tensordot{a, b, ExpandDims(a, b, dims), out} {}
19 const std::vector<std::vector<int64_t>> dims,
Tensor out)
34 [[deprecated(
"Use an overload taking a single `dims` argument instead.")]]
36 const std::vector<int64_t> dims_self,
37 const std::vector<int64_t> dims_other,
Tensor out)
53 const auto expanded_dims = ExpandDims(a, b, dims);
54 (*this)(a, b, expanded_dims[0], expanded_dims[1], out);
58 const std::vector<std::vector<int64_t>> dims,
60 (*this)(a, b, DimGroup(dims, 0), DimGroup(dims, 1), out);
64 [[deprecated(
"Use an overload taking a single `dims` argument instead.")]]
66 const std::vector<int64_t> dims_self,
67 const std::vector<int64_t> dims_other,
96 static std::vector<std::vector<int64_t>> ExpandDims(
const Tensor a,
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 "
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);
110 return expanded_dims;
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 "
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