1#ifndef INFINI_OPS_BASE_COPY_H_
2#define INFINI_OPS_BASE_COPY_H_
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");
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();
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);
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