P16693 Tokitsukaze and Palindrome Border TJ

· · 题解

闲话

我来把没讲的 trie + manacher 做法写了

trie + manacher 还是与 PAM 太相似了,我把 O(n\log n) 做法也说一下。

然而没根号算法跑的快

思路

先将所有字符串正序加入一棵字典树,顺便预处理出 sum_i 表示若节点 i 代表的串是回文的则为串的长度,反之则为 0,这个可以对于每个串跑 manacher,并在 manacher 中顺便将字典树也给插入来解决,大概如下:

cin >> str,ch[cnt = 0] = '!';
for (char j : str) ch[++cnt] = j,ch[++cnt] = '.';
ch[++cnt] = '?';
int mid = 0,rg = 0,now = 1;
for (int j = 1;j < cnt;j++){
    if (j < rg) len[j] = min(len[(mid<<1)-j],len[mid]+mid-j);
    else len[j] = 1;
    for (;ch[j-len[j]] == ch[j+len[j]];len[j]++);
    if (rg < j+len[j]) rg = j+len[j],mid = j;
    if (j&1){//若该字符不是辅助字符
        now = to[now][ch[j]-'a'] = (to[now][ch[j]-'a'] ?: ++trcnt);
        if (j/2+1 == len[j/2+1]) sum[now] = j/2+1;//若该前缀是回文的
    }
}

接下来考虑如何求询问,考虑容斥:

\begin{align*} \sum_{i\notin B}\sum_{j\notin B}f(s_i,s_j)=\quad&\sum_{i}\sum_{j}f(s_i,s_j)\\ -&\sum_{i\in B}\sum_{j}f(s_i,s_j)\\ -&\sum_{i}\sum_{j\in B}f(s_i,s_j)\\ +&\sum_{i\in B}\sum_{j\in B}f(s_i,s_j) \end{align*}

考虑如何求形如 \sum_{i\in B}\sum_{j\in B}f(s_i,s_j) 的式子,发现单个 f(s_i,s_j) 的贡献,发现就是 s_i 在字典树上的结尾与 s_j 倒过来在字典树上的结尾的所有公共祖先的 sum 之和。

根据那个 trick,我们将 1s_i 的结尾的所有节点都打一个标记,那么 f(s_i,s_j) 就是 1s_j 倒过来在字典树上的结尾的所有节点的标记数乘上 sum_i 的和。

对于这个直接树链剖分完线段树维护就可以做到 O(\log^2n)

那么现在两个贡献为正的我们会求了,考虑剩下两个负的。

对于第一个,不难发现可以在求的过程中顺便记录,对于第二个要注意的是 f(s_i,s_j) \not=f(s_j,s_i),所以我们需要反过来求一次即修改 1 到反向结尾,查询 1 到结尾。

然后这道题就可以做到 O(n\log^2n),这一部分代码如下:

for (int i = 1;i <= n;i++) upd(ed[i],1);
for (int i = 1;i <= n;i++) all += sub[i] = find(rved[i]);
for (int i = 1;i <= n;i++) upd(ed[i],-1),upd(rved[i],1);
for (int i = 1;i <= n;i++) sub[i] += find(ed[i]);
for (int i = 1;i <= n;i++) upd(rved[i],-1); 
cin >> q;
while (q--){
    cin >> k,now = 0;
    for (int i = 1;i <= k;i++)
        cin >> a[i],now += sub[a[i]],upd(ed[a[i]],1);
    for (int i = 1;i <= k;i++) now -= find(rved[a[i]]);
    for (int i = 1;i <= k;i++) upd(ed[a[i]],-1);
    cout << all-now << "\n";
}

接下来我们考虑如何将瓶颈树链剖分给优化掉,发现我们一直都是先修改再查询,考虑虚树,为了方便计算虚树上的 sum,我们预处理出 pre_i=pre_{fa_i}+sum_i,那么 sum_i=pre_i-pre_{fa_i},这样子对于虚树直接搞两次树上前缀和即可,如下:

void init(int x,int y){
    for (int i : v[x]){
        if (i == y) continue;
        init(i,x),upd[x] += upd[i];
    }
    ans[x] = upd[x]*(pre[x]-pre[y]);
}
void init2(int x,int y){
    ans[x] += ans[y];
    for (int i : v[x])
        if (i != y) init2(i,x);
}

然后这道题就做完了,似乎理论上如果你把所有查询离线下来以使用计数排序并用单调栈求虚树和四毛子加欧拉序求 lca 的话可以做到 O(n),但这太复杂了我写不来

Code

:::info[O(n \log^2 n)]

