dfs T了但是换种for写法A了,是我的SAM常数太大了嘛
查看原帖
dfs T了但是换种for写法A了,是我的SAM常数太大了嘛
105820
阿尔托莉雅丶楼主2022/9/8 19:54

两者方法求,出每个节点的及其所有后代的权值和

分别放在 f, g 数组里

详见代码里的assert

#include <iostream>
#include <algorithm>
#include <math.h>
#include <vector>
#include <cstdio>
#include <cstring>
#include <assert.h>
#include <map>
using namespace std;
typedef long long ll;
const int N = 5e5 + 5;   //remember to modify the range of the data!!
const int mod = 1e9 + 7;
const int inf = 0x3f3f3f3f;

int n, m, T;
ll t, k;
string ans;
string s;
ll a[N], f[N << 1], vis[N << 1], g[N << 1];

struct SAM
{
    int sz, last;                   //节点数 最后一个节点编号
    vector <int> len, link;         //该状态所含的最长串的长度,和后缀链接
    vector <int> endsz;              //该状态endpos集合的大小
    vector <vector <int> > tran;    //转移
    //vector <map <int, int> > tran;//字符集大的时候用map
    
    //字符串长度 字符集大小
    SAM(int lenth, int sigma) : sz(1), last(1) 
    {
        len.resize(lenth * 2, 0);
        link.resize(lenth * 2 , 0);
        endsz.resize(lenth * 2, 0);
        tran.resize(lenth * 2, vector <int>(sigma, 0));
        // tran.resize(lenth * 2, map <int, int>()) //字符集大的时候用map
    }
    void extend(int c)
    {
        int p, cur = ++ sz;
        len[cur] = len[last] + 1;
        //情况1 直接为每个加上一个转移
        for(p = last; p > 0 && tran[p][c] == 0; p = link[p])
            tran[p][c] = cur;
        if(p == 0)                      //路径上每个点都无到c的转移
            link[cur] = 1;
        else
        {
            int q = tran[p][c];         //情况2 存在一个有到c的转移
            
            if(len[q] == len[p] + 1)    // A 类
                link[cur] = q;
            else                        // B 类
            {
                int clone = ++ sz;      //将q分成两部分
                len[clone] = len[p] + 1;//只保留小于等于len[p] + 1长的串     
                tran[clone] = tran[q];  //复制一个q
                link[clone] = link[q];
                while(p > 0 && tran[p][c] == q)  //将原来转移是q的改为这个拆出来的新的节点
                    tran[p][c] = clone, p = link[p];
                link[cur] = link[q] = clone;
            }
        }
        endsz[cur] = 1;
        last = cur;                     //更新last
    }
    void construct(const string& s)
    {
        for(int i = 0; i < s.size(); i++)
            extend(s[i] - 'a');
    }
    void getendsz(void)
    {
        vector <int> b(sz + 1, 0), res(sz + 1, 0);//按每个节点的len桶排序
        for(int i = 1; i <= sz; i++)
        {
            assert(len[i] <= sz);
            b[len[i]]++;
        }
        for(int i = 1; i <= sz; i++)
            b[i] += b[i - 1];
        for(int i = sz; i >= 1; i--)
        {
            assert(b[len[i]] <= sz);
            res[b[len[i]]--] = i;
        }
        for(int i = sz; i > 0; i--)     //从len大的更新到len小的
        {
            if(t)
                endsz[link[res[i]]] += endsz[res[i]];
            else
                endsz[res[i]] = 1;
        }
        //求f
        endsz[1] = 0;
        for(int i = 1; i <= sz; i++)
            f[i] = endsz[i];
        for(int i = sz; i > 0; i--)
        {
            for(int j = 0; j < 26; j++)
                if(tran[res[i]][j])
                    f[res[i]] += f[tran[res[i]][j]];
        }
    }
        //求g
    void dfs(int u)
    {
        vis[u] == 1;
        g[u] = endsz[u];
        for(int i = 0; i < 26; i++)
        {
            if(!tran[u][i] || vis[tran[u][i]])
                continue;
            dfs(tran[u][i]);
                g[u] += g[tran[u][i]];
        }
    }
    void query(int u, ll kth)
    {
        if(kth <= endsz[u])
            return;
        kth -= endsz[u];
        for(int i = 0; i < 26; i++)
        {
            if(!tran[u][i])
                continue;
            if(f[tran[u][i]] < kth)
                kth -= f[tran[u][i]];
            else
            {
                ans.push_back(i + 'a');
                query(tran[u][i], kth);
                return;
            }
        }
    }
};

int main(void)
{
    //切勿再用scanf();
    std::ios::sync_with_stdio(false);
    std::cin.tie(0);
    cin >> s;
    n = s.size();
    cin >> t >> k;
    SAM sam(n, 30);
    sam.construct(s);
    sam.getendsz();
    sam.dfs(1);
    for(int i = 1; i <= sam.sz; i++) // f 和 g 应该相等
        assert(f[i] == g[i]);

    if(f[1] < k)
        ans = "-1";
    else
        sam.query(1, k);
    cout << ans;
    return 0;
}

2022/9/8 19:54
加载中...