题解:P17094 [ICPC 2017 Qingdao R] Collecting Cents

· · 题解

题意简述

给定初始资金 n、利率分母 d 和年数 y。每笔存款的本金和存期均为整数,存 c 美分共 t 年后得到 c+\lfloor ct/d\rfloor 美分。求 y 年后最多拥有的资金。

解题思路

同一时刻、同一存期的两笔存款应当合并,因为:

\left\lfloor\frac{at}{d}\right\rfloor+\left\lfloor\frac{bt}{d}\right\rfloor\le\left\lfloor\frac{(a+b)t}{d}\right\rfloor

还可以限制多年期存款的规模。若存期 t\ge2 且本金不少于 d,从中拆出 d 美分。原方案中,这部分到期后变为 d+t。改为先存 1 年,再存 t-1 年,可得到:

d+1+\left\lfloor\frac{(d+1)(t-1)}{d}\right\rfloor\ge d+t

因此,每种多年期存款的本金仅需枚举 0\sim d-1。若 t>d,先存 d 年得到 2c,再存剩余年份也不会更差。因此存期仅需考虑到 d。未分配的钱全部存 1 年。

状态记录当前现金 C 和待到期金额 p_i。其中 p_i 会在 i 年后到期,且 1\le i<d。设本次给 t 年期存款分配 z_t,并记:

g_t(z)=z+\left\lfloor\frac{zt}{d}\right\rfloor

设当前还剩 r 年。枚举 z_t\in[0,d-1],其中 2\le t\le\min(d,r)。若多年期本金总和为 s,则其余 C-s 全部存 1 年。下一年现金为:

C-s+\left\lfloor\frac{C-s}{d}\right\rfloor+p_1

随后将待到期序列左移,并加入本轮新存款。

直接保留全部动作仍然偏慢。写成 C=qd+r。定义:

\begin{aligned} h_1 & =(r-s)+\left\lfloor\frac{r-s}{d}\right\rfloor \\ h_k & =h_1+\sum_{t=2}^k g_t(z_t) \end{aligned}

若动作 A 使用的本金不多于动作 B,且每个 h_k^A\ge h_k^B,则 A 能模拟 B。这类支配关系仅依赖 d、剩余年数和 r,可以预处理。

状态也可按支配关系删除。设状态 A 的现金不少于状态 B,初始差额为 \Delta。依次检查未来 d-1 年,并更新差额下界:

\Delta\to\Delta+\left\lfloor\frac{\Delta}{d}\right\rfloor+p_i^A-p_i^B

若差额始终非负,A 可以完整模拟 B。这里使用的是保守下界,因为 \lfloor(x+\Delta)/d\rfloor-\lfloor x/d\rfloor\ge\lfloor\Delta/d\rfloor

每个待到期位置最多汇合 d-1 笔多年期存款。每笔返还不超过 2(d-1),故 p_i\le50。代码用 6 位保存一个 p_i,再用基数排序按待到期序列去重。

设每层保留 P 个状态,每个余数保留 C 个动作。转移共产生 K\le PC 个不同候选。每年转移和排序为 O(PC+K),支配检查为 O(KPd)。本题 d\le6,预处理前也仅有 d^{d-1}\le7776 个动作。答案可能达到 2^{63},所以使用 unsigned long long

参考代码

#include <bits/stdc++.h>
using namespace std;

