#include "ulong_extras.h"
#include "fmpz.h"
#include "arb/impl.h"
#include "padic_radix.h"
#include "gr.h"
static void
_radix_integer_slice_digits(radix_integer_t res, const radix_integer_t x,
slong a, slong b, const radix_t radix)
{
radix_integer_mod_digits(res, x, b, radix);
if (a > 0)
radix_integer_rshift_digits(res, res, a, radix);
}
#define PADIC_RADIX_EXP_BSPLIT_BASECASE 64
#define PADIC_RADIX_EXP_BC_LIMBS (2 * PADIC_RADIX_EXP_BSPLIT_BASECASE + 4)
#define PADIC_RADIX_EXP_SMALLEST_BLOCK 4
static void
_bsplit_basecase_mpn(radix_integer_t Tu, radix_integer_t Q, slong a, slong b,
ulong xl, ulong ru_l, slong l, const radix_t radix)
{
ulong U[PADIC_RADIX_EXP_BC_LIMBS];
ulong QQ[PADIC_RADIX_EXP_BC_LIMBS];
ulong Xp[PADIC_RADIX_EXP_BC_LIMBS];
slong Un, Qn, Xpn, J, need;
ulong hi, cy;
nn_ptr d;
U[0] = 1; Un = 1;
QQ[0] = (ulong) a; Qn = 1;
Xp[0] = xl; Xpn = 1;
for (J = a + 1; J < b; J++)
{
hi = mpn_mul_1(U, U, Un, (ulong) J);
U[Un] = hi; Un += (hi != 0);
if (Un >= Xpn)
cy = mpn_add(U, U, Un, Xp, Xpn);
else
{
cy = mpn_add(U, Xp, Xpn, U, Un);
Un = Xpn;
}
U[Un] = cy; Un += (cy != 0);
hi = mpn_mul_1(QQ, QQ, Qn, (ulong) J);
QQ[Qn] = hi; Qn += (hi != 0);
if (J + 1 < b)
{
hi = mpn_mul_1(Xp, Xp, Xpn, xl);
Xp[Xpn] = hi; Xpn += (hi != 0);
}
}
hi = mpn_mul_1(U, U, Un, ru_l);
U[Un] = hi; Un += (hi != 0);
need = radix_set_mpn_need_alloc(Un, radix);
d = radix_integer_fit_limbs(Tu, need, radix);
Tu->size = radix_set_mpn(d, U, Un, radix);
radix_integer_mod_limbs(Tu, Tu, l, radix);
need = radix_set_mpn_need_alloc(Qn, radix);
d = radix_integer_fit_limbs(Q, need, radix);
Q->size = radix_set_mpn(d, QQ, Qn, radix);
radix_integer_mod_limbs(Q, Q, l, radix);
}
static void
_bsplit(radix_integer_t Tu, radix_integer_t Q, slong a, slong b,
const slong * xexp, const radix_integer_struct * xpow, slong vr, slong l,
ulong xl, ulong ru_l, int xl_fits, const radix_t radix)
{
slong e = radix->exp;
if (xl_fits && (b - a) < PADIC_RADIX_EXP_BSPLIT_BASECASE)
{
_bsplit_basecase_mpn(Tu, Q, a, b, xl, ru_l, l, radix);
}
else if (b - a == 1)
{
radix_integer_set(Tu, xpow + 0, radix);
radix_integer_set_ui(Q, (ulong) a, radix);
}
else if (b - a == 2)
{
slong i2 = _arb_get_exp_pos(xexp, 2);
radix_integer_t tmp;
radix_integer_init(tmp, radix);
radix_integer_set_ui(tmp, (ulong) (a + 1), radix);
radix_integer_mullow_limbs(Tu, xpow + 0, tmp, l, radix);
radix_integer_lshift_digits(tmp, xpow + i2, vr, radix);
radix_integer_mod_limbs(tmp, tmp, l, radix);
radix_integer_add(Tu, Tu, tmp, radix);
radix_integer_mod_limbs(Tu, Tu, l, radix);
radix_integer_set_ui(Q, (ulong) a, radix);
radix_integer_set_ui(tmp, (ulong) (a + 1), radix);
radix_integer_mullow_limbs(Q, Q, tmp, l, radix);
radix_integer_clear(tmp, radix);
}
else
{
slong step = (b - a) / 2;
slong m = a + step;
slong el = e * l;
slong shift;
int include;
radix_integer_t TuR, QR, t2;
radix_integer_init(TuR, radix);
radix_integer_init(QR, radix);
radix_integer_init(t2, radix);
_bsplit(Tu, Q, a, m, xexp, xpow, vr, l, xl, ru_l, xl_fits, radix);
_bsplit(TuR, QR, m, b, xexp, xpow, vr, l, xl, ru_l, xl_fits, radix);
radix_integer_mullow_limbs(Tu, Tu, QR, l, radix);
if (step != 0 && vr > el / step)
include = 0;
else
{
shift = vr * step;
include = (shift < el);
}
if (include)
{
slong is = _arb_get_exp_pos(xexp, step);
radix_integer_mullow_limbs(t2, xpow + is, TuR, l, radix);
radix_integer_lshift_digits(t2, t2, shift, radix);
radix_integer_mod_limbs(t2, t2, l, radix);
radix_integer_add(Tu, Tu, t2, radix);
radix_integer_mod_limbs(Tu, Tu, l, radix);
}
radix_integer_mullow_limbs(Q, Q, QR, l, radix);
radix_integer_clear(TuR, radix);
radix_integer_clear(QR, radix);
radix_integer_clear(t2, radix);
}
}
static void
_bsplit_block_numden(radix_integer_t pnum, radix_integer_t pden,
const radix_integer_t s, slong voff, slong N, const radix_t radix)
{
ulong p = DIGIT_RADIX(radix);
slong e = radix->exp;
slong vs, vr, n, k, l, w, length, i;
slong * xexp;
radix_integer_struct * xpow;
radix_integer_t ru, Q, Tu;
ulong xl = 0, ru_l = 0;
int xl_fits, use_table;
vs = radix_integer_valuation_digits(s, radix);
vr = voff + vs;
if (vr >= N)
{
radix_integer_one(pnum, radix);
radix_integer_one(pden, radix);
return;
}
n = _padic_radix_exp_bound(vr, N, p);
if (n == 1)
{
radix_integer_one(pnum, radix);
radix_integer_one(pden, radix);
return;
}
k = (n >= 2) ? (n - 2) / (slong) (p - 1) : 0;
l = (N + k + e - 1) / e;
radix_integer_init(ru, radix);
radix_integer_init(Q, radix);
radix_integer_init(Tu, radix);
radix_integer_rshift_digits(ru, s, vs, radix);
radix_integer_mod_limbs(ru, ru, l, radix);
ru_l = (ru->size == 0) ? 0 : ru->d[0];
xl_fits = (FLINT_ABS(ru->size) <= 1) && (vr < e) && (ru_l < radix->bpow[e - vr]);
if (xl_fits)
xl = radix->bpow[vr] * ru_l;
else
ru_l = 0;
use_table = !(xl_fits && (n - 1) < PADIC_RADIX_EXP_BSPLIT_BASECASE);
xexp = NULL;
xpow = NULL;
length = 0;
if (use_table)
{
xexp = flint_calloc(2 * FLINT_BITS, sizeof(slong));
length = _arb_compute_bs_exponents(xexp, n - 1);
xpow = flint_malloc(sizeof(radix_integer_struct) * length);
for (i = 0; i < length; i++)
radix_integer_init(xpow + i, radix);
radix_integer_set(xpow + 0, ru, radix);
for (i = 1; i < length; i++)
{
if (xexp[i] == 2 * xexp[i - 1])
radix_integer_mullow_limbs(xpow + i, xpow + i - 1, xpow + i - 1, l, radix);
else if (xexp[i] == 2 * xexp[i - 2])
radix_integer_mullow_limbs(xpow + i, xpow + i - 2, xpow + i - 2, l, radix);
else if (xexp[i] == 2 * xexp[i - 1] + 1)
{
radix_integer_mullow_limbs(xpow + i, xpow + i - 1, xpow + i - 1, l, radix);
radix_integer_mullow_limbs(xpow + i, xpow + i, ru, l, radix);
}
else if (xexp[i] == 2 * xexp[i - 2] + 1)
{
radix_integer_mullow_limbs(xpow + i, xpow + i - 2, xpow + i - 2, l, radix);
radix_integer_mullow_limbs(xpow + i, xpow + i, ru, l, radix);
}
else
flint_throw(FLINT_ERROR, "padic_radix exp: power table malformed\n");
}
}
_bsplit(Tu, Q, 1, n, xexp, xpow, vr, l, xl, ru_l, xl_fits, radix);
w = radix_integer_valuation_digits(Tu, radix);
if (w > 0)
{
radix_integer_rshift_digits(Tu, Tu, w, radix);
radix_integer_rshift_digits(Q, Q, radix_integer_valuation_digits(Q, radix), radix);
}
radix_integer_mod_digits(pden, Q, N, radix);
radix_integer_lshift_digits(pnum, Tu, vr, radix);
radix_integer_add(pnum, pnum, Q, radix);
radix_integer_mod_digits(pnum, pnum, N, radix);
if (use_table)
{
for (i = 0; i < length; i++)
radix_integer_clear(xpow + i, radix);
flint_free(xpow);
flint_free(xexp);
}
radix_integer_clear(ru, radix);
radix_integer_clear(Q, radix);
radix_integer_clear(Tu, radix);
}
void
_padic_radix_exp_balanced(radix_integer_t rop, const radix_integer_t u,
slong v, slong N, const radix_t radix)
{
const slong S = PADIC_RADIX_EXP_SMALLEST_BLOCK;
ulong p = DIGIT_RADIX(radix);
slong e = radix->exp;
slong lN, D, jmax, j, tl, td;
ulong top;
radix_integer_t t, s, Pnum, Pden, pnum, pden;
if (N <= 0)
{
radix_integer_zero(rop, radix);
return;
}
if (v >= N || radix_integer_is_zero(u, radix))
{
radix_integer_one(rop, radix);
return;
}
lN = (N + e - 1) / e;
radix_integer_init(t, radix);
radix_integer_lshift_digits(t, u, v, radix);
radix_integer_mod_digits(t, t, N, radix);
if (t->size == 0)
{
radix_integer_one(rop, radix);
radix_integer_clear(t, radix);
return;
}
tl = FLINT_ABS(t->size);
top = t->d[tl - 1];
td = 0;
while (top != 0)
{
top /= p;
td++;
}
D = (tl - 1) * e + td;
jmax = 1;
while ((S << (jmax - 1)) < D)
jmax++;
radix_integer_init(s, radix);
radix_integer_init(Pnum, radix);
radix_integer_init(Pden, radix);
radix_integer_init(pnum, radix);
radix_integer_init(pden, radix);
radix_integer_one(Pnum, radix);
radix_integer_one(Pden, radix);
for (j = jmax; j >= 1; j--)
{
slong a = (j == 1) ? 0 : (S << (j - 2));
slong b = S << (j - 1);
_radix_integer_slice_digits(s, t, a, b, radix);
if (s->size != 0)
{
_bsplit_block_numden(pnum, pden, s, a, N, radix);
radix_integer_mullow_limbs(Pnum, Pnum, pnum, lN, radix);
radix_integer_mullow_limbs(Pden, Pden, pden, lN, radix);
}
}
radix_integer_mod_digits(Pnum, Pnum, N, radix);
radix_integer_mod_digits(Pden, Pden, N, radix);
if (!radix_integer_divmod_limbs(rop, Pnum, Pden, lN, radix))
flint_throw(FLINT_ERROR, "_padic_radix_exp_balanced: denominator "
"product is not invertible\n");
radix_integer_mod_digits(rop, rop, N, radix);
radix_integer_clear(t, radix);
radix_integer_clear(s, radix);
radix_integer_clear(Pnum, radix);
radix_integer_clear(Pden, radix);
radix_integer_clear(pnum, radix);
radix_integer_clear(pden, radix);
}