#include "radix.h"
#include "padic_radix.h"
#include "gr.h"
#include "longlong.h"
#include "mpn_extras.h"
int
padic_radix_dot_strided_naive(padic_radix_t res, const padic_radix_t initial,
int subtract, const padic_radix_struct * vec1, slong stride1,
const padic_radix_struct * vec2, slong stride2, slong len, gr_ctx_t ctx)
{
if (len <= 0)
{
if (initial == NULL)
return padic_radix_zero(res, ctx);
return padic_radix_set(res, initial, ctx);
}
int status = GR_SUCCESS;
padic_radix_t t;
padic_radix_init(t, ctx);
if (initial == NULL)
{
status |= padic_radix_mul(res, vec1, vec2, ctx);
}
else
{
if (subtract)
status |= padic_radix_neg(res, initial, ctx);
else
status |= padic_radix_set(res, initial, ctx);
status |= padic_radix_mul(t, vec1, vec2, ctx);
status |= padic_radix_add(res, res, t, ctx);
}
for (slong i = 1; i < len; i++)
{
status |= padic_radix_mul(t, vec1 + i * stride1, vec2 + i * stride2, ctx);
status |= padic_radix_add(res, res, t, ctx);
}
if (subtract)
status |= padic_radix_neg(res, res, ctx);
padic_radix_clear(t, ctx);
return status;
}
#define DOT_INF WORD_MAX
#define DOT_MIN2(a, b) ((a) < (b) ? (a) : (b))
#define DOT_MULMID_CUTOFF 140
static void
_radix_dot_normalize_acc(radix_integer_t T, nn_srcptr acc, slong nslots,
const radix_t radix)
{
nmod_t B = radix->B;
nn_ptr d = radix_integer_fit_limbs(T, nslots + 3, radix);
ulong cy[3] = { 0, 0, 0 };
slong s, outn = 0;
if (B.norm == 0)
{
for (s = 0; s < nslots; s++)
{
add_sssaaaaaa(cy[2], cy[1], cy[0], cy[2], cy[1], cy[0],
acc[3 * s + 2], acc[3 * s + 1], acc[3 * s]);
d[outn++] = flint_mpn_divrem_3_1_preinv_norm(cy, cy, B.n, B.ninv);
}
while (cy[0] | cy[1] | cy[2])
d[outn++] = flint_mpn_divrem_3_1_preinv_norm(cy, cy, B.n, B.ninv);
}
else
{
for (s = 0; s < nslots; s++)
{
add_sssaaaaaa(cy[2], cy[1], cy[0], cy[2], cy[1], cy[0],
acc[3 * s + 2], acc[3 * s + 1], acc[3 * s]);
d[outn++] = flint_mpn_divrem_3_1_preinv_unnorm(cy, cy, B.n, B.ninv, B.norm);
}
while (cy[0] | cy[1] | cy[2])
d[outn++] = flint_mpn_divrem_3_1_preinv_unnorm(cy, cy, B.n, B.ninv, B.norm);
}
while (outn > 0 && d[outn - 1] == 0)
outn--;
T->size = outn;
}
int
padic_radix_dot_strided_delayed(padic_radix_t res, const padic_radix_t initial,
int subtract, const padic_radix_struct * vec1, slong stride1,
const padic_radix_struct * vec2, slong stride2, slong len, gr_ctx_t ctx)
{
radix_struct * radix = PADIC_RADIX_CTX_RADIX(ctx);
slong e = radix->exp;
slong prec_abs = PADIC_RADIX_CTX_PREC_ABS(ctx);
slong prec_rel = PADIC_RADIX_CTX_PREC_REL(ctx);
slong k, vmin, maxhi, Nterms, N, Wdig, nslots;
int have_value;
const padic_radix_struct * a;
const padic_radix_struct * b;
int init_nonzero;
nn_ptr * accp;
nn_ptr * accn;
radix_integer_t Mp, Mn, T;
radix_integer_t big, Ptmp;
if (len <= 0)
{
if (initial == NULL)
return padic_radix_zero(res, ctx);
return padic_radix_set(res, initial, ctx);
}
vmin = DOT_INF;
maxhi = -DOT_INF;
Nterms = DOT_INF;
have_value = 0;
init_nonzero = (initial != NULL && initial->u.size != 0);
if (init_nonzero)
{
slong ni = FLINT_ABS(initial->u.size);
vmin = DOT_MIN2(vmin, initial->v);
maxhi = FLINT_MAX(maxhi, initial->v + ni * e);
have_value = 1;
}
if (initial != NULL && initial->N != PADIC_RADIX_EXACT)
Nterms = DOT_MIN2(Nterms, initial->N);
for (k = 0; k < len; k++)
{
slong va, vb, Na, Nb, vk, Nk;
a = vec1 + k * stride1;
b = vec2 + k * stride2;
if ((a->u.size == 0 && a->N == PADIC_RADIX_EXACT)
|| (b->u.size == 0 && b->N == PADIC_RADIX_EXACT))
continue;
va = a->v; vb = b->v; Na = a->N; Nb = b->N;
vk = va + vb;
Nk = PADIC_RADIX_EXACT;
if (Na != PADIC_RADIX_EXACT) Nk = DOT_MIN2(Nk, vb + Na);
if (Nb != PADIC_RADIX_EXACT) Nk = DOT_MIN2(Nk, va + Nb);
if (Nk != PADIC_RADIX_EXACT) Nterms = DOT_MIN2(Nterms, Nk);
if (a->u.size != 0 && b->u.size != 0)
{
slong na = FLINT_ABS(a->u.size), nb = FLINT_ABS(b->u.size);
vmin = DOT_MIN2(vmin, vk);
maxhi = FLINT_MAX(maxhi, vk + (na + nb) * e);
have_value = 1;
}
}
if (!have_value)
{
radix_integer_zero(&res->u, radix);
res->v = 0;
res->N = Nterms;
return _padic_radix_finalize(res, ctx);
}
N = Nterms;
N = DOT_MIN2(N, prec_abs);
if (prec_rel != PADIC_RADIX_PREC_INF)
N = DOT_MIN2(N, vmin + prec_rel);
if (N != DOT_INF && N <= vmin)
{
radix_integer_zero(&res->u, radix);
res->v = 0;
res->N = N;
return _padic_radix_finalize(res, ctx);
}
{
slong nslots_ext = (maxhi - vmin + e - 1) / e + 1;
if (N == DOT_INF)
nslots = nslots_ext;
else
{
slong nslots_prec = (N - vmin + e - 1) / e + 2;
nslots = DOT_MIN2(nslots_prec, nslots_ext);
}
if (nslots < 1)
nslots = 1;
}
Wdig = (N == DOT_INF) ? DOT_INF : (N - vmin);
accp = flint_calloc(e, sizeof(nn_ptr));
accn = flint_calloc(e, sizeof(nn_ptr));
radix_integer_init(big, radix);
radix_integer_init(Ptmp, radix);
if (init_nonzero)
{
slong ni = FLINT_ABS(initial->u.size);
nn_srcptr ui = initial->u.d;
slong d = initial->v - vmin;
slong r = d % e, o = d / e, i;
nn_ptr acc;
if (initial->u.size < 0)
{
if (accn[r] == NULL) accn[r] = flint_calloc(3 * nslots, sizeof(ulong));
acc = accn[r];
}
else
{
if (accp[r] == NULL) accp[r] = flint_calloc(3 * nslots, sizeof(ulong));
acc = accp[r];
}
for (i = 0; i < ni; i++)
{
slong slot = o + i;
if (slot >= nslots)
break;
acc[3 * slot] = ui[i];
}
}
for (k = 0; k < len; k++)
{
slong va, vb, vk, d, r, o, na, nb, m, mmax;
slong win, hi, an_t, bn_t;
nn_srcptr ua, ub;
nn_ptr acc;
int neg;
a = vec1 + k * stride1;
b = vec2 + k * stride2;
if (a->u.size == 0 || b->u.size == 0)
continue;
va = a->v; vb = b->v; vk = va + vb;
d = vk - vmin;
if (Wdig != DOT_INF && d >= Wdig)
continue;
r = d % e; o = d / e;
if (o >= nslots)
continue;
na = FLINT_ABS(a->u.size); nb = FLINT_ABS(b->u.size);
ua = a->u.d; ub = b->u.d;
neg = (a->u.size < 0) ^ (b->u.size < 0) ^ (subtract != 0);
win = nslots - o;
hi = FLINT_MIN(na + nb, win);
an_t = FLINT_MIN(na, hi);
bn_t = FLINT_MIN(nb, hi);
if (FLINT_MIN(an_t, bn_t) >= DOT_MULMID_CUTOFF && hi >= DOT_MULMID_CUTOFF)
{
nn_ptr pd = radix_integer_fit_limbs(Ptmp, hi, radix);
radix_mulmid(pd, ua, an_t, ub, bn_t, 0, hi, radix);
Ptmp->size = hi;
while (Ptmp->size > 0 && pd[Ptmp->size - 1] == 0)
Ptmp->size--;
if (neg)
radix_integer_sublsh(big, big, Ptmp, d, radix);
else
radix_integer_addlsh(big, big, Ptmp, d, radix);
continue;
}
if (neg)
{
if (accn[r] == NULL) accn[r] = flint_calloc(3 * nslots, sizeof(ulong));
acc = accn[r];
}
else
{
if (accp[r] == NULL) accp[r] = flint_calloc(3 * nslots, sizeof(ulong));
acc = accp[r];
}
mmax = na + nb - 2;
if (mmax > nslots - 1 - o)
mmax = nslots - 1 - o;
for (m = 0; m <= mmax; m++)
{
slong slot = o + m;
slong iilo = (m >= nb) ? (m - nb + 1) : 0;
slong iihi = (m < na) ? m : (na - 1);
slong ii;
ulong cy0 = acc[3 * slot];
ulong cy1 = acc[3 * slot + 1];
ulong cy2 = acc[3 * slot + 2];
for (ii = iilo; ii <= iihi; ii++)
{
ulong hilo, lo;
umul_ppmm(hilo, lo, ua[ii], ub[m - ii]);
add_sssaaaaaa(cy2, cy1, cy0, cy2, cy1, cy0, 0, hilo, lo);
}
acc[3 * slot] = cy0;
acc[3 * slot + 1] = cy1;
acc[3 * slot + 2] = cy2;
}
}
radix_integer_init(Mp, radix);
radix_integer_init(Mn, radix);
radix_integer_init(T, radix);
{
slong r;
for (r = 0; r < e; r++)
{
if (accp[r] != NULL)
{
_radix_dot_normalize_acc(T, accp[r], nslots, radix);
radix_integer_addlsh(Mp, Mp, T, r, radix);
flint_free(accp[r]);
}
if (accn[r] != NULL)
{
_radix_dot_normalize_acc(T, accn[r], nslots, radix);
radix_integer_addlsh(Mn, Mn, T, r, radix);
flint_free(accn[r]);
}
}
}
flint_free(accp);
flint_free(accn);
radix_integer_sub(&res->u, Mp, Mn, radix);
if (big->size != 0)
radix_integer_add(&res->u, &res->u, big, radix);
res->v = vmin;
res->N = N;
radix_integer_clear(Mp, radix);
radix_integer_clear(Mn, radix);
radix_integer_clear(T, radix);
radix_integer_clear(big, radix);
radix_integer_clear(Ptmp, radix);
return _padic_radix_finalize(res, ctx);
}
int
padic_radix_dot_strided(padic_radix_t res, const padic_radix_t initial,
int subtract, const padic_radix_struct * vec1, slong stride1,
const padic_radix_struct * vec2, slong stride2, slong len, gr_ctx_t ctx)
{
if (len <= 7)
return padic_radix_dot_strided_naive(res, initial, subtract, vec1, stride1, vec2, stride2, len, ctx);
else
return padic_radix_dot_strided_delayed(res, initial, subtract, vec1, stride1, vec2, stride2, len, ctx);
}
int
padic_radix_dot(padic_radix_t res, const padic_radix_t initial,
int subtract, const padic_radix_struct * vec1,
const padic_radix_struct * vec2, slong len, gr_ctx_t ctx)
{
return padic_radix_dot_strided(res, initial, subtract, vec1, 1, vec2, 1, len, ctx);
}
int
padic_radix_dot_rev(padic_radix_t res, const padic_radix_t initial,
int subtract, const padic_radix_struct * vec1,
const padic_radix_struct * vec2, slong len, gr_ctx_t ctx)
{
return padic_radix_dot_strided(res, initial, subtract, vec1, 1, vec2 + len - 1, -1, len, ctx);
}