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

Method Compute

tensorflow/core/kernels/matrix_diag_op.cc:153–255  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls 15

InvalidArgumentFunction · 0.85
allocate_outputMethod · 0.80
ComputeFunction · 0.70
IsScalarFunction · 0.50
minFunction · 0.50
maxFunction · 0.50
inputMethod · 0.45
num_inputsMethod · 0.45
shapeMethod · 0.45
DebugStringMethod · 0.45
NumElementsMethod · 0.45
dim_sizeMethod · 0.45

Tested by

no test coverage detected