static*/
| 1511 | } |
| 1512 | |
| 1513 | /*static*/ absl::optional<std::vector<int64>> ShapeUtil::FindTranspose021( |
| 1514 | const Shape& a, const Shape& b) { |
| 1515 | if (!CompatibleIgnoringElementType(a, b)) { |
| 1516 | return absl::nullopt; |
| 1517 | } |
| 1518 | |
| 1519 | std::vector<int64> permutation(a.dimensions().size()); |
| 1520 | absl::Span<const int64> minor_to_major_a = LayoutUtil::MinorToMajor(a); |
| 1521 | std::vector<int64> major_to_minor_a(minor_to_major_a.rbegin(), |
| 1522 | minor_to_major_a.rend()); |
| 1523 | absl::Span<const int64> minor_to_major_b = LayoutUtil::MinorToMajor(b); |
| 1524 | std::vector<int64> major_to_minor_b(minor_to_major_b.rbegin(), |
| 1525 | minor_to_major_b.rend()); |
| 1526 | for (size_t i = 0; i < permutation.size(); ++i) { |
| 1527 | permutation[i] = PositionInContainer(major_to_minor_b, major_to_minor_a[i]); |
| 1528 | } |
| 1529 | |
| 1530 | std::vector<size_t> segments = ConsecutiveSegments(permutation); |
| 1531 | if ((3 == segments.size() && 0 == permutation[0]) || 2 == segments.size()) { |
| 1532 | Shape descending_layout_shape = |
| 1533 | ShapeUtil::MakeShapeWithDescendingLayoutAndSamePhysicalLayout(a); |
| 1534 | Shape normalized_shape = MergeDimensions(segments, descending_layout_shape); |
| 1535 | absl::Span<const int64> normalized_dims = |
| 1536 | AsInt64Slice(normalized_shape.dimensions()); |
| 1537 | std::vector<int64> dims_021; |
| 1538 | if (2 == segments.size()) { |
| 1539 | // The logical component-0 is of size one. |
| 1540 | dims_021 = {1, normalized_dims[1], normalized_dims[0]}; |
| 1541 | } else { |
| 1542 | dims_021 = {normalized_dims[0], normalized_dims[2], normalized_dims[1]}; |
| 1543 | } |
| 1544 | |
| 1545 | return dims_021; |
| 1546 | } |
| 1547 | |
| 1548 | return absl::nullopt; |
| 1549 | } |
| 1550 | |
| 1551 | } // namespace xla |
nothing calls this directly
no test coverage detected