题解 P4008 【[NOI2003]文本编辑器】

· · 题解

splay裸题。

看到区间插入删除,想到splay。

但是在说splay之前先说几个别的坑点。

(另外我的splay是跟rtldalao学的,代码可能有点异样但大家应该看得懂的)

1.insert的输入

我的输入方法是这样的:一开始输入了一个长度,接下来一直读入字符,如果是合法字符就接到临时字符串末尾,如果输入的合法字符已经够了,那么直接退出。

至于毒瘤的空格和换行符的问题,不用管,题面里说了,合法 \text{ASCII} 值在 [32,126]。只要判断 \text{ASCII} 即可。

strcpy(ins,"");
flag=false; //初始化
scanf("%d",&x); //长度
for(int i=0;i<x;i++){
    char ch=getchar();  //一直读入
    if(ch>=32 && ch<=126) flag=true;    //如果读到合法字符,那后面的空格都是合法的,否则可能是非法的(flag)
    if(!flag || ch<32 || ch>126){   //如果不合法
        i--;continue;   //读到的字符数--
    }
    ins[i]=ch;  //放到末尾
}

2.光标的设置

光标可以强行 O(1) 模拟,也不会对splay的后续操作发生影响。至于从0还是从1开始我们待会会讲。

3.内存的设置

由于总共只会插入2097152个字符,所以splay的节点个数和临时字符串大小只需要开这么大。(或许是数据水吧,我开2000200就过了)

接下来开始重头戏。

1.insert操作

假设光标在 x 处,字符串长度为 len。

那么仿照模板(其实大多数题目都是这样),将排名为 x-1 的点splay到根,x 的点splay到根的右儿子,然后在根的右儿子的左儿子build一颗完美的splay。可以再把这个子树splay到根上以确定随机性,但是顶多也就多2的高度,所以可以不splay。

好的,假设 x=0 怎么办?

一般来说解决方法都是头尾加哨兵元素,然后就可以避免查询排名为非正数的点了。在这里我们加入哨兵元素后,splay的应该是排名为 x 和 x+1 的点。光标位置一开始设成1而不是0。

(从前我试着偷懒,不splay上去就在下面操作,然后发现特别难实现,好不容易写出来的又wa了……大家千万不要这样,有区间操作最好就splay上去再弄)

Splay():cnt(0),pos(1){  //初始化以及两个哨兵元素('\n')
    memset(e,0,sizeof(e));
    root=addnode('\n',0);
    e[root].son[1]=addnode('\n',root);
}
int build(char *str,int fa,int l,int r){    //完美splay
    int mid=l+r>>1;
    char ch=str[mid-1];
    int x=addnode(ch,fa);   //自己
    if(l<mid) e[x].son[0]=build(str,x,l,mid-1);
    if(mid<r) e[x].son[1]=build(str,x,mid+1,r); //递归
    pushup(x);
    return x;
}
void insert(int len,char *str){
    int l=kth(pos),r=kth(pos+1);
    splay(l,0);
    splay(r,l); //寻找节点然后splay
    e[e[root].son[1]].son[0]=build(str,e[root].son[1],1,len);   //建造完美splay
    pushup(e[root].son[1]);pushup(root);
}

2.move,prev,next操作

也就move注意一下,光标改成 x+1 而不是 x。其他没什么好说的。

inline void move(int x){pos=x+1;}
inline void prv(){pos--;}
inline void nxt(){pos++;}

3.delete操作

同样的道理,把 x splay到根, (x+len-1)+2=x+len+1 splay到根的右儿子,把根的右儿子的左儿子直接断掉关系。

void remove(int len){
    int l=kth(pos),r=kth(pos+len+1);
    splay(l,0);
    splay(r,l); //寻找节点然后splay
    e[e[root].son[1]].son[0]=0; //断掉关系
    pushup(e[root].son[1]);pushup(root);
}

4.get操作

把 x splay到根, (x+len-1)+2=x+len+1 splay到根的右儿子,然后输出根的右儿子的左儿子的中序遍历。

void output(int x){ //中序遍历
    if(e[x].son[0]) output(e[x].son[0]);    //左儿子
    if(e[x].val!='\n') putchar(e[x].val);   //自己
    if(e[x].son[1]) output(e[x].son[1]);    //右儿子
}
void get(int len){
    int l=kth(pos),r=kth(pos+len+1);
    splay(l,0);
    splay(r,l); //寻找节点然后splay
    output(e[e[root].son[1]].son[0]);   //中序遍历输出
    putchar(10);
}

