| 991 | } |
| 992 | |
| 993 | Status AlgebraicSimplifierVisitor::HandleSubtract(HloInstruction* sub) { |
| 994 | HloInstruction *lhs, *rhs; |
| 995 | CHECK(Match(sub, m::Subtract(m::Op(&lhs), m::Op(&rhs)))); |
| 996 | // A - 0 => A |
| 997 | VLOG(10) << "trying transform [A - 0 => A]: " << sub->ToString(); |
| 998 | if (IsAll(rhs, 0) && ReplaceInstructionIfSameShape(sub, lhs)) { |
| 999 | return Status::OK(); |
| 1000 | } |
| 1001 | |
| 1002 | // Canonicalize subtraction of a constant to addition. |
| 1003 | VLOG(10) << "trying transform [A - Const => A + (-Const)]"; |
| 1004 | if (Match(sub, m::Subtract(m::NonConstant(&lhs), m::Constant(&rhs))) || |
| 1005 | Match(sub, m::Subtract(m::NonConstant(&lhs), |
| 1006 | m::Broadcast(m::Constant(&rhs))))) { |
| 1007 | HloInstruction* negative_const = computation_->AddInstruction( |
| 1008 | HloInstruction::CreateUnary(rhs->shape(), HloOpcode::kNegate, rhs)); |
| 1009 | if (const HloInstruction* broadcast = |
| 1010 | DynCast<HloBroadcastInstruction>(sub->operand(1))) { |
| 1011 | negative_const = |
| 1012 | computation_->AddInstruction(HloInstruction::CreateBroadcast( |
| 1013 | broadcast->shape(), negative_const, broadcast->dimensions())); |
| 1014 | } |
| 1015 | return ReplaceWithNewInstruction( |
| 1016 | sub, HloInstruction::CreateBinary(sub->shape(), HloOpcode::kAdd, lhs, |
| 1017 | negative_const)); |
| 1018 | } |
| 1019 | |
| 1020 | return Status::OK(); |
| 1021 | } |
| 1022 | namespace { |
| 1023 | template <typename T> |
| 1024 | Status InvertConstant(const HloInstruction& constant, Literal* result) { |
no test coverage detected