题解:P17330 [ICPC 2018 Nanjing R] Mediocre String Problem

· · 题解

题意简述

s 中选一个非空子串,从 t 中选一个非空前缀。前者长度必须严格大于后者,且两者依次拼接后为回文串。求合法的 (i,j,k) 数量。

解题思路

设从 t 中选择的前缀长度为 k,并将 s[i\mathbin{\ldotp\ldotp}j] 与这个前缀拼接。因为 j-i+1>k,整个串的回文中心一定落在来自 s 的部分中。

s 的子串按长度 k 分成两段:

s[i\mathbin{\ldotp\ldotp}j] =s[i\mathbin{\ldotp\ldotp}i+k-1]+s[i+k\mathbin{\ldotp\ldotp}j]

拼接结果为回文串,当且仅当同时满足:

\begin{aligned} \operatorname{rev}(s[i\mathbin{\ldotp\ldotp}i+k-1]) & = t[0\mathbin{\ldotp\ldotp}k-1] \\ s[i+k\mathbin{\ldotp\ldotp}j] & \text{ 为回文串} \end{aligned}

枚举第二段的起点 x=i+k。接下来只需分别统计两件事。

a_x 为以 x 开头的非空回文子串数量。对 s 运行 Manacher 算法。每个回文中心能产生一段连续的左端点:奇回文中心 c、字符半径 d 对区间 [c-d+1,c] 各贡献一次;偶回文中心右侧位置为 c、字符半径为 d 时,对区间 [c-d,c-1] 各贡献一次。使用差分数组即可在线性时间内求出所有 a_x

再记 b_x 为满足第一项条件的 k 的数量。令 r=\operatorname{rev}(s),则:

\operatorname{rev}(s[x-k\mathbin{\ldotp\ldotp}x-1]) =r[n-x\mathbin{\ldotp\ldotp}n-x+k-1]

因此 b_x 就是 tr[n-x\mathbin{\ldotp\ldotp}n-1] 的最长公共前缀长度。用扩展 KMP 一次求出 tr 的每个后缀的最长公共前缀,即可得到全部 b_x

对固定的 x,每个可行的 k 与每个以 x 开头的回文子串可以独立组合,并且会唯一确定原来的 (i,j,k)。所以答案为:

\sum_{x=1}^{n-1}a_xb_x

Manacher、扩展 KMP 和差分前缀和均为线性处理。时间复杂度为 O(n+m),空间复杂度为 O(n+m)

正确性证明

先证明上述拆分条件的充要性。

若拼接串为回文串,它末尾来自 tk 个字符会与开头来自 sk 个字符逐一对应,所以这两段互为反串。删去这两段后,剩余的 s[i+k\mathbin{\ldotp\ldotp}j] 仍关于原回文中心对称,因而也是回文串。

反过来,若两端的两段互为反串,且中间一段是回文串,那么将三段依次拼接后,从外到内的字符都两两相同,所得字符串必为回文串。

固定中间回文串的左端点 x。扩展 KMP 得到的 b_x 是最大可匹配长度。最长公共前缀的任意非空前缀仍然相同,因此合法的 k 恰有 b_x 个。差分统计得到的 a_x 则恰好枚举所有以 x 开头的非空回文子串。

任取其中一个 k 和一个回文子串 s[x\mathbin{\ldotp\ldotp}j],令 i=x-k,便唯一得到一组合法方案。反之,每组合法方案也会按 x=i+k 唯一落入这一对选择中。因此乘积 a_xb_x 不重不漏,求和所得即为答案。

参考代码

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

using ll=long long;
const int N=1000005;
const int M=2000005;
int z[N],ext[N],rad[M],cnt[N];
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    string s,t;
    cin>>s>>t;
    int n=s.size(),m=t.size();
    z[0]=m;
    int l=0,r=-1;
    for(int i=1;i<m;i++)
    {
        if(i<=r)z[i]=min(z[i-l],r-i+1);
        while(i+z[i]<m&&t[z[i]]==t[i+z[i]])z[i]++;
        if(i+z[i]-1>r)
        {
            l=i;
            r=i+z[i]-1;
        }
    }
    string rev=s;
    reverse(rev.begin(),rev.end());
    l=0;
    r=-1;
    for(int i=0;i<n;i++)
    {
        if(i<=r)ext[i]=min(z[i-l],r-i+1);
        while(ext[i]<m&&i+ext[i]<n&&t[ext[i]]==rev[i+ext[i]])ext[i]++;
        if(i+ext[i]-1>r)
        {
            l=i;
            r=i+ext[i]-1;
        }
    }
    string u;
    u+='#';
    for(auto c:s)
    {
        u+=c;
        u+='#';
    }
    l=0;
    r=-1;
    for(int i=0;i<u.size();i++)
    {
        if(i<=r)rad[i]=min(rad[l+r-i],r-i);
        while(i-rad[i]-1>=0&&i+rad[i]+1<u.size()&&u[i-rad[i]-1]==u[i+rad[i]+1])rad[i]++;
        if(i+rad[i]>r)
        {
            l=i-rad[i];
            r=i+rad[i];
        }
        if(i&1)
        {
            int j=i/2,k=(rad[i]+1)/2;
            cnt[j-k+1]++;
            cnt[j+1]--;
        }
        else
        {
            int j=i/2,k=rad[i]/2;
            if(k)
            {
                cnt[j-k]++;
                cnt[j]--;
            }
        }
    }
    for(int i=1;i<n;i++)cnt[i]+=cnt[i-1];
    ll ans=0;
    for(int i=1;i<n;i++)ans+=1LL*cnt[i]*ext[n-i];
    cout<<ans<<'\n';
    return 0;
}