#include "radix.h"
FLINT_FORCE_INLINE mp_limb_t
_radix_divrem_2_1_unnorm1(mp_ptr qp, mp_srcptr up, mp_limb_t d, mp_limb_t dinv, unsigned int norm)
{
mp_limb_t u0, u1, r;
FLINT_ASSERT(norm >= 1);
u1 = up[1];
u0 = up[0];
if (u1 < d)
{
d <<= norm;
qp[1] = 0;
r = (u1 << norm) | (u0 >> (FLINT_BITS - norm));
}
else
{
d <<= norm;
r = (u1 >> (FLINT_BITS - norm));
udiv_qrnnd_preinv(qp[1], r, r, (u1 << norm) | (u0 >> (FLINT_BITS - norm)), d, dinv);
}
udiv_qrnnd_preinv(qp[0], r, r, u0 << norm, d, dinv);
return r >> norm;
}
int
radix_divmod_bn_1(nn_ptr q, nn_ptr rem, nn_srcptr a, slong an, ulong b,
slong n, const radix_t radix)
{
nmod_t mod = radix->B;
ulong binv, bb;
slong i;
slong sbits;
FLINT_ASSERT(an >= 1);
FLINT_ASSERT(n >= 1);
FLINT_ASSERT(b < mod.n);
bb = b;
if (!radix_invmod_bn(&binv, &bb, 1, 1, radix))
return 0;
sbits = 2 * NMOD_BITS(mod);
if (sbits <= FLINT_BITS)
{
ulong cy = 0;
FLINT_ASSERT(mod.norm != 0);
for (i = 0; i < n; i++)
{
ulong ai = (i < an) ? a[i] : 0;
ulong w = nmod_sub(ai, cy, mod);
ulong qi = nmod_mul(w, binv, mod);
q[i] = qi;
cy += qi * b;
(void) n_divrem_preinv_unnorm(&cy, cy, mod.n, mod.ninv, mod.norm);
}
FLINT_ASSERT(cy < mod.n);
if (rem != NULL)
rem[0] = nmod_sub((n < an) ? a[n] : 0, cy, mod);
}
else if (mod.norm == 0)
{
ulong hi, lo, r;
ulong cy0 = 0, cy1 = 0;
for (i = 0; i < n; i++)
{
ulong ai = (i < an) ? a[i] : 0;
ulong w = nmod_sub(ai, cy0, mod);
ulong qi = nmod_mul(w, binv, mod);
q[i] = qi;
umul_ppmm(hi, lo, qi, b);
add_ssaaaa(cy1, cy0, cy1, cy0, hi, lo);
r = n_divrem_norm(&cy1, cy1, mod.n);
udiv_qrnnd_preinv(cy0, r, r, cy0, mod.n, mod.ninv);
(void) r;
}
FLINT_ASSERT(cy0 < mod.n);
FLINT_ASSERT(cy1 == 0);
if (rem != NULL)
rem[0] = nmod_sub((n < an) ? a[n] : 0, cy0, mod);
}
else
{
ulong hi, lo;
ulong cy[2] = { 0, 0 };
for (i = 0; i < n; i++)
{
ulong ai = (i < an) ? a[i] : 0;
ulong w = nmod_sub(ai, cy[0], mod);
ulong qi = nmod_mul(w, binv, mod);
q[i] = qi;
umul_ppmm(hi, lo, qi, b);
add_ssaaaa(cy[1], cy[0], cy[1], cy[0], hi, lo);
(void) _radix_divrem_2_1_unnorm1(cy, cy, mod.n, mod.ninv, mod.norm);
}
FLINT_ASSERT(cy[0] < mod.n);
FLINT_ASSERT(cy[1] == 0);
if (rem != NULL)
rem[0] = nmod_sub((n < an) ? a[n] : 0, cy[0], mod);
}
return 1;
}