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

Function EinsumShape

tensorflow/core/framework/common_shape_fns.cc:297–446  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

295} // namespace
296
297Status EinsumShape(shape_inference::InferenceContext* c) {
298 // We assume that the equation has a valid format. Either (x),(y)->(z)
299 // or (x)->(z), where each of (x), (y) and (z) are concatenation of zero or
300 // more latin alphabets and contains at most one ellipsis ('...').
301 string equation;
302 TF_RETURN_IF_ERROR(c->GetAttr("equation", &equation));
303 gtl::InlinedVector<string, 2> input_labels;
304 string output_labels;
305 TF_RETURN_IF_ERROR(
306 ParseEinsumEquation(equation, &input_labels, &output_labels));
307
308 if (c->num_inputs() == 0 || c->num_inputs() > 2) {
309 return errors::InvalidArgument("Expected either 1 or 2 inputs but got: ",
310 c->num_inputs());
311 }
312 if (c->num_inputs() != input_labels.size()) {
313 return errors::InvalidArgument("Expected ", input_labels.size(),
314 " inputs for equation ", equation,
315 " but got: ", c->num_inputs());
316 }
317
318 // Validate input subscripts, build the label to dimension mapping and obtain
319 // the broadcast shapes that map to ellipsis.
320 absl::flat_hash_map<char, DimensionHandle> label_to_dimension;
321 gtl::InlinedVector<ShapeHandle, 2> input_bcast_shapes(c->num_inputs());
322 for (int i = 0; i < c->num_inputs(); ++i) {
323 bool has_ellipsis = false;
324 TF_RETURN_IF_ERROR(ValidateEinsumEllipsis(input_labels[i], &has_ellipsis));
325 ShapeHandle input_shape = c->input(i);
326 // Validate that the input rank is sufficient for the given number of named
327 // labels.
328 if (c->RankKnown(input_shape)) {
329 if (has_ellipsis) {
330 const int num_named_labels =
331 static_cast<int>(input_labels[i].size()) - 3;
332 TF_RETURN_WITH_CONTEXT_IF_ERROR(
333 c->WithRankAtLeast(input_shape, num_named_labels, &input_shape),
334 " for ", i, "th input and equation: ", equation);
335 } else {
336 const int num_named_labels = static_cast<int>(input_labels[i].size());
337 TF_RETURN_WITH_CONTEXT_IF_ERROR(
338 c->WithRank(input_shape, num_named_labels, &input_shape), " for ",
339 i, "th input and equation: ", equation);
340 }
341 }
342
343 bool seen_ellipsis = false;
344 input_bcast_shapes[i] = c->Scalar();
345 // Run through the input labels; populate label_to_dimension mapping and
346 // compute the broadcast shapes corresponding to the ellipsis (if present).
347 for (int label_idx = 0; label_idx < input_labels[i].size(); ++label_idx) {
348 const char label = input_labels[i][label_idx];
349 // Calculate the input axis that the current label is referring to. After
350 // the ellipsis, the axis may be found by using negative indices; i.e the
351 // (rank - k)th dimension corresponds to the (num_labels - k)th label.
352 const int64 axis_before_ellipsis = label_idx;
353 const int64 axis_after_ellipsis =
354 c->RankKnown(input_shape)

Callers

nothing calls this directly

Calls 15

ParseEinsumEquationFunction · 0.85
InvalidArgumentFunction · 0.85
ValidateEinsumEllipsisFunction · 0.85
RankKnownMethod · 0.80
WithRankAtLeastMethod · 0.80
UnknownShapeMethod · 0.80
SubshapeMethod · 0.80
UnknownDimMethod · 0.80
containsMethod · 0.80
WithRankAtMostMethod · 0.80
GetAttrMethod · 0.45

Tested by

no test coverage detected