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

Method Compute

tensorflow/core/kernels/matrix_diag_op.cc:52–142  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls 14

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