#include "pari.h"
#include "paripriv.h"
static GEN
psi(GEN c, ulong q, long prec)
{
GEN a = divru(c, q), ea = mpexp(a), invea = invr(ea);
GEN cha = shiftr(addrr(ea, invea), -1);
GEN sha = shiftr(subrr(ea, invea), -1);
return mulrr(sqrtr(stor(q,prec)), subrr(mulrr(a,cha), sha));
}
static GEN
L(GEN n, ulong k, long bitprec)
{
ulong r, l, m;
long pr = nbits2prec(bitprec / k + k);
GEN s = stor(0,pr), pi = mppi(pr);
pari_sp av = avma;
r = 2; m = umodiu(n,k);
for (l = 0; l < 2*k; l++)
{
if (m == 0)
{
GEN c = mpcos(divru(mulru(pi, 6*l+1), 6*k));
if (odd(l)) subrrz(s, c, s); else addrrz(s, c, s);
avma = av;
}
m += r; if (m >= k) m -= k;
r += 3; if (r >= k) r -= k;
}
return mulrr(s, sqrtr((k % 3)? rdivss(k,3,pr): utor(k/3,pr)));
}
static GEN
estim(GEN n)
{
pari_sp av = avma;
GEN p1, pi = mppi (DEFAULTPREC);
p1 = divru( itor(shifti(n,1), DEFAULTPREC), 3 );
p1 = mpexp( mulrr(pi, sqrtr(p1)) );
p1 = divri (shiftr(p1,-2), n);
p1 = divrr(p1, sqrtr( stor(3,DEFAULTPREC) ));
return gerepileupto(av, mplog(p1));
}
static void
pinit(GEN n, GEN *c, GEN *d, ulong prec)
{
GEN b = divru( itor( subiu(muliu(n,24), 1), prec ), 24 );
GEN sqrtb = sqrtr(b), Pi = mppi(prec), pi2sqrt2, pisqrt2d3;
pisqrt2d3 = mulrr(Pi, sqrtr( divru(stor(2, prec), 3) ));
pi2sqrt2 = mulrr(Pi, sqrtr( stor(8, prec) ));
*c = mulrr(pisqrt2d3, sqrtb);
*d = invr( mulrr(pi2sqrt2, mulrr(b,sqrtb)) );
}
GEN
numbpart(GEN n)
{
pari_sp ltop = avma, av;
GEN sum, est, C, D, p1, p2;
long prec, bitprec;
ulong q;
if (typ(n) != t_INT) pari_err_TYPE("partition function",n);
if (signe(n) < 0) return gen_0;
if (abscmpiu(n, 2) < 0) return gen_1;
if (cmpii(n, uu32toi(0x38d7e, 0xa4c68000)) >= 0)
pari_err_OVERFLOW("numbpart [n < 10^15]");
est = estim(n);
bitprec = (long)(rtodbl(est)/M_LN2) + 32;
prec = nbits2prec(bitprec);
pinit(n, &C, &D, prec);
sum = cgetr (prec); affsr(0, sum);
av = avma; togglesign(est);
for (q = (ulong)(sqrt(gtodouble(n))*0.24 + 5); q >= 3; q--, avma=av)
{
GEN t = L(n, q, bitprec);
if (abscmprr(t, mpexp(divru(est,q))) < 0) continue;
t = mulrr(t, psi(gprec_w(C, nbits2prec(bitprec / q + 32)), q, prec));
affrr(addrr(sum, t), sum);
}
p1 = addrr(sum, psi(C, 1, prec));
p2 = psi(C, 2, prec);
affrr(mod2(n)? subrr(p1,p2): addrr(p1,p2), sum);
return gerepileuptoint (ltop, roundr(mulrr(D,sum)));
}
static void
parse_interval(GEN a, long *amin, long *amax)
{
switch (typ(a))
{
case t_INT:
*amax = itos(a);
break;
case t_VEC:
if (lg(a) != 3)
pari_err_TYPE("forpart [expect vector of type [amin,amax]]",a);
*amin = gtos(gel(a,1));
*amax = gtos(gel(a,2));
if (*amin>*amax || *amin<0 || *amax<=0)
pari_err_TYPE("forpart [expect 0<=min<=max, 0<max]",a);
break;
default:
pari_err_TYPE("forpart",a);
}
}
void
forpart_init(forpart_t *T, long k, GEN abound, GEN nbound)
{
T->amin=1;
if (abound) parse_interval(abound,&T->amin,&T->amax);
else T->amax = k;
T->strip = (T->amin > 0) ? 1 : 0;
T->nmin=0;
if (nbound) parse_interval(nbound,&T->nmin,&T->nmax);
else T->nmax = k;
if ( T->amin*T->nmin > k || k > T->amax * T->nmax )
{
T->nmin = T->nmax = 0;
}
else
{
if ( T->nmin * T->amax < k )
T->nmin = 1 + (k - 1) / T->amax;
if (T->strip && T->nmax > k/T->amin)
T->nmax = k / T->amin;
if ( T->amax + (T->nmin-1)* T->amin > k )
T->amax = k - (T->nmin-1)* T->amin;
}
if ( T->amax < T->amin )
T->nmin = T->nmax = 0;
T->v = zero_zv(T->nmax);
T->k = k;
}
GEN
forpart_next(forpart_t *T)
{
GEN v = T->v;
long n = lg(v)-1;
long i, s, a, k, vi, vn;
if (n>0 && v[n])
{
s = a = v[n];
for(i = n-1; i>0 && v[i]+1 >= a; s += v[i--]);
if (i == 0) {
if ((n+1) * T->amin > s || n == T->nmax) return NULL;
i = 1; n++;
setlg(v, n+1);
vi = T->amin;
} else {
s += v[i];
vi = v[i]+1;
}
} else {
s = T->k;
if (T->amin == 0) T->amin = 1;
if (T->strip) { n = T->nmin; setlg(T->v, n+1); }
if (s==0)
{
if (n==0 && T->nmin==0) {T->nmin++; return v;}
return NULL;
}
if (n==0) return NULL;
vi = T->amin;
i = T->strip ? 1 : n + 1 - T->nmin;
if (s <= (n-i)*vi) return NULL;
}
vn = s - (n-i)*vi;
if (T->amax && vn > T->amax)
{
long ai, q, r;
vn -= vi;
ai = T->amax - vi;
q = vn / ai;
r = vn % ai;
while ( q-- ) v[n--] = T->amax;
if ( n >= i ) v[n--] = vi + r;
while ( n >= i ) v[n--] = vi;
} else {
for ( k=i; k<n; v[k++] = vi );
v[n] = vn;
}
return v;
}
GEN
forpart_prev(forpart_t *T)
{
GEN v = T->v;
long n = lg(v)-1;
long j, ni, q, r;
long i, s;
if (n>0 && v[n])
{
i = n-1; s = v[n];
while (i>1 && (v[i-1]==v[i] || v[i+1]==T->amax))
s+= v[i--];
if (!i) return NULL;
if ( v[i+1] == T->amax ) return NULL;
if (v[i] == T->amin) {
if (!T->strip) return NULL;
s += v[i]; v[i] = 0;
} else {
v[i]--; s++;
}
if (v[i] == 0)
{
if (T->nmin > n-i) return NULL;
if (T->strip) {
i = 0; n--;
setlg(v, n+1);
}
}
} else
{
s = T->k;
i = 0;
if (s==0)
{
if (n==0 && T->nmin==0) {T->nmin++; return v;}
return NULL;
}
if (n*T->amax < s || s < T->nmin*T->amin) return NULL;
}
ni = n-i;
q = s / ni;
r = s % ni;
for(j=i+1; j<=n-r; j++) v[j]=q;
for(j=n-r+1; j<=n; j++) v[j]=q + 1;
return v;
}
static long
countpart(long k, GEN abound, GEN nbound)
{
pari_sp av = avma;
long n;
forpart_t T;
if (k<0) return 0;
forpart_init(&T, k, abound, nbound);
for (n=0; forpart_next(&T); n++)
avma = av;
return n;
}
GEN
partitions(long k, GEN abound, GEN nbound)
{
GEN v;
forpart_t T;
long i, n = countpart(k,abound,nbound);
if (n==0) return cgetg(1, t_VEC);
forpart_init(&T, k, abound, nbound);
v = cgetg(n+1, t_VEC);
for (i=1; i<=n; i++)
gel(v,i)=zv_copy(forpart_next(&T));
return v;
}
void
forpart(void *E, long call(void*, GEN), long k, GEN abound, GEN nbound)
{
pari_sp av = avma;
GEN v;
forpart_t T;
forpart_init(&T, k, abound, nbound);
while ((v=forpart_next(&T)))
if (call(E, v)) break;
avma=av;
}
void
forpart0(GEN k, GEN code, GEN abound, GEN nbound)
{
pari_sp av = avma;
if (typ(k) != t_INT) pari_err_TYPE("forpart",k);
if (signe(k)<0) return;
push_lex(gen_0, code);
forpart((void*)code, &gp_evalvoid, itos(k), abound, nbound);
pop_lex(1);
avma=av;
}