using ull=unsigned long long;
const int D=11;
const int B=6;
const ull base=1ULL<<B;
struct choice
{
    ull key;
    int sum,cst;
};
struct state
{
    ull key,csh;
    array<unsigned char,5> p;
};
struct full
{
    ull key;
    int sum,cst;
    int pre[D];
};
int d,y,lim,rem;
int z[D];
bool vis[D];
vector<choice> ch[D][D][D];
vector<full> all;
vector<state> f,nf,st,buf;
vector<pair<ull,ull>> cv,tmp;
void dfs(int t,int sum)
{
    if(t>lim)
    {
        int x=rem-sum,val=x>=0?x+x/d:x-(-x+d-1)/d;
        full cur={0,sum,0,{}};
        cur.pre[0]=val;
        cur.cst=rem-val;
        for(int i=2;i<=d;i++)
        {
            int q=i<=lim?z[i]+z[i]*i/d:0;
            cur.key|=ull(q)<<(B*(i-2));
            val+=q;
            cur.pre[i-1]=val;
        }
        all.push_back(cur);
        return;
    }
    for(int i=0;i<d;i++)
    {
        z[t]=i;
        dfs(t+1,sum+i);
    }
}
bool cover(const full &a,const full &b)
{
    if(a.sum>b.sum)return 0;
    for(int i=0;i<d;i++)if(a.pre[i]<b.pre[i])return 0;
    return 1;
}
void init()
{
    if(vis[d])return;
    vis[d]=1;
    for(lim=1;lim<=d;lim++)
    {
        for(rem=0;rem<d;rem++)
        {
            all.clear();
            dfs(2,0);
            vector<full> kp;
            for(auto &cur:all)
            {
                bool bad=0;
                for(auto &q:kp)
                {
                    if(cover(q,cur))
                    {
                        bad=1;
                        break;
                    }
                }
                if(bad)continue;
                kp.erase(remove_if(kp.begin(),kp.end(),[&](const full &q)
                {
                    return cover(cur,q);
                }),kp.end());
                kp.push_back(cur);
            }
            for(auto &cur:kp)ch[d][lim][rem].push_back({cur.key,cur.sum,cur.cst});
        }
    }
}
void radix_sort()
{
    int k=(B*(d-1)+9)/10,siz=cv.size();
    tmp.resize(siz);
    int cnt[1024];
    for(int i=0;i<k;i++)
    {
        memset(cnt,0,sizeof cnt);
        int sh=i*10;
        for(auto &x:cv)cnt[(x.first>>sh)&1023]++;
        for(int j=1;j<1024;j++)cnt[j]+=cnt[j-1];
        for(int j=siz-1;j>=0;j--)tmp[--cnt[(cv[j].first>>sh)&1023]]=cv[j];
        cv.swap(tmp);
    }
}
void prune()
{
    radix_sort();
    st.clear();
    int siz=cv.size();
    for(int i=0;i<siz;)
    {
        int r=i+1;
        ull key=cv[i].first,csh=cv[i].second;
        while(r<siz&&cv[r].first==key)
        {
            csh=max(csh,cv[r].second);
            r++;
        }
        state cur={key,csh,{}};
        ull x=key;
        for(int j=0;j<d-1;j++)
        {
            cur.p[j]=x%base;
            x>>=B;
        }
        st.push_back(cur);
        i=r;
    }
    ull mx=0,mn=ULLONG_MAX;
    for(auto &cur:st)
    {
        mx=max(mx,cur.csh);
        mn=min(mn,cur.csh);
    }
    vector<int> cnt(mx-mn+1);
    for(auto &cur:st)cnt[mx-cur.csh]++;
    int sum=0;
    for(auto &x:cnt)
    {
        int cur=x;
        x=sum;
        sum+=cur;
    }
    buf.resize(st.size());
    for(auto &cur:st)buf[cnt[mx-cur.csh]++]=cur;
    st.swap(buf);
    nf.clear();
    for(auto &cur:st)
    {
        bool bad=0;
        for(auto &q:nf)
        {
            ull dif=q.csh-cur.csh;
            bool ok=1;
            for(int i=0;i<d-1;i++)
            {
                dif+=dif/d;
                if(q.p[i]>=cur.p[i])dif+=q.p[i]-cur.p[i];
                else
                {
                    ull val=ull(cur.p[i]-q.p[i]);
                    if(dif>=val)dif-=val;
                    else
                    {
                        ok=0;
                        break;
                    }
                }
            }
            if(ok)
            {
                bad=1;
                break;
            }
        }
        if(!bad)nf.push_back(cur);
    }
}
ull solve(ull n)
{
    f.clear();
    f.push_back({0,n,{}});
    for(int i=0;i<y;i++)
    {
        cv.clear();
        int m=min(d,y-i);
        for(auto &cur:f)
        {
            ull q=cur.csh/d,bck=cur.key%base;
            int r=int(cur.csh-q*d);
            for(auto &x:ch[d][m][r])
            {
                if(cur.csh<ull(x.sum))continue;
                ull csh=cur.csh+q-x.cst+bck,key=(cur.key>>B)+x.key;
                cv.push_back({key,csh});
            }
        }
        prune();
        f.swap(nf);
    }
    return f[0].csh;
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int T;
    cin>>T;
    while(T--)
    {
        ull n;
        cin>>n>>d>>y;
        init();
        cout<<solve(n)<<'\n';
    }
    return 0;
}