题解 P2293【高精度开根】

· · 个人记录

我的算法之前题解似乎都没用过。

前言:本算法适用于想要稍微了解比朴素高深一点的算法的朋友们,需要的数学知识并不多。

法一:二分+高精+快速幂。

复杂度 O(\log n\log m\log ^2n)

(这个我不说,想知道可以翻其他题解)。

法二:递推。

我们发现,在法一中,快速幂每次都要处理所有的数,也就是 O(\log ^2n)

我们发现这东西太慢了!!!

而且我们可以发现:在二分中,每次计算的结果只用来比个大小,对后来的计算并没有任何帮助。

于是考虑换一种思路,破除二分的局限,按位常试。

由于知道原数,也知道开多少次根,那我们就能知道结果的位数。

举个栗子,\sqrt[3]{239847529},那一定是个 3 位数。

规律:设 n 为原数位数, m 为根次,则结果一定是 \lceil \frac{n} {m}\rceil

也就是:int w=(n+(m-1))/m;

然后,从高到低一位一位试。

还拿上面的栗子,\sqrt[3]{239847529},从百位开始试,发现 6^3=216,7^3=343,216<239<343,所以,百位是 6。

第一位很好做,第二位也只需要处理两位,并不慢,但位数多了呢?

接下来怎么实现——牛顿二项式定理:

此过程的重点为如何将前面算出的结果与填上一位后的结果联系起来。

假设开 m 次根,已经试出来的高精度结果为 a ,现在正在试的那位试一个数字 b ,我们试出的结果就会是:\sum \limits _{i=0}^m C_m^i \times a^{m-i} \times b^i

于是神奇的一件事发生了:我们可以通过前一位推出填上后的结果了!

具体一点:

其中组合数 C 预处理成高精,b 跟着循环一次一次乘,都很简单。

但是确定出一位后还要算 a^i,这怎么办?

我们在算某一位试数的结果时,发现无论填什么,前面结果不会变(废话)。

考虑预处理出前一位的 a^i

现在,让我们看看这两位的对话:

上一位:“我已经试出来正确结果啦!”

下一位:“真不错,但是我需要你的结果的 1~m 次方,你能求一下吗?谢谢!”

上一位:“我想想……诶,我有主意了,我在尝试结果时计算 m 次方用到了牛顿二项式定理的公式,那计算我的多少次方,也可以这么算!”

于是,上一位开心地在求出正确结果后迭代地求出了它的 1~m 次方。

假设填的是 y , 上一位之前所有数拼接起来结果是 x

对于每个 i 次方,具体来说上一位求的就是:\sum\limits _{j=0}\limits ^i C_i^j\times x^{i-j}\times y^j

同理,y 随着循环一次一次乘;不同的是,x 已经有预处理好的结果摆着了。

时间复杂度分析

经“简单”计算,C 数组平均不超过 10 位,b^i 显然平均也只有 20 位左右,所以我们发现,如果用高精乘它们,常数应该不大。

回顾一下,一位都干了哪些事:

1.从低到高一位一位尝试,并计算结果:O(10 \, m\log n)

2.处理出自己的1~m次方:O(m^2\log n)

由于有 \log n 位,此算法时间复杂度为 O(m^2\log^2 n),比直接二分的 O(\log^2 n \log n\log m),还是要快好多倍。

关于压位

我们知道,压位高精可以大幅度提高高精度的效率。我们设压 d 位:

乍一看复杂度是 O(\frac {m^2\log^2 n}{d^2}),但是会遇到问题。

还记得之前一位是怎么确定的吗?需要一点一点尝试。但 既然压了 d 位,那尝试的时候肯定要 dd 位试,就要算 10^d 次,明显地浪费时间!

还是拿我的栗子:\sqrt[3]{239847529}6^3=216 是答案,那么大于 6 算出的结果就一定偏大,反之偏小。

所以它符合单调性,可以 d 位内二分。

因为 dd 位地试,所以只需要 \frac{\log n}{d} 次,总的复杂度为:O(\frac {(m+d)\log^2 n}{d^3}),但因为这样做会增大常数,所以和 O(\frac {m^2\log^2 n}{d^2}) 复杂度其实是同一级别的。

这里我选择压 4 位,即 d=4 。(下面我将把 d 换成 4 )。

此时复杂度为 O(\frac {m^2\log^2 n}{16})

算法比较:

1.二分+快速幂+压位: O(\frac{\log^2 n \log n\log m}{16}) ,简单好写,速度慢。

2.逐位递推(本题解):约 O(\frac {m^2\log^2 n}{16}) ,好理解,比较难写,速度稍快。

3.迭代:不会,不知道

放上代码:

