题解:P15495 [ICPC 2025 APC] Tower of Hanoi

· · 题解

题意简述

每个圆盘初始位于三根柱子之一。支持单点修改圆盘所在柱,并询问一个连续大小区间内的圆盘全部移到 1 号柱所需的最少步数。

解题思路

先从小到大加入圆盘。设已经处理了 k 个较小圆盘,d_t 表示把它们从给定初态全部移到柱 t 的最少步数,柱子从 0 开始编号。

现在加入一个位于柱 x 的更大圆盘。若目标柱 t=x,这个圆盘不必移动,答案仍是 d_t

t\ne x,移动大圆盘前,全部小圆盘必须先到第三根柱:

s=3\mathbin{\mathrm{xor}}x\mathbin{\mathrm{xor}}t

这一步需要 d_s 次。随后移动大圆盘一次,再把 k 个小圆盘从 s 整体移到 t,需要 2^k-1 次。因此:

d'_t= \begin{cases} d_t & x=t \\ d_s+2^k & x\ne t \end{cases}

递推中除了三个 d_t,还需要维护 2^k。把状态写成列向量:

v=(d_0,d_1,d_2,2^k)^{\mathsf T}

加入任意一个圆盘,都是对 v 的线性变换。圆盘位于 x 时,构造一个 4\times4 矩阵 M_x。对前三行:若 x=t,就在第 t 列放 1;否则在第 s 列和最后一列各放 1。最后一行只在最后一列放 2,表示 2^{k+1}=2\cdot2^k

设较小圆盘区间的变换矩阵为 A,紧接着的较大圆盘区间为 B。先应用 A,再应用 B,合并矩阵就是:

BA

矩阵乘法满足结合律,所以可以直接用线段树维护。代码中的 merge(x,y) 表示把区间 x 放在区间 y 前面,返回 M_yM_x。这个方向不能交换,因为矩阵乘法一般不满足交换律。

一个询问区间之前没有已处理圆盘,初始状态为:

v_0=(0,0,0,1)^{\mathsf T}

设询问得到的总矩阵为 M。题目要求移到编号为 0 的柱,所以答案是 (Mv_0)_0,也就是 M_{0,3}

线段树采用迭代写法。区间查询时,左侧累积器按从左到右追加节点,右侧累积器则把新节点放在已有结果之前,最后再按顺序合并两侧。补齐到线段树底层的空位置使用单位矩阵。

矩阵大小固定为 4,一次合并是常数时间。建树需要 O(n),每次修改和询问需要 O(\log n),空间复杂度为 O(n)

正确性证明

对已经处理的圆盘数 k 归纳。没有圆盘时,移到任意柱都需要零步,最后一维为 2^0=1,初始向量正确。

加入位于 x 的最大圆盘后,若目标就是 x,最优方案只需处理较小圆盘。若目标不是 x,在移动最大圆盘前,较小圆盘必须全部离开起点柱和目标柱,只能位于第三根柱 s。完成这一步后,移动最大圆盘并把整叠小圆盘转移到目标共需 2^k 次。故矩阵的前三行准确实现最优递推,最后一行也准确把 2^k 倍增。

区间矩阵按圆盘从小到大的顺序复合。merge(A,B) 返回先执行 A、再执行 B 的矩阵 BA。矩阵乘法的结合律保证线段树无论如何划分区间,得到的总变换都与逐个加入圆盘相同。

询问以空圆盘集合的状态 v_0 为输入。总矩阵作用后,第零个分量就是把整个区间移到 1 号柱的最少步数。因此程序输出的 M_{0,3} 正确。

参考代码

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

const int N=1<<18;
const int mod=998244353;
struct node
{
    int a[4][4];
}tr[N<<1];
node one()
{
    node res={};
    for(int i=0;i<4;i++)res.a[i][i]=1;
    return res;
}
node gen(int x)
{
    node res={};
    for(int i=0;i<3;i++)
    {
        int p=x==i?i:3^x^i;
        res.a[i][p]=1;
        res.a[i][3]=x!=i;
    }
    res.a[3][3]=2;
    return res;
}
node merge(const node &x,const node &y)
{
    node res={};
    for(int i=0;i<4;i++)
    {
        for(int j=0;j<4;j++)
        {
            for(int k=0;k<4;k++)res.a[i][j]=(res.a[i][j]+(long long)y.a[i][k]*x.a[k][j])%mod;
        }
    }
    return res;
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int n,q;
    cin>>n>>q;
    for(int i=0;i<N;i++)tr[N+i]=one();
    for(int i=0;i<n;i++)
    {
        int x;
        cin>>x;
        tr[N+i]=gen(x-1);
    }
    for(int i=N-1;i;i--)tr[i]=merge(tr[i<<1],tr[i<<1|1]);
    while(q--)
    {
        char op;
        int x,y;
        cin>>op>>x>>y;
        if(op=='c')
        {
            int p=N+x-1;
            tr[p]=gen(y-1);
            while(p>>=1)tr[p]=merge(tr[p<<1],tr[p<<1|1]);
        }
        else
        {
            node l=one(),r=one();
            x+=N-1;
            y+=N;
            while(x<y)
            {
                if(x&1)l=merge(l,tr[x++]);
                if(y&1)r=merge(tr[--y],r);
                x>>=1;
                y>>=1;
            }
            cout<<merge(l,r).a[0][3]<<'\n';
        }
    }
    return 0;
}