12 Cat(
const std::vector<Tensor> tensors,
const int64_t dim,
Tensor out)
17 assert(!tensors.empty() &&
"`Cat` requires a non-empty tensor list");
19 auto ndim =
static_cast<int64_t
>(out.ndim());
20 dim_ = dim < 0 ? dim + ndim : dim;
21 assert(
dim_ >= 0 &&
dim_ < ndim &&
"`Cat` dim out of range");
23 Tensor::Size cat_size = 0;
24 for (
const auto& tensor : tensors) {
25 assert(tensor.ndim() == out.ndim() &&
26 "`Cat` requires all tensors to have the output rank");
27 assert(tensor.dtype() == out.dtype() &&
28 "`Cat` requires all tensors to have the output dtype");
30 for (Tensor::Size axis = 0; axis < out.ndim(); ++axis) {
31 if (axis !=
static_cast<Tensor::Size
>(
dim_)) {
32 assert(tensor.size(axis) == out.size(axis) &&
33 "`Cat` input dimensions must match outside `dim`");
36 cat_size += tensor.size(
dim_);
38 assert(cat_size == out.size(
dim_) &&
39 "`Cat` output size along `dim` must equal the input sum");