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

Method Compute

tensorflow/core/kernels/list_kernels.h:688–773  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

686 }
687
688 void Compute(OpKernelContext* c) override {
689 const TensorList* tensor_list = nullptr;
690 OP_REQUIRES_OK(c, GetInputList(c, 0, &tensor_list));
691 OP_REQUIRES(
692 c, element_dtype_ == tensor_list->element_dtype,
693 errors::InvalidArgument(
694 "Invalid data types; op elements ", DataTypeString(element_dtype_),
695 " but list elements ", DataTypeString(tensor_list->element_dtype)));
696 const Tensor& indices = c->input(1);
697 PartialTensorShape partial_element_shape;
698 OP_REQUIRES_OK(c, GetElementShapeFromInput(c, *tensor_list, 2,
699 &partial_element_shape));
700 OP_REQUIRES(
701 c, partial_element_shape.IsFullyDefined() || indices.NumElements() > 0,
702 errors::InvalidArgument("Tried to gather 0-elements from "
703 "a list with non-fully-defined shape: ",
704 partial_element_shape.DebugString()));
705
706 // Check that `element_shape` input tensor is compatible with the shapes of
707 // element tensors.
708 if (!tensor_list->element_shape.IsFullyDefined()) {
709 for (int index = 0; index < indices.NumElements(); ++index) {
710 const int i = indices.flat<int32>()(index);
711 const Tensor& t = tensor_list->tensors()[i];
712 if (t.dtype() != DT_INVALID) {
713 PartialTensorShape tmp = partial_element_shape;
714 OP_REQUIRES_OK(c, tmp.MergeWith(t.shape(), &partial_element_shape));
715 }
716 }
717 }
718
719 // Compute the shape of the output tensor by pre-pending the leading dim to
720 // the element_shape.
721 TensorShape element_shape;
722 OP_REQUIRES(
723 c, partial_element_shape.AsTensorShape(&element_shape),
724 errors::InvalidArgument("Tried to gather uninitialized tensors from a ",
725 "list with non-fully-defined element_shape: ",
726 partial_element_shape.DebugString()));
727 TensorShape output_shape = element_shape;
728 output_shape.InsertDim(0, indices.NumElements());
729 Tensor* output;
730 OP_REQUIRES_OK(c, c->allocate_output(0, output_shape, &output));
731 if (output->NumElements() == 0) {
732 return;
733 }
734
735 ConstMatrixVector inputs_flat;
736 inputs_flat.reserve(indices.NumElements());
737 Tensor zeros;
738 for (int index = 0; index < indices.NumElements(); ++index) {
739 const int i = indices.flat<int32>()(index);
740 OP_REQUIRES(
741 c, i < tensor_list->tensors().size(),
742 errors::InvalidArgument("Index ", i, " out o range; list only has ",
743 tensor_list->tensors().size(), " elements."));
744 const Tensor& t = tensor_list->tensors()[i];
745 if (t.dtype() != DT_INVALID) {

Callers

nothing calls this directly

Calls 15

GetInputListFunction · 0.85
InvalidArgumentFunction · 0.85
GetElementShapeFromInputFunction · 0.85
IsFullyDefinedMethod · 0.80
tensorsMethod · 0.80
MergeWithMethod · 0.80
AsTensorShapeMethod · 0.80
allocate_outputMethod · 0.80
set_on_hostMethod · 0.80
DataTypeStringFunction · 0.50
inputMethod · 0.45
NumElementsMethod · 0.45

Tested by

no test coverage detected