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

Method BinaryOp

tensorflow/compiler/xla/client/xla_builder.cc:510–581  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

508}
509
510XlaOp 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));

Callers 15

CompareFunction · 0.45
ComplexFunction · 0.45
AddFunction · 0.45
SubFunction · 0.45
MulFunction · 0.45
DivFunction · 0.45
RemFunction · 0.45
MaxFunction · 0.45
MinFunction · 0.45
AndFunction · 0.45
OrFunction · 0.45
XorFunction · 0.45

Calls 14

InvalidArgumentFunction · 0.85
HloOpcodeStringFunction · 0.85
MakeShapeFunction · 0.85
is_dynamic_dimensionMethod · 0.80
TF_ASSIGN_OR_RETURNFunction · 0.50
mutable_shapeMethod · 0.45
ToProtoMethod · 0.45
has_valueMethod · 0.45
rankMethod · 0.45
emptyMethod · 0.45
push_backMethod · 0.45

Tested by

no test coverage detected