| 151 | explicit MatrixDiagOp(OpKernelConstruction* context) : OpKernel(context) {} |
| 152 | |
| 153 | void Compute(OpKernelContext* context) override { |
| 154 | const Tensor& diagonal = context->input(0); |
| 155 | |
| 156 | // MatrixDiag and MatrixDiagV2 both use this OpKernel. MatrixDiag only has |
| 157 | // one input, so we have to check the number of inputs before reading |
| 158 | // additional parameters in MatrixDiagV2. |
| 159 | int32 lower_diag_index = 0; |
| 160 | int32 upper_diag_index = 0; |
| 161 | int32 num_rows = -1; |
| 162 | int32 num_cols = -1; |
| 163 | T padding_value(0); |
| 164 | |
| 165 | // MatrixDiagOpV2-specific. |
| 166 | if (context->num_inputs() > 1) { |
| 167 | auto& diag_index = context->input(1); |
| 168 | OP_REQUIRES(context, |
| 169 | TensorShapeUtils::IsScalar(diag_index.shape()) || |
| 170 | TensorShapeUtils::IsVector(diag_index.shape()), |
| 171 | errors::InvalidArgument( |
| 172 | "diag_index must be a scalar or vector, received shape: ", |
| 173 | diag_index.shape().DebugString())); |
| 174 | OP_REQUIRES(context, diag_index.NumElements() > 0, |
| 175 | errors::InvalidArgument( |
| 176 | "Expected diag_index to have at least 1 element")); |
| 177 | lower_diag_index = diag_index.flat<int32>()(0); |
| 178 | upper_diag_index = lower_diag_index; |
| 179 | if (TensorShapeUtils::IsVector(diag_index.shape())) { |
| 180 | auto diag_index_size = diag_index.dim_size(0); |
| 181 | OP_REQUIRES( |
| 182 | context, 0 < diag_index_size && diag_index_size <= 2, |
| 183 | errors::InvalidArgument( |
| 184 | "diag_index must have only one or two elements, received ", |
| 185 | diag_index_size, " elements.")); |
| 186 | if (diag_index_size > 1) { |
| 187 | upper_diag_index = diag_index.flat<int32>()(1); |
| 188 | } |
| 189 | } |
| 190 | num_rows = context->input(2).flat<int32>()(0); |
| 191 | num_cols = context->input(3).flat<int32>()(0); |
| 192 | padding_value = context->input(4).flat<T>()(0); |
| 193 | } |
| 194 | |
| 195 | // Size validations. |
| 196 | const TensorShape& diagonal_shape = diagonal.shape(); |
| 197 | const int diag_rank = diagonal_shape.dims(); |
| 198 | const Eigen::Index num_diags = upper_diag_index - lower_diag_index + 1; |
| 199 | OP_REQUIRES(context, TensorShapeUtils::IsVectorOrHigher(diagonal_shape), |
| 200 | errors::InvalidArgument( |
| 201 | "diagonal must be at least 1-dim, received shape: ", |
| 202 | diagonal.shape().DebugString())); |
| 203 | OP_REQUIRES( |
| 204 | context, lower_diag_index <= upper_diag_index, |
| 205 | errors::InvalidArgument( |
| 206 | "lower_diag_index must not be larger than upper_diag_index: ", |
| 207 | lower_diag_index, " > ", upper_diag_index)); |
| 208 | OP_REQUIRES(context, |
| 209 | lower_diag_index == upper_diag_index || |
| 210 | diagonal_shape.dim_size(diag_rank - 2) == num_diags, |
nothing calls this directly
no test coverage detected