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

Function IgammacContinuedFraction

tensorflow/compiler/xla/client/lib/math.cc:775–904  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

773// Helper function for computing Igammac using a continued fraction.
774template <kIgammaMode mode>
775XlaOp 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)));

Callers

nothing calls this directly

Calls 15

ScalarLikeFunction · 0.85
FullLikeFunction · 0.85
EpsilonFunction · 0.85
WhileLoopHelperFunction · 0.85
ReportErrorOrReturnMethod · 0.80
AnyFunction · 0.70
ReciprocalFunction · 0.70
DigammaFunction · 0.70
AndFunction · 0.50
LtFunction · 0.50
NeFunction · 0.50
SelectFunction · 0.50

Tested by

no test coverage detected