InfiniOps
Operator Library for Accelerators
Loading...
Searching...
No Matches
internal_chunk_cat.h
Go to the documentation of this file.
1#ifndef INFINI_OPS_BASE_INTERNAL_CHUNK_CAT_H_
2#define INFINI_OPS_BASE_INTERNAL_CHUNK_CAT_H_
3
4#include <vector>
5
6#include "operator.h"
7
8namespace infini::ops::internal {
9
10class ChunkCat : public Operator<ChunkCat> {
11 public:
12 ChunkCat(const std::vector<Tensor> tensors, const int64_t dim,
13 const int64_t num_chunks, Tensor out)
14 : out_shape_{out.shape()},
15 out_strides_{out.strides()},
16 out_type_{out.dtype()},
17 tensors_{tensors},
18 dim_{dim},
19 num_chunks_{num_chunks},
20 device_index_{out.device().index()} {}
21
22 virtual void operator()(const std::vector<Tensor> tensors, const int64_t dim,
23 const int64_t num_chunks, Tensor out) const = 0;
24
25 protected:
26 Tensor::Shape out_shape_;
27
28 Tensor::Strides out_strides_;
29
30 DataType out_type_;
31
32 std::vector<Tensor> tensors_{};
33
34 int64_t dim_{};
35
36 int64_t num_chunks_{};
37
39};
40
41} // namespace infini::ops::internal
42
43#endif
Definition generated/include/operator.h:282
Definition internal_chunk_cat.h:10
std::vector< Tensor > tensors_
Definition internal_chunk_cat.h:32
Tensor::Strides out_strides_
Definition internal_chunk_cat.h:28
Tensor::Shape out_shape_
Definition internal_chunk_cat.h:26
int64_t num_chunks_
Definition internal_chunk_cat.h:36
int device_index_
Definition internal_chunk_cat.h:38
virtual void operator()(const std::vector< Tensor > tensors, const int64_t dim, const int64_t num_chunks, Tensor out) const =0
int64_t dim_
Definition internal_chunk_cat.h:34
DataType out_type_
Definition internal_chunk_cat.h:30
ChunkCat(const std::vector< Tensor > tensors, const int64_t dim, const int64_t num_chunks, Tensor out)
Definition internal_chunk_cat.h:12
Definition internal_add_relu.h:6
infini::rt::TensorView Tensor
Definition tensor.h:8