| 937 | std::vector<phi::dtype::complex<double>>* dst); |
| 938 | |
| 939 | DenseTensor ReshapeToMatrix(const DenseTensor& src, int num_col_dims) { |
| 940 | int rank = src.dims().size(); |
| 941 | PADDLE_ENFORCE_GE( |
| 942 | rank, |
| 943 | 2, |
| 944 | common::errors::InvalidArgument( |
| 945 | "'ReshapeToMatrix()' is only used for flatten high rank " |
| 946 | "tensors to matrixs. The dimensions of DenseTensor must be " |
| 947 | "greater or equal than 2. " |
| 948 | "But received dimensions of DenseTensor is %d", |
| 949 | rank)); |
| 950 | if (rank == 2) { |
| 951 | return src; |
| 952 | } |
| 953 | DenseTensor res; |
| 954 | res.ShareDataWith(src); |
| 955 | res.Resize(common::flatten_to_2d(src.dims(), num_col_dims)); |
| 956 | return res; |
| 957 | } |
| 958 | |
| 959 | template <typename T> |
| 960 | T GetValue(const DenseTensor* x) { |