#include<stdio.h>
#include<iostream>
using namespace std;
struct hi{int a[10005],w;}a,t1,ans,C[52][52],fang[52],fang1[52];
int m;
string s;
hi chan(int x)
{
    hi a;
    a.w=0;
    while(x)
        a.a[++a.w]=x%10000,
        x/=10000;
    return a;
}
void print(hi a)
{
    if(a.w==0) {printf("0\n");return;}
    printf("%d",a.a[a.w]);
    for(int i=a.w-1;i>=1;i--)
    {
        if(a.a[i]>=1000) printf("%d",a.a[i]);
        else if(a.a[i]>=100) printf("0%d",a.a[i]);
        else if(a.a[i]>=10) printf("00%d",a.a[i]);
        else if(a.a[i]>=1) printf("000%d",a.a[i]);
        else printf("0000");
    }
    printf("\n");
} 
bool cmp(hi a,hi b)
{
    if(a.w>b.w) return 0;
    if(b.w>a.w) return 1;
    for(int i=a.w;i>=1;i--)
    {
        if(a.a[i]>b.a[i]) return 0;
        if(b.a[i]>a.a[i]) return 1;
    }
    return 1;
}
hi pl(hi a,hi b)
{
    hi c;
    c.w=max(a.w,b.w);
    c.a[c.w+1]=0;
    if(a.w>b.w) for(int i=b.w+1;i<=a.w;i++) b.a[i]=0;
    if(a.w<b.w) for(int i=a.w+1;i<=b.w;i++) a.a[i]=0;
    for(int i=1;i<=c.w;i++) c.a[i]=a.a[i]+b.a[i];
    for(int i=1;i<=c.w;i++) c.a[i+1]+=c.a[i]/10000,c.a[i]%=10000;
    if(c.a[c.w+1]>0) c.w++;
    return c;
}
hi ti(hi a,int b)
{
    for(int i=a.w+1;i<=a.w+7;i++) a.a[i]=0;
    for(int i=1;i<=a.w;i++) a.a[i]*=b;
    for(int i=1;i<=a.w+5;i++) a.a[i+1]+=a.a[i]/10000,a.a[i]%=10000;
    while(a.a[a.w+1]>0) a.w++;
    return a;
}
hi times(hi a,hi b)
{
    hi c;
    c.w=a.w+b.w-1;
    for(int i=1;i<=c.w+10;i++) c.a[i]=0;
    if(a.w==0||b.w==0) {c.w=0;return c;}
    for(int i=1;i<=a.w;i++)
        for(int j=1;j<=b.w;j++)
            c.a[i+j-1]+=a.a[i]*b.a[j];
    for(int i=1;i<=c.w+9;i++) c.a[i+1]+=c.a[i]/10000,c.a[i]%=10000;
    while(c.a[c.w+1]>0) c.w++;
    return c;
}
void init()
{
    for(int i=0;i<=m;i++)
        for(int j=0;j<=m;j++) C[i][j].w=0;
    C[0][0]=chan(1);
    for(int i=1;i<=m+1;i++)
        for(int j=0;j<=i-1;j++) C[i][j]=pl(C[i-1][j-1],C[i-1][j]);
    return;
}
hi f(int b,int m)//(a+b)^m
{
    hi res;
    hi t1;
    res=chan(0);
    t1=chan(1);
    //for(int i=1;i<=10;i++) res.a[i]=0;
    for(int i=m;i>=0;i--) res=pl(res,times(times(C[m+1][m-i],fang[i]),t1)),t1=ti(t1,b);
    return res;
}
int main()
{
    scanf("%d",&m);
    init();
    cin>>s;
    int L=s.length();
    while(s[L-1]==' '||s[L-1]=='EOF') L--;
    //L--;
    a.w=(L+3)/4;
    for(int i=1;i<=a.w-1;i++)
        a.a[i]=s[(L-i*4)]*1000+s[(L-i*4)+1]*100+s[(L-i*4)+2]*10+s[(L-i*4)+3]-'0'*1111;
    if(a.w*4-L==3) a.a[a.w]=s[0]-'0';
    if(a.w*4-L==2) a.a[a.w]=s[0]*10+s[1]-'0'*11;
    if(a.w*4-L==1) a.a[a.w]=s[0]*100+s[1]*10+s[2]-'0'*111;
    if(a.w*4-L==0) a.a[a.w]=s[0]*1000+s[1]*100+s[2]*10+s[3]-'0'*1111;
    ans.w=((a.w*4+m-1)/m+3)/4;
//    printf("%d %d\n",ans.w,L);
//    print(a);
//  for(int i=0;i<=m;i++) {for(int j=0;j<=i;j++) print(C[i][j]);printf("\n");}
    for(int i=ans.w;i>=1;i--)
    {
        t1.w=0;
        for(int j=(i-1)*m+1;j<=a.w;j++) t1.a[++t1.w]=a.a[j];
        //print(t1);
        //for(int i=1;i<=m;i++) fang[i]=chan(0);
        int l=0,r=9999;
        //fang[m]=chan(0);
        fang[0]=chan(1);
        while(l<r)
        {
            fang1[m]=chan(0);
            fang1[0]=chan(1);
            int mid=(l+r+1)/2;
            fang1[m]=f(mid,m);
            int t=cmp(fang1[m],t1);
            //printf("::::::::::::::::::::::::::::::::::%d %d %d %d:",i,mid,fang1[m].w,t);
            //print(fang1[m]);
            if(t) l=mid;
            else r=mid-1;
        }
        for(int k=1;k<=m;k++)
            fang1[k]=f(r,k);
        for(int k=1;k<=m;k++)
        {
            int t2=fang1[k].w+k;
            for(int l=t2;l>=t2-fang1[k].w;l--) fang1[k].a[l]=fang1[k].a[fang1[k].w-t2+l];
            for(int l=1;l<=t2-fang1[k].w;l++) fang1[k].a[l]=0;
            fang1[k].w=t2;
//              printf("%d:",t2);
//          print(fang1[k]);
        }
        for(int k=1;k<=m;k++) fang[k]=fang1[k];
        ans.a[i]=r;
    }
    print(ans);
}

非常好理解对不对?