MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / Split

Method Split

tensorflow/core/util/sparse/sparse_tensor.h:580–660  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

578
579template <typename T>
580Status 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]++;

Callers 2

mainFunction · 0.80
camelCaseFunction · 0.80

Calls 15

GetSliceIndexFunction · 0.85
GetSliceShapeFunction · 0.85
GetDimensionInSliceFunction · 0.85
CreateFunction · 0.70
StatusClass · 0.50
StrCatFunction · 0.50
TensorShapeClass · 0.50
reserveMethod · 0.45
valuesMethod · 0.45
indicesMethod · 0.45
sizeMethod · 0.45
shapeMethod · 0.45

Tested by

no test coverage detected