| 773 | // Helper function for computing Igammac using a continued fraction. |
| 774 | template <kIgammaMode mode> |
| 775 | XlaOp IgammacContinuedFraction(XlaOp ax, XlaOp x, XlaOp a, XlaOp enabled, |
| 776 | xla::PrimitiveType type) { |
| 777 | // vals: enabled, ans, t, y, z, c, pkm1, qkm1, pkm2, qkm2 |
| 778 | auto cond = [&](absl::Span<const XlaOp> vals, |
| 779 | XlaBuilder* builder) -> StatusOr<XlaOp> { |
| 780 | XlaOp enabled = vals[0]; |
| 781 | XlaOp c = vals[5]; |
| 782 | return And(Lt(c, ScalarLike(c, 2000)), Any(enabled)); |
| 783 | }; |
| 784 | auto body = [&](absl::Span<const XlaOp> vals, |
| 785 | XlaBuilder* builder) -> StatusOr<std::vector<XlaOp>> { |
| 786 | XlaOp enabled = vals[0]; |
| 787 | XlaOp ans = vals[1]; |
| 788 | XlaOp t = vals[2]; |
| 789 | XlaOp y = vals[3]; |
| 790 | XlaOp z = vals[4]; |
| 791 | XlaOp c = vals[5]; |
| 792 | XlaOp pkm1 = vals[6]; |
| 793 | XlaOp qkm1 = vals[7]; |
| 794 | XlaOp pkm2 = vals[8]; |
| 795 | XlaOp qkm2 = vals[9]; |
| 796 | |
| 797 | XlaOp dpkm2_da = vals[10]; |
| 798 | XlaOp dqkm2_da = vals[11]; |
| 799 | XlaOp dpkm1_da = vals[12]; |
| 800 | XlaOp dqkm1_da = vals[13]; |
| 801 | XlaOp dans_da = vals[14]; |
| 802 | |
| 803 | c = c + ScalarLike(c, 1); |
| 804 | y = y + ScalarLike(y, 1); |
| 805 | z = z + ScalarLike(z, 2); |
| 806 | XlaOp yc = y * c; |
| 807 | XlaOp pk = pkm1 * z - pkm2 * yc; |
| 808 | XlaOp qk = qkm1 * z - qkm2 * yc; |
| 809 | XlaOp qk_is_nonzero = Ne(qk, ScalarLike(qk, 0)); |
| 810 | XlaOp r = pk / qk; |
| 811 | |
| 812 | t = Select(qk_is_nonzero, Abs((ans - r) / r), FullLike(t, 1)); |
| 813 | ans = Select(qk_is_nonzero, r, ans); |
| 814 | |
| 815 | XlaOp dpk_da = dpkm1_da * z - pkm1 - dpkm2_da * yc + pkm2 * c; |
| 816 | XlaOp dqk_da = dqkm1_da * z - qkm1 - dqkm2_da * yc + qkm2 * c; |
| 817 | XlaOp dans_da_new = |
| 818 | Select(qk_is_nonzero, (dpk_da - ans * dqk_da) / qk, dans_da); |
| 819 | XlaOp grad_conditional = |
| 820 | Select(qk_is_nonzero, Abs(dans_da_new - dans_da), FullLike(dans_da, 1)); |
| 821 | |
| 822 | pkm2 = pkm1; |
| 823 | pkm1 = pk; |
| 824 | qkm2 = qkm1; |
| 825 | qkm1 = qk; |
| 826 | |
| 827 | dpkm2_da = dpkm1_da; |
| 828 | dqkm2_da = dqkm1_da; |
| 829 | dpkm1_da = dpk_da; |
| 830 | dqkm1_da = dqk_da; |
| 831 | |
| 832 | XlaOp rescale = Gt(Abs(pk), Reciprocal(Epsilon(builder, type))); |
nothing calls this directly
no test coverage detected