| 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) { |
nothing calls this directly
no test coverage detected