[Str记录]P5211 [ZJOI2017]字符串
command_block
·
·
个人记录
the 500th Blog.
题意 : 维护一个长度为 n 的动态字符串 s ,字符集为 [1,10^9]∩Z。
支持下列操作 :
------------
配合 [Border理论小记](https://www.luogu.com.cn/blog/command-block/border-li-lun-xiao-ji) 食用。
记 ${\rm suf}(S)$ 为 $S$ 的后缀集合,${\rm minsuf}(S)$ 为 $S$ 的最小后缀。
定义 ${\rm Ssuf}(S)=\{T\in{\rm suf}(S)\big|\exists V,TV={\rm minsuf}(SV)\}$。
即在 $S$ 后面加上一个串后,可能成为最小后缀的后缀集合。
- **结论** :对于同串的两个 ${\rm Ssuf}\ U,V$ ,若 $|U|<|V|$ 则 $2|U|\leq |V|$。
**推论** : $|{\rm Ssuf}(S)|\leq O(\log n)$。
考虑使用线段树维护 ${\rm Ssuf}$。
在合并两个儿子的信息时,假设左侧的串为 $S_1$ ,右侧的串为 $S_2$。
则我们要通过 ${\rm Ssuf}(S_1),{\rm Ssuf}(S_2)$ 得到 ${\rm Ssuf}(S_1S_2)$。
由于在线段树上有 $|S_2|\leq |S_1|\leq |S_2|+1$ ,由结论得 $S_1$ 中至多有一个 $S_1S_2$ 的 $\rm Ssuf$。
如何从 ${\rm Ssuf}(S_1)$ 中挑出那个可能的 $\rm Ssuf$?
先在所有 $U\in{\rm Ssuf}(S_1)$ 的后面加上 $S_2$ ,然后比较。
比较 $U,V$ 时,不妨设 $US_2<VS_2$ ,若 $US_2$ 不是 $VS_2$ 的前缀,则保留 $US_2$。
否则,保留较长的那个。(原因可见结论的证明)
不难发现,就算我们挑出来的不是真正的 $\rm Ssuf$ (即 ${\rm Ssuf}(S_1)$ 没能导出 ${\rm Ssuf}(S_1S_2)$ 中的任一个元素),线段树上维护的集合大小仍然是 $O(\log n)$ 的。
若要维护真正的 ${\rm Ssuf}(S)$ ,则要利用下述结论 :
- **结论** :对于同串的两个 ${\rm Ssuf}\ U,V$ ,若 $|U|<|V|$ 则 $U$ 是 $V$ 的前缀。
当合并 ${\rm Ssuf}(S_1),{\rm Ssuf}(S_2)$ 时,先将 ${\rm Ssuf}(S_1)$ 简化成一个后缀 $P$,然后和 ${\rm Ssuf}(S_2)$ 中的每个 $V$ 串依次做如下操作 :
- 若 $V$ 是 $P$ 的前缀,不做任何操作。
- 若 $V$ 不是 $P$ 的前缀,且 $V<P$ ,抛弃 $P$。
- 若 $V$ 不是 $P$ 的前缀,且 $V>P$ ,抛弃 $V$。
这样可以有效减小常数。
查询时,将得到的大小总和为 $O(\log^2 n)$ 的集合全部检查一遍。
在上面的维护过程中,我们需要比较两个字符串的大小。动态字符串比较大小一般采用 $\rm Hash$。
整个维护中涉及了 $O(n\log n+m\log^2n)$ 次比较。
若使用线段树维护 $\rm Hash$ ,一次查询是 $O(\log n)$ 的,一次比较是 $O(\log^2n)$ 的,复杂度为 $O(n\log^3 n+m\log^4 n)$ ,无法通过。
改为用分块维护 $\rm Hash$ ,这样一次修改是 $O(\sqrt{n})$ ,查询优化到了 $O(1)$,复杂度为 $O(n\log^2n+m\log^3n+m\sqrt{n})$。
注意,在修改区间 $[l,r]$ 后,还要在维护 $\rm Ssuf$ 的线段树上进行更新。
自然溢出被卡傻了,最后单模过了……
```cpp
#include<algorithm>
#include<cstring>
#include<cstdio>
#include<vector>
#include<cmath>
#define pb push_back
#define ll long long
#define MaxN 205000
using namespace std;
const int det=300000000,mod=1000000007,buf=998244353;
int n,BS,BS2,s[MaxN],tag[666];
int pw[MaxN],spw[666],o0[MaxN],o1[666],s1[666];
void Init()
{
pw[0]=spw[0]=1;
for (int i=1;i<=n;i++)pw[i]=1ll*pw[i-1]*buf%mod;
for (int i=1;i<BS;i++)spw[i]=(spw[i-1]+pw[i])%mod;
for (int p=0;p<=n;p+=BS){
o0[p]=s[p];
for (int i=p+1;i<p+BS;i++)o0[i]=(1ll*o0[i-1]*buf+s[i])%mod;
}
for (int t=0;t*BS<=n;t++){
o1[t]=o0[t*BS+BS-1];
s1[t]=(1ll*s1[t-1]*pw[BS]+o1[t])%mod;
}
}
void upds1()
{for (int t=0;t*BS<=n;t++)s1[t]=(1ll*s1[t-1]*pw[BS]+o1[t])%mod;}
void add(int p,int c)
{
int bp=p/BS,bl=bp*BS;
for (int i=0;i<bp;i++){
tag[i]+=c;
o1[i]=(o1[i]+1ll*c*spw[BS-1])%mod;
}for (int i=bl;i<=p;i++)s[i]+=c;
o0[bl]=s[bl];
for (int i=bl+1;i<bl+BS;i++)o0[i]=(1ll*o0[i-1]*buf+s[i])%mod;
o1[bp]=(o0[bp*BS+BS-1]+1ll*tag[bp]*spw[BS-1])%mod;
}
ll get(int p)
{return 1ll*s1[(p>>BS2)-1]*pw[(p&(BS-1))+1]+o0[p]+1ll*spw[p&(BS-1)]*tag[p>>BS2];}
int gets(int p){return tag[p/BS]+s[p];}
bool tr(int u,int v,int len)
{return ((get(u+len-1)-get(v+len-1))+(get(v-1)-get(u-1))%mod*pw[len])%mod==0;}
int ext(int u,int v,int lim)
{
int r=lim-max(u,v)+1,l=0,mid;
while(l<r){
mid=(l+r+1)>>1;
if (tr(u,v,mid))l=mid;
else r=mid-1;
}return l;
}
int cmp0(int u,int v,int lim){
int l=ext(u,v,lim);
if (l==lim-max(u,v)+1)return 2;
return gets(u+l)<gets(v+l);
}
bool cmp1(int u,int v,int lim){
int l=ext(u,v,lim);
return u+l>lim||(v+l<=lim&&gets(u+l)<gets(v+l));
}
vector<int> merge(const vector<int> &ls,const vector<int> &rs,int lim)
{
vector<int> ret(rs);
int p=ls[0];
for (int i=1;i<ls.size();i++){
int fl=cmp0(p,ls[i],lim);
if (fl==2)p=min(p,ls[i]);
if (fl==0)p=ls[i];
}
while(!ret.empty()){
int fl=cmp0(p,ret.back(),lim);
if (fl==2)break;
if (fl==0)return ret;
if (fl==1)ret.pop_back();
}ret.pb(p);
return ret;
}
vector<int> p[MaxN<<2];
void up(int u,int r)
{p[u]=merge(p[u<<1],p[u<<1|1],r);}
void build(int l=1,int r=n,int u=1)
{
if (l==r){p[u].pb(l);return ;}
int mid=(l+r)>>1;
build(l,mid,u<<1);
build(mid+1,r,u<<1|1);
up(u,r);
}
int wfl,wfr,ret;
void upd(int l=1,int r=n,int u=1)
{
if (wfl<=l&&r<=wfr)return ;
int mid=(l+r)>>1;
if (wfl<=mid)upd(l,mid,u<<1);
if (mid<wfr)upd(mid+1,r,u<<1|1);
up(u,r);
}
void qry(int l=1,int r=n,int u=1)
{
if (wfl<=l&&r<=wfr){
for (int i=0;i<p[u].size();i++)
if (cmp1(p[u][i],ret,wfr))
ret=p[u][i];
return ;
}int mid=(l+r)>>1;
if (wfl<=mid)qry(l,mid,u<<1);
if (mid<wfr)qry(mid+1,r,u<<1|1);
}
int m;
int main()
{
scanf("%d%d",&n,&m);
for (BS2=1;(1<<2*BS2)<n;BS2++);
BS=1<<BS2;
for (int i=1;i<=n;i++)
{scanf("%d",&s[i]);s[i]+=det;}
Init();build();
for (int i=1,op;i<=m;i++){
scanf("%d%d%d",&op,&wfl,&wfr);
if (op==1){
int c;
scanf("%d",&c);
add(wfl-1,-c);add(wfr,c);
upds1();upd();
}else {
ret=wfl++;qry();
printf("%d\n",ret);
}
}return 0;
}
```