| 295 | } // namespace |
| 296 | |
| 297 | Status 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) |
nothing calls this directly
no test coverage detected