#include <bits/stdc++.h>
#define int long long
#define pii pair<int,int>
#define fi first
#define se second
#define maxn 600005
using namespace std;
struct sgnd{
    int sum,rsum,flg;
}e[maxn<<2];
int to[maxn][26],sum[maxn],trcnt;
int fa[maxn],siz[maxn],son[maxn],top[maxn],dfn[maxn],pos[maxn],dfcnt;
int n,q,k,cnt,now,all,sub[maxn],ed[maxn],rved[maxn],len[maxn],a[maxn];
string str;
char ch[maxn];
vector<pii> v[maxn];
void giveflg(int x,int y){
    e[x].rsum += e[x].sum*y,e[x].flg += y;
}
void push_down(int x){
    giveflg(x<<1,e[x].flg),giveflg(x<<1|1,e[x].flg),e[x].flg = 0;
}
void push_up(int x){
    e[x].sum = e[x<<1].sum+e[x<<1|1].sum;
    e[x].rsum = e[x<<1].rsum+e[x<<1|1].rsum; 
}
void update(int l,int r,int x,int y,int id){
    if (l == r) e[id].sum += y;
    else{
        int mid = l + (r - l >> 1);
        push_down(id);
        if (x <= mid) update(l,mid,x,y,id<<1);
        else update(mid+1,r,x,y,id<<1|1);
        push_up(id);
    }
}
void change(int l,int r,int sl,int sr,int x,int id){
    if (sl <= l && r <= sr) giveflg(id,x);
    else{
        int mid = l + (r - l >> 1);
        push_down(id);
        if (sl <= mid) change(l,mid,sl,sr,x,id<<1);
        if (sr > mid) change(mid+1,r,sl,sr,x,id<<1|1);
        push_up(id);
    }
}
int find(int l,int r,int sl,int sr,int id){
    if (sl <= l && r <= sr) return e[id].rsum;
    else{
        int mid = l + (r - l >> 1),sum = 0;
        push_down(id);
        if (sl <= mid) sum += find(l,mid,sl,sr,id<<1);
        if (sr > mid) sum += find(mid+1,r,sl,sr,id<<1|1);
        return sum;
    }
}
void dfs1(int x,int y){
    fa[x] = y,siz[x] = 1;
    for (int p = 0;p < 26;p++){
        int i = to[x][p];if (!i) continue;
        dfs1(i,x);
        siz[x] += siz[i];
        if (siz[son[x]] < siz[i]) son[x] = i;
    }
}
void dfs2(int x,int y){
    top[x] = y,dfn[++dfcnt] = x,pos[x] = dfcnt;
    if (son[x]) dfs2(son[x],y);
    for (int p = 0;p < 26;p++){
        int i = to[x][p];
        if (i && i != son[x]) dfs2(i,i);
    }
}
void show(int x,int y){
    while (x)
        change(1,dfcnt,pos[top[x]],pos[x],y,1),x = fa[top[x]];
}
int find(int x){
    int ans = 0;
    while (x)
        ans += find(1,dfcnt,pos[top[x]],pos[x],1),x = fa[top[x]];
    return ans;
}
signed main(){
    ios::sync_with_stdio(0);
    cin.tie(0),cout.tie(0);
    cin >> n,trcnt = 1;
    for (int i = 1;i <= n;i++){
        cin >> str,ch[cnt = 0] = '!';
        for (char j : str) ch[++cnt] = j,ch[++cnt] = '.';
        ch[++cnt] = '?';
        int mid = 0,rg = 0,now = 1;
        for (int j = 1;j < cnt;j++){
            if (j < rg) len[j] = min(len[(mid<<1)-j],len[mid]+mid-j);
            else len[j] = 1;
            for (;ch[j-len[j]] == ch[j+len[j]];len[j]++);
            if (rg < j+len[j]) rg = j+len[j],mid = j;
            if (j&1){
                now = to[now][ch[j]-'a'] = (to[now][ch[j]-'a'] ?: ++trcnt);
                if (j/2+1 == len[j/2+1]) sum[now] = j/2+1;
            }
        }
        ed[i] = now;
        reverse(str.begin(),str.end());
        now = 1;
        for (char j : str) now = to[now][j-'a'] = (to[now][j-'a'] ?: ++trcnt);
        rved[i] = now;
    }
    dfs1(1,0),dfs2(1,1);
    for (int i = 1;i <= trcnt;i++) update(1,dfcnt,pos[i],sum[i],1);
    for (int i = 1;i <= n;i++) show(ed[i],1);
    for (int i = 1;i <= n;i++) all += sub[i] = find(rved[i]);
    for (int i = 1;i <= n;i++) show(ed[i],-1),show(rved[i],1);
    for (int i = 1;i <= n;i++) sub[i] += find(ed[i]);
    for (int i = 1;i <= n;i++) show(rved[i],-1); 
    cin >> q;
    while (q--){
        cin >> k,now = 0;
        for (int i = 1;i <= k;i++)
            cin >> a[i],now += sub[a[i]],show(ed[a[i]],1);
        for (int i = 1;i <= k;i++) now -= find(rved[a[i]]);
        for (int i = 1;i <= k;i++) show(ed[a[i]],-1);
        cout << all-now << "\n";
    }
    return 0;
}

