| 578 | |
| 579 | template <typename T> |
| 580 | Status SparseTensor::Split(const SparseTensor& input_tensor, |
| 581 | const int split_dim, const int num_split, |
| 582 | std::vector<SparseTensor>* result) { |
| 583 | std::vector<Tensor> output_indices; |
| 584 | std::vector<Tensor> output_values; |
| 585 | std::vector<TensorShape> output_shapes; |
| 586 | output_indices.reserve(num_split); |
| 587 | output_values.reserve(num_split); |
| 588 | output_shapes.reserve(num_split); |
| 589 | |
| 590 | std::vector<typename TTypes<int64>::Matrix> output_indices_t; |
| 591 | std::vector<typename TTypes<T>::Vec> output_values_t; |
| 592 | output_indices_t.reserve(num_split); |
| 593 | output_values_t.reserve(num_split); |
| 594 | auto input_values_t = input_tensor.values().vec<T>(); |
| 595 | auto input_indices_t = input_tensor.indices().matrix<int64>(); |
| 596 | |
| 597 | std::vector<int> num_values(num_split, 0); |
| 598 | const int num_dim = input_tensor.shape().size(); |
| 599 | const int split_dim_size = input_tensor.shape()[split_dim]; |
| 600 | const int split_size = split_dim_size / num_split; |
| 601 | |
| 602 | if (!(num_split > 0 && num_split <= split_dim_size)) { |
| 603 | return Status(error::INVALID_ARGUMENT, |
| 604 | strings::StrCat("num_split must be in the interval (0, ", |
| 605 | split_dim_size, "]")); |
| 606 | } |
| 607 | if (!(split_dim >= 0 && split_dim < num_dim)) { |
| 608 | return Status( |
| 609 | error::INVALID_ARGUMENT, |
| 610 | strings::StrCat("num_dim must be in the interval [0, ", num_dim, ")")); |
| 611 | } |
| 612 | |
| 613 | const int residual = split_dim_size % num_split; |
| 614 | for (int i = 0; i < input_tensor.indices().dim_size(0); ++i) { |
| 615 | const int dim = input_tensor.indices().matrix<int64>()(i, split_dim); |
| 616 | int slice_index = GetSliceIndex(dim, split_size, residual); |
| 617 | num_values[slice_index]++; |
| 618 | } |
| 619 | |
| 620 | for (int i = 0; i < num_split; ++i) { |
| 621 | // TODO(ataei): Pass an allocator to avoid allocating large memory buffer. |
| 622 | output_indices.emplace_back(DT_INT64, |
| 623 | TensorShape({num_values[i], num_dim})); |
| 624 | output_values.emplace_back(DataTypeToEnum<T>::v(), |
| 625 | TensorShape({num_values[i]})); |
| 626 | output_shapes.emplace_back(input_tensor.shape()); |
| 627 | output_indices_t.emplace_back(output_indices[i].matrix<int64>()); |
| 628 | output_values_t.emplace_back(output_values[i].vec<T>()); |
| 629 | const int size = GetSliceShape(i, split_size, residual); |
| 630 | output_shapes[i].set_dim(split_dim, size); |
| 631 | } |
| 632 | |
| 633 | std::vector<int> values_inserted_in_slice(num_split, 0); |
| 634 | for (int i = 0; i < input_tensor.indices().dim_size(0); ++i) { |
| 635 | const int dim = input_indices_t(i, split_dim); |
| 636 | const int slice_index = GetSliceIndex(dim, split_size, residual); |
| 637 | const int slice_dim = values_inserted_in_slice[slice_index]++; |
no test coverage detected