大致就这些啦,那么就上代码,刚刚没有详解过的地方都加了注释。开O2总用时756ms,空间42.08MB,不开O2总用时1164ms。虽说不是很优秀,但是看着题解没有这么详细的,于是就写一篇题解帮帮大家。

#include<bits/stdc++.h>
using namespace std;
struct Splay{
    #define root e[0].son[1]
    struct node{    //节点
        char val;
        int size,fa,son[2];
    }e[2100000];
    int cnt,pos;
    inline void pushup(int x){  //上传节点信息
        e[x].size=e[e[x].son[0]].size+e[e[x].son[1]].size+1;
    }
    inline int atson(int x){    //是左儿子还是右儿子
        return e[e[x].fa].son[1]==x;
    }
    inline void connect(int x,int fa,int type){ //建立关系(注:并不断开矛盾的关系)
        e[x].fa=fa;
        e[fa].son[type]=x;
    }
    inline int addnode(char val,int fa){    //添加节点
        e[++cnt].val=val;
        e[cnt].size=1;
        e[cnt].fa=fa;
        return cnt;
    }
    void rotate(int x){ //旋转
        int fa=e[x].fa,ffa=e[fa].fa;
        int type=atson(x),ttype=atson(fa);
        int B=e[x].son[type^1];
        connect(B,fa,type);connect(fa,x,type^1);connect(x,ffa,ttype);
        pushup(fa);pushup(x);
    }
    void splay(int at,int to){
        if(!to) root=at;
        while(e[at].fa!=to){
            int fa=e[at].fa;
            if(e[fa].fa!=to) rotate(atson(at)^atson(fa)?at:fa);
            rotate(at);
        }
    }
    Splay():cnt(0),pos(1){
        memset(e,0,sizeof(e));
        root=addnode('\n',0);
        e[root].son[1]=addnode('\n',root);
    }
    inline void move(int x){pos=x+1;}
    int build(char *str,int fa,int l,int r){
        int mid=l+r>>1;
        char ch=str[mid-1];
        int x=addnode(ch,fa);
        if(l<mid) e[x].son[0]=build(str,x,l,mid-1);
        if(mid<r) e[x].son[1]=build(str,x,mid+1,r);
        pushup(x);
        return x;
    }
    int kth(int x){
        int now=root;
        while(now){
            int sum=e[e[now].son[0]].size+1;
            if(x==sum) return now;
            if(x<sum) now=e[now].son[0];
            else{
                x-=sum;
                now=e[now].son[1];
            }
        }
    }
    void insert(int len,char *str){
        int l=kth(pos),r=kth(pos+1);
        splay(l,0);
        splay(r,l);
        e[e[root].son[1]].son[0]=build(str,e[root].son[1],1,len);
        pushup(e[root].son[1]);pushup(root);
    }
    void remove(int len){
        int l=kth(pos),r=kth(pos+len+1);
        splay(l,0);
        splay(r,l);
        e[e[root].son[1]].son[0]=0;
        pushup(e[root].son[1]);pushup(root);
    }
    void output(int x){
        if(e[x].son[0]) output(e[x].son[0]);
        if(e[x].val!='\n') putchar(e[x].val);
        if(e[x].son[1]) output(e[x].son[1]);
    }
    void get(int len){
        int l=kth(pos),r=kth(pos+len+1);
        splay(l,0);
        splay(r,l);
        output(e[e[root].son[1]].son[0]);
        putchar(10);
    }
    inline void prv(){pos--;}
    inline void nxt(){pos++;}
}sp;
int t;
char op[7],ins[2100000];
int main(){
    scanf("%d",&t);
    while(t--){
        strcpy(ins,"");
        int x;
        bool flag=false;
        scanf("%s",op);
        switch(op[0]){  //哪个操作
            case 'M':   //move
                scanf("%d",&x);sp.move(x);break;
            case 'I':   //insert
                scanf("%d",&x);
                for(int i=0;i<x;i++){
                    char ch=getchar();
                    if(ch>=32 && ch<=126) flag=true;
                    if(!flag || ch<32 || ch>126){
                        i--;continue;
                    }
                    ins[i]=ch;
                }
                sp.insert(x,ins);break;
            case 'D':   //delete
                scanf("%d",&x);sp.remove(x);break;
            case 'G':   //get
                scanf("%d",&x);sp.get(x);break;
            case 'P':   //prev
                sp.prv();break;
            case 'N':   //next
                sp.nxt();break;
        }
    }
}