| 508 | } |
| 509 | |
| 510 | XlaOp XlaBuilder::BinaryOp(HloOpcode binop, XlaOp lhs, XlaOp rhs, |
| 511 | absl::Span<const int64> broadcast_dimensions, |
| 512 | absl::optional<ComparisonDirection> direction) { |
| 513 | return ReportErrorOrReturn([&]() -> StatusOr<XlaOp> { |
| 514 | HloInstructionProto instr; |
| 515 | TF_ASSIGN_OR_RETURN(const Shape* lhs_shape, GetShapePtr(lhs)); |
| 516 | TF_ASSIGN_OR_RETURN(const Shape* rhs_shape, GetShapePtr(rhs)); |
| 517 | TF_ASSIGN_OR_RETURN( |
| 518 | Shape shape, ShapeInference::InferBinaryOpShape( |
| 519 | binop, *lhs_shape, *rhs_shape, broadcast_dimensions)); |
| 520 | *instr.mutable_shape() = shape.ToProto(); |
| 521 | if (binop == HloOpcode::kCompare) { |
| 522 | if (!direction.has_value()) { |
| 523 | return InvalidArgument( |
| 524 | "kCompare expects a ComparisonDirection, but none provided."); |
| 525 | } |
| 526 | instr.set_comparison_direction(ComparisonDirectionToString(*direction)); |
| 527 | } else if (direction.has_value()) { |
| 528 | return InvalidArgument( |
| 529 | "A comparison direction is provided for a non-compare opcode: %s.", |
| 530 | HloOpcodeString(binop)); |
| 531 | } |
| 532 | |
| 533 | const int64 lhs_rank = lhs_shape->rank(); |
| 534 | const int64 rhs_rank = rhs_shape->rank(); |
| 535 | |
| 536 | XlaOp updated_lhs = lhs; |
| 537 | XlaOp updated_rhs = rhs; |
| 538 | |
| 539 | if (!broadcast_dimensions.empty() && lhs_rank != rhs_rank) { |
| 540 | const bool should_broadcast_lhs = lhs_rank < rhs_rank; |
| 541 | XlaOp from = should_broadcast_lhs ? lhs : rhs; |
| 542 | const Shape& from_shape = should_broadcast_lhs ? *lhs_shape : *rhs_shape; |
| 543 | |
| 544 | std::vector<int64> to_size; |
| 545 | std::vector<bool> to_size_is_dynamic; |
| 546 | for (int i = 0; i < shape.rank(); i++) { |
| 547 | to_size.push_back(shape.dimensions(i)); |
| 548 | to_size_is_dynamic.push_back(shape.is_dynamic_dimension(i)); |
| 549 | } |
| 550 | for (int64 from_dim = 0; from_dim < from_shape.rank(); from_dim++) { |
| 551 | int64 to_dim = broadcast_dimensions[from_dim]; |
| 552 | to_size[to_dim] = from_shape.dimensions(from_dim); |
| 553 | to_size_is_dynamic[to_dim] = from_shape.is_dynamic_dimension(from_dim); |
| 554 | } |
| 555 | |
| 556 | const Shape& broadcasted_shape = ShapeUtil::MakeShape( |
| 557 | from_shape.element_type(), to_size, to_size_is_dynamic); |
| 558 | TF_ASSIGN_OR_RETURN( |
| 559 | XlaOp broadcasted_operand, |
| 560 | InDimBroadcast(broadcasted_shape, from, broadcast_dimensions)); |
| 561 | |
| 562 | updated_lhs = should_broadcast_lhs ? broadcasted_operand : lhs; |
| 563 | updated_rhs = !should_broadcast_lhs ? broadcasted_operand : rhs; |
| 564 | } |
| 565 | |
| 566 | TF_ASSIGN_OR_RETURN(const Shape* updated_lhs_shape, |
| 567 | GetShapePtr(updated_lhs)); |
no test coverage detected