P2371 [国家集训队]墨墨的等式

· · 题解

考察的是『同余最短路』,很巧妙的一种处理方法。这里说一下自己的理解。

题意看上去很简单,但感觉无从下手,所以我们要一层一层地分析。

1.首先求区间[l,r]的满足条件的个数,我们可以转化为求[1,r]-[1,l-1]的答案,问题变成了:小于等于x的b有多少个。

2.观察数据范围:发现n极小,不知道有什么用,a[i]的范围很有特点,10^5这个数字和很多的算法有关,而l,r的范围巨大,从l,r入手不太可取。因此我们想从n,a[i]中找点性质。

3.我们先给a[i]排个序,然后我们注意到,如果一个值val它可以被表示出来,那么所有的val+k*a[1]都可以被表示出来。

4.进一步分析:确定了一个val,val+k*a[1]都能被表示,它们对a[1]取膜后的结果是一样的!那我有一个很棒的想法:余数0,1.....a[1]-1,我想对于每一个余数p,都找到用所有的a[i],能够表示出的最小的%a[1]=p的值是多少,由此就能知道在小于等于x的所有值中%a[1]=p的值一共有多少个。

5.感觉找到能够表示出的最小的%a[1]=p的值是多少很难办?我们可以假设dp[i]表示最小的%a[1]=i的值是多少。那么可以想到一个转移方程:dp[(a[j]+i)%a[1]]=min(dp[(a[j]+i)%a[1]],dp[i]+a[j]). 你把dp换成dis,是不是感觉这很像最短路的形式?我们把dp[i]放到最短路中完成这件事情,于是本题得到完美解决。

6.总结:同余最短路是一种非常神奇的技巧,并可以在一些题目中得到完美的展现。实际上,我认为它是从dp方程转移推过来的一种优化的处理方式,它的实质是dp。 转化为最短路后,就可以使复杂度降为O(a_{1}log \ a_{1}),效率大大提高。

code:
#include<bits/stdc++.h>
#define int long long
#define reg register
using namespace std;
const int maxn=1e6+5;
const int mod=31011;
const int INF=2e16;
inline int read()
{
    int x=0,f=1;
    char ch=getchar();
    while(ch<'0' || ch>'9')
    {
        if(ch=='-')
        f=-1;
        ch=getchar();
    }
    while(ch>='0' && ch<='9')
    {
        x=x*10+ch-'0';
        ch=getchar();
    }
    return x*f;
}
int n,l,r,a[maxn],k,vis[maxn],dist[maxn];
queue<int> q;
int Work(int x)
{
    int res=0;
    for(reg int i=0;i<k;i++)
    {
        if(dist[i]!=INF)
        {
            int wns=max((x-i)/k+1-dist[i],(int)0);
            res+=wns;
        }
    }
    return res;
}
signed main()
{
    n=read();
    l=read(),r=read();
    for(reg int i=1;i<=n;i++)
    {
        a[i]=read();
    }
    sort(a+1,a+1+n);
    k=a[1];
    for(reg int i=0;i<=k;i++)
    dist[i]=2e16;
    dist[0]=0;
    vis[0]=1;
    q.push(0);
    while(!q.empty())
    {
        int x=q.front();
        q.pop();
        vis[x]=0;
        for(reg int i=1;i<=n;++i)
        {
            int t=(x+a[i])%k;
            if(dist[t]>dist[x]+a[i])
            {
                dist[t]=dist[x]+a[i];
                if(vis[t]==0)
                {
                    vis[t]=1;
                    q.push(t);
                }
            }
        }
    }
    for(reg int i=0;i<k;i++)
    {
        if(dist[i]!=INF)
        {
            dist[i]=(dist[i]-i)/k;
        }
    }
    printf("%lld\n",Work(r)-Work(l-1));
    return 0;
}