| 698 | // Helper function for computing Igamma using a power series. |
| 699 | template <kIgammaMode mode> |
| 700 | XlaOp IgammaSeries(XlaOp ax, XlaOp x, XlaOp a, XlaOp enabled, |
| 701 | xla::PrimitiveType type) { |
| 702 | // vals: (enabled, r, c, ans, x) |
| 703 | // 'enabled' is a predication mask that says for which elements we should |
| 704 | // execute the loop body. Disabled elements have no effect in the loop body. |
| 705 | // TODO(phawkins): in general this isn't an optimal implementation on any |
| 706 | // backend. For example, on GPU, we should probably vectorize to the warp |
| 707 | // size, and then run independent loops for each warp's worth of |
| 708 | // data. |
| 709 | auto cond = [&](absl::Span<const XlaOp> vals, |
| 710 | XlaBuilder* builder) -> StatusOr<XlaOp> { |
| 711 | XlaOp enabled = vals[0]; |
| 712 | return Any(enabled); |
| 713 | }; |
| 714 | auto body = [&](absl::Span<const XlaOp> vals, |
| 715 | XlaBuilder* builder) -> StatusOr<std::vector<XlaOp>> { |
| 716 | XlaOp enabled = vals[0]; |
| 717 | XlaOp r = vals[1]; |
| 718 | XlaOp c = vals[2]; |
| 719 | XlaOp ans = vals[3]; |
| 720 | XlaOp x = vals[4]; |
| 721 | XlaOp dc_da = vals[5]; |
| 722 | XlaOp dans_da = vals[6]; |
| 723 | |
| 724 | r = r + ScalarLike(r, 1); |
| 725 | dc_da = dc_da * (x / r) + (ScalarLike(r, -1) * c * x) / (r * r); |
| 726 | dans_da = dans_da + dc_da; |
| 727 | c = c * (x / r); |
| 728 | ans = ans + c; |
| 729 | XlaOp conditional; |
| 730 | if (mode == VALUE) { |
| 731 | conditional = And(enabled, Gt(c / ans, Epsilon(builder, type))); |
| 732 | } else { |
| 733 | conditional = |
| 734 | And(enabled, Gt(Abs(dc_da / dans_da), Epsilon(builder, type))); |
| 735 | } |
| 736 | |
| 737 | return std::vector<XlaOp>{ |
| 738 | conditional, |
| 739 | Select(enabled, r, vals[1]), |
| 740 | Select(enabled, c, vals[2]), |
| 741 | Select(enabled, ans, vals[3]), |
| 742 | Select(enabled, x, vals[4]), |
| 743 | Select(enabled, dc_da, vals[5]), |
| 744 | Select(enabled, dans_da, vals[6]), |
| 745 | }; |
| 746 | }; |
| 747 | auto& b = *ax.builder(); |
| 748 | return b.ReportErrorOrReturn([&]() -> StatusOr<XlaOp> { |
| 749 | std::vector<XlaOp> vals = { |
| 750 | enabled, a, FullLike(a, 1), FullLike(a, 1), x, FullLike(a, 0), |
| 751 | FullLike(a, 0), |
| 752 | }; |
| 753 | |
| 754 | TF_ASSIGN_OR_RETURN(vals, WhileLoopHelper(cond, body, vals, "igamma", &b)); |
| 755 | XlaOp ans = vals[3]; |
| 756 | XlaOp dans_da = vals[6]; |
| 757 | if (mode == VALUE) { |
nothing calls this directly
no test coverage detected