#include "ulong_extras.h"
#include "padic_radix.h"
#include "arb/impl.h"
#include "gr.h"
#define PADIC_RADIX_LOG_BSPLIT_BASECASE 64
#define PADIC_RADIX_LOG_BC_LIMBS (2 * PADIC_RADIX_LOG_BSPLIT_BASECASE + 4)
#define PADIC_RADIX_LOG_SMALLEST_BLOCK 4
static void
_padic_radix_log_bsplit_basecase_mpn(radix_integer_t Tu, radix_integer_t B,
slong a, slong b, ulong xl, slong vr, slong l, const radix_t radix)
{
ulong T[PADIC_RADIX_LOG_BC_LIMBS];
ulong BB[PADIC_RADIX_LOG_BC_LIMBS];
ulong BXp[PADIC_RADIX_LOG_BC_LIMBS];
slong Tn, Bn, BXpn, J, need;
ulong hi, lo, cy;
nn_ptr d;
T[0] = xl; Tn = 1;
BB[0] = (ulong) a; Bn = 1;
umul_ppmm(hi, lo, (ulong) a, xl);
BXp[0] = lo;
BXp[1] = hi;
BXpn = 1 + (hi != 0);
for (J = a + 1; J < b; J++)
{
hi = mpn_mul_1(T, T, Tn, (ulong) J);
T[Tn] = hi; Tn += (hi != 0);
if (Tn < BXpn)
{
flint_mpn_zero(T + Tn, BXpn - Tn);
Tn = BXpn;
}
cy = mpn_addmul_1(T, BXp, BXpn, xl);
if (cy != 0)
{
if (Tn > BXpn)
cy = mpn_add_1(T + BXpn, T + BXpn, Tn - BXpn, cy);
if (cy != 0)
{
T[Tn] = cy; Tn++;
}
}
hi = mpn_mul_1(BB, BB, Bn, (ulong) J);
BB[Bn] = hi; Bn += (hi != 0);
if (J + 1 < b)
{
umul_ppmm(hi, lo, (ulong) J, xl);
if (hi == 0)
{
hi = mpn_mul_1(BXp, BXp, BXpn, lo);
BXp[BXpn] = hi; BXpn += (hi != 0);
}
else
{
hi = mpn_mul_1(BXp, BXp, BXpn, (ulong) J);
BXp[BXpn] = hi; BXpn += (hi != 0);
hi = mpn_mul_1(BXp, BXp, BXpn, xl);
BXp[BXpn] = hi; BXpn += (hi != 0);
}
}
}
need = radix_set_mpn_need_alloc(Tn, radix);
d = radix_integer_fit_limbs(Tu, need, radix);
Tu->size = radix_set_mpn(d, T, Tn, radix);
if (vr > 0)
radix_integer_rshift_digits(Tu, Tu, vr, radix);
radix_integer_mod_limbs(Tu, Tu, l, radix);
need = radix_set_mpn_need_alloc(Bn, radix);
d = radix_integer_fit_limbs(B, need, radix);
B->size = radix_set_mpn(d, BB, Bn, radix);
radix_integer_mod_limbs(B, B, l, radix);
}
static void
_padic_radix_log_bsplit(radix_integer_t Tu, radix_integer_t B,
slong a, slong b, const slong * xexp, const radix_integer_struct * xpow,
slong vr, slong l, ulong xl, int xl_fits, const radix_t radix)
{
slong e = radix->exp;
slong el = e * l;
if (xl_fits && (b - a) < PADIC_RADIX_LOG_BSPLIT_BASECASE)
{
_padic_radix_log_bsplit_basecase_mpn(Tu, B, a, b, xl, vr, l, radix);
}
else if (b - a == 1)
{
radix_integer_set(Tu, xpow + 0, radix);
radix_integer_set_ui(B, (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);
if (vr < el)
{
radix_integer_set_ui(tmp, (ulong) a, radix);
radix_integer_mullow_limbs(tmp, xpow + i2, tmp, l, radix);
radix_integer_lshift_digits(tmp, tmp, 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(B, (ulong) a, radix);
radix_integer_set_ui(tmp, (ulong) (a + 1), radix);
radix_integer_mullow_limbs(B, B, tmp, l, radix);
radix_integer_clear(tmp, radix);
}
else
{
slong step = (b - a) / 2;
slong m = a + step;
slong shift;
int include;
radix_integer_t TuR, BR, t2;
radix_integer_init(TuR, radix);
radix_integer_init(BR, radix);
radix_integer_init(t2, radix);
_padic_radix_log_bsplit(Tu, B, a, m, xexp, xpow, vr, l, xl, xl_fits, radix);
_padic_radix_log_bsplit(TuR, BR, m, b, xexp, xpow, vr, l, xl, xl_fits, radix);
radix_integer_mullow_limbs(Tu, Tu, BR, 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, B, l, radix);
radix_integer_mullow_limbs(t2, t2, 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(B, B, BR, l, radix);
radix_integer_clear(TuR, radix);
radix_integer_clear(BR, radix);
radix_integer_clear(t2, radix);
}
}
static void
_padic_radix_log_bsplit_block(radix_integer_t rop, const radix_integer_t x,
slong w, slong N, const radix_t radix)
{
ulong p = DIGIT_RADIX(radix);
slong e = radix->exp;
slong vr, Nr, n, k, l, lNr, kk, length, i;
slong * xexp;
radix_integer_struct * xpow;
radix_integer_t ru, Tu, B, ub;
ulong xl = 0, ru_l;
int xl_fits, use_table;
vr = radix_integer_valuation_digits(x, radix);
if (vr >= N)
{
radix_integer_zero(rop, radix);
return;
}
Nr = N - vr;
n = _padic_radix_log_bound(FLINT_MAX(w, vr), N, p);
n = FLINT_MAX(n, 2);
k = (n >= 2) ? (n - 2) / (slong) (p - 1) : 0;
lNr = (Nr + e - 1) / e;
l = lNr + (k + e - 1) / e;
radix_integer_init(ru, radix);
radix_integer_init(Tu, radix);
radix_integer_init(B, radix);
radix_integer_init(ub, radix);
radix_integer_rshift_digits(ru, x, vr, 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;
use_table = !(xl_fits && (n - 1) < PADIC_RADIX_LOG_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 log: power table malformed\n");
}
}
_padic_radix_log_bsplit(Tu, B, 1, n, xexp, xpow, vr, l, xl, xl_fits, radix);
kk = radix_integer_valuation_digits(B, radix);
if (kk > 0)
{
radix_integer_rshift_digits(B, B, kk, radix);
radix_integer_rshift_digits(Tu, Tu, kk, radix);
}
if (!radix_integer_divmod_limbs(ub, Tu, B, lNr, radix))
flint_throw(FLINT_ERROR, "_padic_radix_log_bsplit_block: series "
"denominator is not invertible after stripping\n");
radix_integer_mod_digits(ub, ub, Nr, radix);
radix_integer_lshift_digits(rop, ub, vr, radix);
radix_integer_mod_digits(rop, rop, 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(Tu, radix);
radix_integer_clear(B, radix);
radix_integer_clear(ub, radix);
}
void
_padic_radix_log_balanced(radix_integer_t rop, const radix_integer_t y, slong N,
const radix_t radix)
{
slong e = radix->exp;
slong lN, bound, w;
radix_integer_t t, r, om, blk, z;
if (N <= 0)
{
radix_integer_zero(rop, radix);
return;
}
lN = (N + e - 1) / e;
radix_integer_init(t, radix);
radix_integer_init(r, radix);
radix_integer_init(om, radix);
radix_integer_init(blk, radix);
radix_integer_init(z, radix);
radix_integer_mod_digits(t, y, N, radix);
radix_integer_zero(z, radix);
bound = PADIC_RADIX_LOG_SMALLEST_BLOCK;
w = 1;
while (!radix_integer_is_zero(t, radix))
{
radix_integer_mod_digits(r, t, bound, radix);
radix_integer_sub(t, t, r, radix);
if (!radix_integer_is_zero(t, radix))
{
radix_integer_one(om, radix);
radix_integer_sub(om, om, r, radix);
if (!radix_integer_divmod_limbs(t, t, om, lN, radix))
flint_throw(FLINT_ERROR, "_padic_radix_log_balanced: 1 - r is "
"not invertible\n");
radix_integer_mod_digits(t, t, N, radix);
}
if (!radix_integer_is_zero(r, radix))
{
_padic_radix_log_bsplit_block(blk, r, w, N, radix);
radix_integer_sub(z, z, blk, radix);
}
w = bound;
bound *= 2;
}
radix_integer_mod_digits(rop, z, N, radix);
radix_integer_clear(t, radix);
radix_integer_clear(r, radix);
radix_integer_clear(om, radix);
radix_integer_clear(blk, radix);
radix_integer_clear(z, radix);
}