::: :::info[O(n \log n)]

#include <bits/stdc++.h>
#define int long long
#define maxn 600005
using namespace std;
int to[maxn][26],sum[maxn],pre[maxn],upd[maxn],ans[maxn],trcnt;
int dfn[maxn],pos[maxn],dfcnt;
int st[20][maxn];
int n,q,k,cnt,now,all,sub[maxn],ed[maxn],rved[maxn],len[maxn],a[maxn],num[maxn<<2];
string str;
char ch[maxn];
vector<int> v[maxn];
int cmpdfn(int x,int y){
    return (pos[x] < pos[y] ? x : y);
}
void dfs(int x,int y){
    dfn[pos[x]=++dfcnt] = x;
    st[0][pos[x]] = y;
    pre[x] = sum[x]+pre[y];
    for (int p = 0;p < 26;p++)
        if (to[x][p]) dfs(to[x][p],x);
}
void init(){
    for (int i = 1;i <= 19;i++)
        for (int j = 1;j+(1<<i)-1 <= dfcnt;j++)
            st[i][j] = cmpdfn(st[i-1][j],st[i-1][j+(1<<i-1)]);
}
int lca(int x,int y){
    if (x == y) return x;
    if ((x=pos[x]) > (y=pos[y])) swap(x,y);
    int k = __lg(y-x++);
    return cmpdfn(st[k][x],st[k][y-(1<<k)+1]);
}
bool cmp(int x,int y){
    return (pos[x] < pos[y]);
}
void build(){
    for (int i = 1;i <= k;i++) num[i*2-1] = ed[a[i]],num[i*2] = rved[a[i]];
    num[2*k+1] = 1;
    int len = 2*k+1;
    sort(num+1,num+len+1,cmp);
    for (int i = 1,rg = len;i < rg;i++) num[++len] = lca(num[i],num[i+1]);
    sort(num+1,num+len+1,cmp),len = unique(num+1,num+len+1)-num-1;
    for (int i = 1;i <= len;i++) v[num[i]].clear(),upd[num[i]] = 0;
    for (int i = 1;i < len;i++) v[lca(num[i],num[i+1])].push_back(num[i+1]);
}
void init(int x,int y){
    for (int i : v[x]){
        if (i == y) continue;
        init(i,x),upd[x] += upd[i];
    }
    ans[x] = upd[x]*(pre[x]-pre[y]);
}
void init2(int x,int y){
    ans[x] += ans[y];
    for (int i : v[x])
        if (i != y) init2(i,x);
}
signed main(){
    ios::sync_with_stdio(0);
    cin.tie(0),cout.tie(0);
    cin >> n,trcnt = 1;
    for (int i = 1;i <= n;i++){
        cin >> str,ch[cnt = 0] = '!';
        for (char j : str) ch[++cnt] = j,ch[++cnt] = '.';
        ch[++cnt] = '?';
        int mid = 0,rg = 0,now = 1;
        for (int j = 1;j < cnt;j++){
            if (j < rg) len[j] = min(len[(mid<<1)-j],len[mid]+mid-j);
            else len[j] = 1;
            for (;ch[j-len[j]] == ch[j+len[j]];len[j]++);
            if (rg < j+len[j]) rg = j+len[j],mid = j;
            if (j&1){
                now = to[now][ch[j]-'a'] = (to[now][ch[j]-'a'] ?: ++trcnt);
                if (j/2+1 == len[j/2+1]) sum[now] = j/2+1;
            }
        }
        ed[i] = now;
        reverse(str.begin(),str.end());
        now = 1;
        for (char j : str) now = to[now][j-'a'] = (to[now][j-'a'] ?: ++trcnt);
        rved[i] = now;
    }
    dfs(1,0),init();
    k = n;for (int i = 1;i <= n;i++) a[i] = i;
    build();for (int i = 1;i <= n;i++) upd[ed[i]]++;
    init(1,0),init2(1,0);
    for (int i = 1;i <= n;i++) all += sub[i] = ans[rved[i]];

    k = n;for (int i = 1;i <= n;i++) a[i] = i;
    build();for (int i = 1;i <= n;i++) upd[rved[i]]++;
    init(1,0),init2(1,0);
    for (int i = 1;i <= n;i++) sub[i] += ans[ed[i]];

    cin >> q;
    while (q--){
        cin >> k;for (int i = 1;i <= k;i++) cin >> a[i];
        build();for (int i = 1;i <= k;i++) upd[ed[a[i]]]++;
        init(1,0),init2(1,0),now = 0;
        for (int i = 1;i <= k;i++) now += sub[a[i]]-ans[rved[a[i]]];
        cout << all-now << "\n";
    }

    return 0;
}

:::