| 50 | : OpKernel(context) {} |
| 51 | |
| 52 | void Compute(OpKernelContext* context) override { |
| 53 | const Tensor& input = context->input(0); |
| 54 | |
| 55 | // MatrixDiagPart and MatrixDiagPartV2 both use this OpKernel. |
| 56 | // MatrixDiagPart only has one input, so we have to check the number of |
| 57 | // inputs before reading additional parameters in MatrixDiagV2. |
| 58 | int32 lower_diag_index = 0; |
| 59 | int32 upper_diag_index = 0; |
| 60 | T padding_value(0); |
| 61 | |
| 62 | // MatrixDiagPartV2-specific. |
| 63 | if (context->num_inputs() > 1) { |
| 64 | auto& diag_index = context->input(1); |
| 65 | OP_REQUIRES(context, |
| 66 | TensorShapeUtils::IsScalar(diag_index.shape()) || |
| 67 | TensorShapeUtils::IsVector(diag_index.shape()), |
| 68 | errors::InvalidArgument( |
| 69 | "diag_index must be a scalar or vector, received shape: ", |
| 70 | diag_index.shape().DebugString())); |
| 71 | OP_REQUIRES(context, diag_index.NumElements() > 0, |
| 72 | errors::InvalidArgument( |
| 73 | "Expected diag_index to have at least 1 element")); |
| 74 | lower_diag_index = diag_index.flat<int32>()(0); |
| 75 | upper_diag_index = lower_diag_index; |
| 76 | if (TensorShapeUtils::IsVector(diag_index.shape())) { |
| 77 | auto diag_index_size = diag_index.dim_size(0); |
| 78 | OP_REQUIRES( |
| 79 | context, 0 < diag_index_size && diag_index_size <= 2, |
| 80 | errors::InvalidArgument( |
| 81 | "diag_index must have only one or two elements, received ", |
| 82 | diag_index_size, " elements.")); |
| 83 | if (diag_index_size > 1) { |
| 84 | upper_diag_index = diag_index.flat<int32>()(1); |
| 85 | } |
| 86 | } |
| 87 | const Tensor& padding_in = context->input(2); |
| 88 | OP_REQUIRES(context, padding_in.NumElements() == 1, |
| 89 | errors::InvalidArgument("Padding must be scalar.")); |
| 90 | padding_value = padding_in.flat<T>()(0); |
| 91 | } |
| 92 | const TensorShape& input_shape = input.shape(); |
| 93 | |
| 94 | // Preliminary validation of sizes. |
| 95 | OP_REQUIRES(context, TensorShapeUtils::IsMatrixOrHigher(input_shape), |
| 96 | errors::InvalidArgument( |
| 97 | "input must be at least 2-dim, received shape: ", |
| 98 | input.shape().DebugString())); |
| 99 | |
| 100 | // Make sure lower_diag_index and upper_diag_index is valid. |
| 101 | const int rank = input_shape.dims(); |
| 102 | const Eigen::Index num_rows = input_shape.dim_size(rank - 2); |
| 103 | const Eigen::Index num_cols = input_shape.dim_size(rank - 1); |
| 104 | OP_REQUIRES( // Checks lower_diag_index == 0 for when matrix shape = 0. |
| 105 | context, |
| 106 | (-num_rows < lower_diag_index && lower_diag_index < num_cols) || |
| 107 | lower_diag_index == 0, |
| 108 | errors::InvalidArgument( |
| 109 | "lower_diag_index is out of bound: ", lower_diag_index, |
nothing calls this directly
no test coverage detected