#include "radix.h"
static void
_radix2_mask_bits(nn_ptr a, slong d, slong alimbs, slong e)
{
slong full = d / e, r = d % e, i;
if (r)
{
a[full] &= (UWORD(1) << r) - 1;
for (i = full + 1; i < alimbs; i++)
a[i] = 0;
}
else
{
for (i = full; i < alimbs; i++)
a[i] = 0;
}
}
static void
_radix_halve_odd(nn_ptr W, nn_srcptr V, slong k, const radix_t radix)
{
ulong B = LIMB_RADIX(radix);
ulong Bhalf = B >> 1;
ulong rem, par;
slong i;
par = 0;
for (i = 0; i < k; i++)
par ^= V[i];
par &= 1;
rem = par;
for (i = k - 1; i >= 0; i--)
{
ulong v = V[i];
if (rem)
{
W[i] = Bhalf + (v >> 1) + (v & 1);
rem = 1 - (v & 1);
}
else
{
W[i] = v >> 1;
rem = v & 1;
}
}
}
int
radix_sqrtmod_bn(nn_ptr res, nn_srcptr x, slong xn, slong n, const radix_t radix)
{
slong m, nm;
nn_ptr y;
TMP_INIT;
FLINT_ASSERT(xn >= 1);
FLINT_ASSERT(n >= 1);
m = (n + 1) / 2;
nm = n - m;
TMP_START;
if (DIGIT_RADIX(radix) == 2)
{
slong e = radix->exp;
slong mm = m + 1;
slong w = n + 1;
slong xc;
ulong bo;
nn_ptr b, b2, c, yc;
y = TMP_ALLOC(mm * sizeof(ulong));
if (!radix_rsqrtmod_bn(y, x, xn, mm, radix))
{
TMP_END;
return 0;
}
b = TMP_ALLOC(w * sizeof(ulong));
b2 = TMP_ALLOC(w * sizeof(ulong));
c = TMP_ALLOC(w * sizeof(ulong));
yc = TMP_ALLOC(w * sizeof(ulong));
radix_mulmid(b, x, FLINT_MIN(xn, mm), y, mm, 0, mm, radix);
flint_mpn_zero(b + mm, w - mm);
radix_mulmid(b2, b, mm, b, mm, 0, w, radix);
xc = FLINT_MIN(xn, w);
bo = 0;
if (xc > 0)
bo = radix_sub(c, x, xc, b2, xc, radix);
if (xc < w)
{
radix_neg(c + xc, b2 + xc, w - xc, radix);
if (bo)
radix_sub(c + xc, c + xc, w - xc, &bo, 1, radix);
}
radix_mulmid(yc, y, mm, c, w, 0, w, radix);
radix_rshift_digits(yc, yc, w, 1, radix);
radix_add(b, b, w, yc, w, radix);
_radix2_mask_bits(b, e * n, w, e);
flint_mpn_copyi(res, b, n);
TMP_END;
return 1;
}
else
{
y = TMP_ALLOC(m * sizeof(ulong));
if (!radix_rsqrtmod_bn(y, x, xn, m, radix))
{
TMP_END;
return 0;
}
radix_mulmid(res, x, FLINT_MIN(xn, m), y, m, 0, m, radix);
if (nm > 0)
{
nn_ptr b2h, dh, scr;
slong ah;
ulong bo;
b2h = TMP_ALLOC(nm * sizeof(ulong));
dh = TMP_ALLOC(nm * sizeof(ulong));
scr = TMP_ALLOC(n * sizeof(ulong));
_radix_mulhigh_known_low(b2h, res, m, res, m, x, FLINT_MIN(xn, m),
m, n, scr, radix);
ah = (xn > m) ? FLINT_MIN(xn - m, nm) : 0;
bo = 0;
if (ah > 0)
bo = radix_sub(dh, x + m, ah, b2h, ah, radix);
if (ah < nm)
{
radix_neg(dh + ah, b2h + ah, nm - ah, radix);
if (bo)
radix_sub(dh + ah, dh + ah, nm - ah, &bo, 1, radix);
}
radix_mulmid(res + m, y, m, dh, nm, 0, nm, radix);
_radix_halve_odd(res + m, res + m, nm, radix);
}
TMP_END;
return 1;
}
}