九条可怜今天玩一款关于环境、污染和交罚款的德式桌游。

游戏的地图是一条河的河道。因为河流会分叉,所以其河道以有根树的形式呈现。在河道边上有 n 个小镇,其中编号为 1 的小镇对应了河流的源头,而第 i(i>1) 个小镇的直接上游是编号为 fi​(1≤fi​<i) 的小镇。

游戏开始时河流完全没有受到污染。而在游戏的每一轮,可怜需要选择一个还没有被污染的小镇修建一个新的工厂。在工厂开工后,它会立刻对该小镇以及其下游的所有小镇产生永久污染,也就意味着在游戏的剩余时间内,这些小镇再也不能作为工厂的地址。可怜需要不断重复这一过程,直到河道中的所有小镇都被污染为止 —— 因为每一轮一定会有一个新的小镇受到污染,所以游戏的轮数不会超过 n 轮。

既然造成了污染,可怜自然地也要缴纳罚款。在游戏的第 i 轮,如果可怜的工厂对 k 个小镇造成了污染,那么她将会收到 k 张数值为 ci​ 的罚单,其中 ci​ 是预先给定的罚款系数;在游戏结束时,可怜需要缴纳的罚款数额为所有罚单上数值的乘积注意,一些处在下游的小镇可能会在游戏的不同轮内多次被造成污染。

下面是一局游戏的例子,假设 n=4,小镇 2, 3, 4 的直接上游分别为小镇 1, 2, 2,且常数 c1​ 到 c4​ 分别为 [1,2,3,4]。

  1. 在第一轮,所有小镇都没有受到污染,于是可怜可以任选一个小镇新建工厂。如果她选择了小镇 3 ,那么她将对这一个小镇造成污染,收到一张数值为 1 的罚单。
  2. 在第二轮,目前只有小镇 3 受到了污染,于是可怜可以在小镇 1,2,4 中选择下一个工厂的地址。如果她选择了小镇 2,那么她将对小镇 2,3,4 同时造成污染,收到三张数值为 2 的罚单。
  3. 在第三轮,目前只有小镇 1 还没有受到污染,于是可怜只能选择在这一小镇新建下一个工厂。此时,她将同时污染所有小镇,收到四张数值为 3 的罚单。
  4. 此时,所有小镇都已经受到了污染,游戏结束。可怜一共需要缴纳的罚款为 1×23×34=648 元。

现在,给定游戏的地图以及常数 ci​,你需要帮助可怜计算对于所有可能的游戏情况,可怜需要缴纳的罚款总和是多少。这个答案可能很大,所以你只需输出对 998244353 取模后的结果。

输入格式:

第一行一个整数 n(1≤n≤40),表示小镇数量。

第二行 n−1 个整数,依次对应 f2​ 至 fn​,即每个非源头小镇的直接上游。输入保证 1≤fi​<i。

第三行 n 个整数,表示 c1​ 至 cn​,即每一轮中的罚款系数。输入保证 0≤ci​<106。

输出格式:

输出一行一个整数,表示对于所有可能的游戏情况,可怜需要缴纳的罚款总和对 998244353 取模后的结果。

输入样例 1:

4
1 2 2
1 2 3 4

输出样例 1:

29317

样例解释:

下表展示了所有可能的游戏情况与对应的罚款数额,其中我们用一个数组 [a1​,…,ak​] 代表一个 k 天的游戏情况,ki​ 表示第 i 天可怜选择的小镇。

游戏情况 罚款 游戏情况 罚款
[1] 14=1 [2, 1] 13×24=16
[3, 1] 11×24=16 [4, 1] 11×24=16
[3, 2, 1] 11×23×34=648 [3, 4, 1] 11×21×34=162
[4, 2, 1] 11×23×34=648 [4, 3, 1] 11×21×34=162
[3, 4, 2, 1] 11×21×33×44=13824 [4, 3, 2, 1] 11×21×33×44=13824

输入样例 2:

8
1 1 1 4 5 1 4
1 1 9 1 9 8 1 0

输出样例 2:

314366430

代码长度限制

16 KB

时间限制

5000 ms

内存限制

512 MB

栈限制

8192 KB

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

static const int MOD = 998244353;
static inline int addmod(int a,int b){ a+=b; if(a>=MOD) a-=MOD; return a; }
static inline int mulmod(long long a,long long b){ return (int)(a*b%MOD); }

static inline uint64_t splitmix64(uint64_t x){
    x += 0x9e3779b97f4a7c15ULL;
    x = (x ^ (x>>30)) * 0xbf58476d1ce4e5b9ULL;
    x = (x ^ (x>>27)) * 0x94d049bb133111ebULL;
    return x ^ (x>>31);
}


struct Key {
    uint64_t w[4];
    uint64_t mask;
};
static inline bool operator==(const Key& a,const Key& b) noexcept{
    return a.w[0]==b.w[0] && a.w[1]==b.w[1] && a.w[2]==b.w[2] && a.w[3]==b.w[3] && a.mask==b.mask;
}
static inline uint64_t key_hash(const Key& k){
    uint64_t h = 0x123456789abcdef0ULL;
    h ^= splitmix64(k.w[0] + 0x1111111111111111ULL);
    h ^= splitmix64(k.w[1] + 0x2222222222222222ULL);
    h ^= splitmix64(k.w[2] + 0x3333333333333333ULL);
    h ^= splitmix64(k.w[3] + 0x4444444444444444ULL);
    h ^= splitmix64(k.mask + 0x5555555555555555ULL);
    return h;
}


static inline int get6(const Key& k, int idx){
    int bit = idx*6, word = bit>>6, off = bit&63;
    if(off<=58) return (int)((k.w[word]>>off)&63ULL);
    int low = 64-off;
    uint64_t part1 = k.w[word]>>off;
    uint64_t part2 = k.w[word+1] & ((1ULL<<(6-low))-1ULL);
    return (int)((part1 | (part2<<low)) & 63ULL);
}

struct TSHashMap {
    size_t cap=0, mask=0;
    vector<Key> keys;
    vector<int> vals;
    vector<uint32_t> vis;
    vector<uint32_t> used;
    uint32_t stamp=1;
    size_t sz=0;

    void init_fixed(size_t cap_pow2){
        cap=cap_pow2; mask=cap-1;
        keys.resize(cap);
        vals.resize(cap);
        vis.assign(cap,0);
        used.reserve(cap/2);
        stamp=1; sz=0;
    }
    inline void clear_fast(){
        stamp++;
        if(stamp==0){ fill(vis.begin(), vis.end(), 0); stamp=1; }
        used.clear();
        sz=0;
    }
    inline void add(const Key& k,int v){
        size_t i = (size_t)key_hash(k) & mask;
        while(true){
            if(vis[i]!=stamp){
                vis[i]=stamp;
                keys[i]=k;
                vals[i]=v;
                used.push_back((uint32_t)i);
                sz++;
                return;
            }
            if(keys[i]==k){
                vals[i]=addmod(vals[i], v);
                return;
            }
            i = (i+1) & mask;
        }
    }
    template<class F>
    inline void for_each(F&& f) const{
        for(uint32_t idx: used) f(keys[idx], vals[idx]);
    }
    inline void swap_with(TSHashMap& o){
        swap(cap,o.cap); swap(mask,o.mask);
        keys.swap(o.keys); vals.swap(o.vals); vis.swap(o.vis); used.swap(o.used);
        swap(stamp,o.stamp); swap(sz,o.sz);
    }
};


struct Op {
    uint8_t idx;
    uint8_t delta;
    uint8_t word;
    uint8_t off;    
    uint8_t cross;   
    uint8_t lowBits; 
    uint64_t maskLow;
    uint64_t maskHigh;
};

static inline Op make_op(int idx, int delta){
    Op o{};
    o.idx = (uint8_t)idx;
    o.delta = (uint8_t)delta;
    int bit = idx*6;
    o.word = (uint8_t)(bit>>6);
    o.off  = (uint8_t)(bit & 63);
    if(o.off<=58){
        o.cross = 0;
        o.maskLow = 63ULL << o.off;
        o.maskHigh = 0;
        o.lowBits = 0;
    }else{
        o.cross = 1;
        o.lowBits = (uint8_t)(64 - o.off); // 1..5
        uint64_t lowMask = ((1ULL<<o.lowBits)-1ULL) << o.off;
        o.maskLow = lowMask;
        int highBits = 6 - o.lowBits;
        o.maskHigh = (1ULL<<highBits) - 1ULL;
    }
    return o;
}

static inline void apply_add(Key& k, const Op& o){
    if(!o.cross){
        uint64_t cur = (k.w[o.word] & o.maskLow) >> o.off;
        cur += o.delta;
        k.w[o.word] = (k.w[o.word] & ~o.maskLow) | ((cur & 63ULL) << o.off);
        k.mask |= (1ULL<<o.idx);
    }else{
        // split bits
        uint64_t lowPart = (k.w[o.word] & o.maskLow) >> o.off;
        uint64_t highPart = (k.w[o.word+1] & o.maskHigh);
        uint64_t cur = lowPart | (highPart << o.lowBits);
        cur += o.delta;
        // write back
        uint64_t newLow = (cur & ((1ULL<<o.lowBits)-1ULL)) << o.off;
        uint64_t newHigh = (cur >> o.lowBits) & o.maskHigh;
        k.w[o.word] = (k.w[o.word] & ~o.maskLow) | newLow;
        k.w[o.word+1] = (k.w[o.word+1] & ~o.maskHigh) | newHigh;
        k.mask |= (1ULL<<o.idx);
    }
}

static inline void apply_dec_known(Key& k, const Op& op_for_idx, int newVal){
    const Op& o = op_for_idx;
    if(!o.cross){
        k.w[o.word] = (k.w[o.word] & ~o.maskLow) | ((uint64_t)newVal << o.off);
    }else{
        uint64_t newLow = ((uint64_t)newVal & ((1ULL<<o.lowBits)-1ULL)) << o.off;
        uint64_t newHigh = ((uint64_t)newVal >> o.lowBits) & o.maskHigh;
        k.w[o.word] = (k.w[o.word] & ~o.maskLow) | newLow;
        k.w[o.word+1] = (k.w[o.word+1] & ~o.maskHigh) | newHigh;
    }
    uint64_t bit = 1ULL<<o.idx;
    if(newVal==0) k.mask &= ~bit;
    else k.mask |= bit;
}

struct TransItem {
    uint32_t off;
    uint16_t len;
    int weight;
};

int main(){
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    int n; cin>>n;
    vector<int> parent(n+1,0);
    vector<vector<int>> ch(n+1);
    for(int i=2;i<=n;i++){
        cin>>parent[i];
        ch[parent[i]].push_back(i);
    }
    vector<int> c(n+1);
    for(int i=1;i<=n;i++) cin>>c[i];

    vector<int> sub(n+1,1);
    function<void(int)> dfs_sz = [&](int u){
        sub[u]=1;
        for(int v: ch[u]){ dfs_sz(v); sub[u]+=sub[v]; }
    };
    dfs_sz(1);

    vector<int> type(n+1,-1);
    map<vector<int>,int> sig2id;
    vector<int> rep_root;
    function<int(int)> dfs_type = [&](int u)->int{
        vector<int> sig; sig.reserve(ch[u].size());
        for(int v: ch[u]) sig.push_back(dfs_type(v));
        sort(sig.begin(), sig.end());
        auto it = sig2id.find(sig);
        int id;
        if(it==sig2id.end()){
            id=(int)sig2id.size();
            sig2id[sig]=id;
            rep_root.push_back(u);
        }else id=it->second;
        type[u]=id;
        return id;
    };
    dfs_type(1);
    int T = (int)sig2id.size();

    // powC[day][s]
    vector<vector<int>> powC(n+1, vector<int>(n+1,1));
    for(int day=1; day<=n; day++){
        long long base = c[day] % MOD;
        for(int s=1; s<=n; s++){
            powC[day][s] = (int)(1LL*powC[day][s-1]*base%MOD);
        }
    }

    vector<Op> op_idx(T);
    for(int i=0;i<T;i++) op_idx[i] = make_op(i, 0);
    struct AddKey {
        uint64_t w[4];
    };
    struct AddKeyHash {
        size_t operator()(AddKey const& a) const noexcept{
            uint64_t h=0xabcdef0123456789ULL;
            h ^= splitmix64(a.w[0]+0x1111111111111111ULL);
            h ^= splitmix64(a.w[1]+0x2222222222222222ULL);
            h ^= splitmix64(a.w[2]+0x3333333333333333ULL);
            h ^= splitmix64(a.w[3]+0x4444444444444444ULL);
            return (size_t)h;
        }
    };
    struct AddKeyEq {
        bool operator()(AddKey const& a, AddKey const& b) const noexcept{
            return a.w[0]==b.w[0]&&a.w[1]==b.w[1]&&a.w[2]==b.w[2]&&a.w[3]==b.w[3];
        }
    };
    auto addkey_get6 = [&](const AddKey& k, int idx)->int{
        int bit = idx*6, word = bit>>6, off = bit&63;
        if(off<=58) return (int)((k.w[word]>>off)&63ULL);
        int low = 64-off;
        uint64_t part1 = k.w[word]>>off;
        uint64_t part2 = k.w[word+1] & ((1ULL<<(6-low))-1ULL);
        return (int)((part1 | (part2<<low)) & 63ULL);
    };
    auto addkey_inc = [&](AddKey& k, int idx, int delta){
        int v = addkey_get6(k, idx) + delta;
        int bit = idx*6, word = bit>>6, off = bit&63;
        if(off<=58){
            uint64_t mask = 63ULL<<off;
            k.w[word] = (k.w[word] & ~mask) | ((uint64_t)v<<off);
        }else{
            int low = 64-off;
            uint64_t m1 = (1ULL<<low)-1ULL;
            k.w[word] &= ~(m1<<off);
            k.w[word] |= (uint64_t(v)&m1)<<off;
            int high = 6-low;
            uint64_t m2 = (1ULL<<high)-1ULL;
            k.w[word+1] &= ~m2;
            k.w[word+1] |= ((uint64_t(v)>>low)&m2);
        }
    };

    vector<vector<Op>> addops_pool(T);             
    vector<vector<uint32_t>> off(T, vector<uint32_t>(n+2,0));
    vector<vector<TransItem>> items(T);

    for(int t=0;t<T;t++){
        int r = rep_root[t];

        vector<int> nodes;
        function<void(int)> collect = [&](int u){
            nodes.push_back(u);
            for(int v: ch[u]) collect(v);
        };
        collect(r);

        unordered_map<AddKey, vector<int>, AddKeyHash, AddKeyEq> mp;
        mp.reserve(256);

        for(int x: nodes){
            AddKey ak{}; ak.w[0]=ak.w[1]=ak.w[2]=ak.w[3]=0;

            // children(x)
            for(int v: ch[x]) addkey_inc(ak, type[v], 1);

            // path to root, add siblings at each level
            int cur=x;
            while(cur!=r){
                int p = parent[cur];
                for(int s: ch[p]) if(s!=cur) addkey_inc(ak, type[s], 1);
                cur=p;
            }

            auto it = mp.find(ak);
            if(it==mp.end()){
                vector<int> coef(n+1,0);
                coef[sub[x]] = 1;
                mp.emplace(ak, std::move(coef));
            }else{
                it->second[sub[x]] = addmod(it->second[sub[x]], 1);
            }
        }

        vector<pair<AddKey, vector<int>>> vec;
        vec.reserve(mp.size());
        for(auto &kv: mp) vec.push_back({kv.first, std::move(kv.second)});

        addops_pool[t].clear();
        items[t].clear();
        fill(off[t].begin(), off[t].end(), 0);

        struct DeltaOpsRef { uint32_t op_off; uint16_t op_len; vector<int> coef; };
        vector<DeltaOpsRef> deltas;
        deltas.reserve(vec.size());

        for(auto &kv: vec){
            const AddKey& ak = kv.first;
            const vector<int>& coef = kv.second;

            uint32_t op_off = (uint32_t)addops_pool[t].size();
            uint16_t op_len = 0;
            for(int i=0;i<T;i++){
                int v = addkey_get6(ak, i);
                if(v){
                    addops_pool[t].push_back(make_op(i, v));
                    op_len++;
                }
            }
            deltas.push_back(DeltaOpsRef{op_off, op_len, coef});
        }
        
        for(int day=1; day<=n; day++){
            off[t][day] = (uint32_t)items[t].size();
            for(auto &d : deltas){
                long long sum=0;
                const auto &coef = d.coef;
                for(int s=1;s<=n;s++){
                    if(coef[s]) sum += 1LL * coef[s] * powC[day][s] % MOD;
                }
                int w = (int)(sum % MOD);
                if(!w) continue;
                items[t].push_back(TransItem{d.op_off, d.op_len, w});
            }
            off[t][day+1] = (uint32_t)items[t].size();
        }
    }

    Key init{}; init.w[0]=init.w[1]=init.w[2]=init.w[3]=0; init.mask=0;
    for(int v: ch[1]){
        int tp = type[v];
        int old = get6(init, tp);
        Op o = make_op(tp, 1);
        apply_add(init, o);
        (void)old;
    }

    const size_t CAP = (1u<<21); 
    TSHashMap dp, ndp;
    dp.init_fixed(CAP);
    ndp.init_fixed(CAP);

    long long ans=0;

    for(int k=1;k<=n;k++){
        int last_factor = powC[k][n];
        if(k==1){
            ans = (ans + last_factor) % MOD;
            continue;
        }

        dp.clear_fast();
        dp.add(init, 1);

        for(int day=k-1; day>=1; day--){
            ndp.clear_fast();

            dp.for_each([&](const Key& st, int curv){
                uint64_t m = st.mask;
                while(m){
                    int t = __builtin_ctzll(m);
                    m &= (m-1);

                    int q = get6(st, t);
                    int baseMul = mulmod(curv, q % MOD);

                    uint32_t L = off[t][day];
                    uint32_t R = off[t][day+1];

                    // prepare base state with count[t] decremented once
                    Key base = st;
                    apply_dec_known(base, op_idx[t], q-1);

                    const auto &vec = items[t];
                    const Op* pool = addops_pool[t].data();

                    for(uint32_t p=L; p<R; p++){
                        const TransItem &it = vec[p];
                        int w = it.weight;
                        int val = mulmod(baseMul, w);

                        Key nxt = base; // 40B copy
                        const Op* ops = pool + it.off;
                        for(int i=0;i<it.len;i++){
                            apply_add(nxt, ops[i]);
                        }
                        ndp.add(nxt, val);
                    }
                }
            });

            dp.swap_with(ndp);
        }

        long long sum_dp=0;
        dp.for_each([&](const Key&, int v){
            sum_dp += v;
            if(sum_dp>=MOD) sum_dp-=MOD;
        });

        ans = (ans + 1LL * last_factor % MOD * (sum_dp % MOD)) % MOD;
    }

    cout << ans % MOD << "\n";
    return 0;
}

改吐了 佬就是佬 出的题非我辈能解的

2026.2.2 更新AC算法,仔细拜读了下吉佬的解析,发现要把反向还原纠正为 正向消除剩余子树,并且把最后选择根节点清空全场作为单独的结算步骤,避免逻辑上的死胡同。正赛的环境下肯定想不出正解的……

应该是全网第一篇AC的题解吧

这题是典型的状态压缩 DP 的极限变种,n=40 这个范围非常阴险,它大到让你没法直接状压,小到让你觉得一定存在某种搜索或者缩减状态的方法。

这道题最核心的突破点在于:尽管节点有 40 个,但在森林里长得一模一样的子树是可以合并的。

1. 状态的本质:从 节点 到 子树类型

如果直接记录哪些点被污染了,那是 2^40。但仔细想,第 i 天我们要从剩下的森林里选一棵树的一个点。如果森林里有两棵一模一样的树,选哪棵的结果是一样的。

所以,我们先做树哈希。把所有同构的子树映射成同一个 Type。n=40 时,不同的子树类型 T 其实很少(通常几十个)。

现在的状态就变成了:{Type 0 有 a 棵, Type 1 有 b 棵, ...}

预处理:静态图构建

这题的 c_{i} 是随天数变化的,这意味着我们不能直接做一个简单的计数。但“从一棵 Type A 的树里删掉一个点,会变成哪些新树”这个逻辑是固定的。

  • 做法:把每种 Type 的树拎出来,枚举删掉它里面的每一个点 x。

  • 产物:删掉 x 后,这棵树会消失,同时释放出 x 的子节点组成的树,以及 x 到原根路径上所有“旁系”子树。

  • 记录:把这个转移存下来。为了 DP 跑得快,我们要把这些状态(森林的组成)通过 BFS 全部跑出来,存进一个静态邻接表里。

倒推 DP:处理变动的系数 c_{i}

题目要求的是所有方案的罚款乘积之和。

设总共玩了 k 天。第 k 天肯定选的是根节点(因为选了根节点游戏就结束了)。

由于 c_{i} 每天不一样,我们需要外层枚举总天数 k,内层跑 DP。

dp[day][state]表示第 day 天 ,森林状态为 state 时的所有方案贡献和

转移方程

dp[day][next_state] += dp[day-1][state] * 数量 * 结构方案数 * c_{day}^{size}

这里 size 是那一次操作瞬间污染的节点数。

/**
 * L3-042 污染大亨
 * * 核心思路:
 * 1. 树哈希缩点 + 状态压缩 (Key)。
 * 2. 静态邻接表存储转移 (BFS)。
 * 3. 逆向 DP (O(N*|E|)):一次计算所有天数 k 的结果,消除外层循环冗余。
 */

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

// --- 常数与模运算 ---
static const int MOD = 998244353;
static inline int addmod(int a, int b) { a += b; if (a >= MOD) a -= MOD; return a; }
static inline int mulmod(long long a, long long b) { return (int)(a * b % MOD); }

// --- 状态 Key 定义 ---
struct Key {
    uint64_t w[4];  // 存储每种类型的数量
    uint64_t mask;  // 标记哪些类型数量不为 0
    inline bool operator==(const Key& o) const noexcept {
        return w[0] == o.w[0] && w[1] == o.w[1] && w[2] == o.w[2] && w[3] == o.w[3] && mask == o.mask;
    }
};

static inline uint64_t key_hash(const Key& k) {
    uint64_t h = k.mask * 0x9e3779b97f4a7c15ULL;
    h = (h ^ k.w[0]) * 0xbf58476d1ce4e5b9ULL;
    h = (h ^ k.w[1]) * 0x94d049bb133111ebULL;
    h = (h ^ k.w[2]) * 0xbf58476d1ce4e5b9ULL;
    h = (h ^ k.w[3]) * 0x94d049bb133111ebULL;
    return h;
}

// 获取第 idx 种树类型的当前数量
static inline int get6(const Key& k, int idx) {
    int bit = idx * 6;
    int word = bit >> 6;
    int off = bit & 63;
    if (off <= 58) return (int)((k.w[word] >> off) & 63ULL);
    int low = 64 - off;
    uint64_t p1 = k.w[word] >> off;
    uint64_t p2 = k.w[word + 1] & ((1ULL << (6 - low)) - 1ULL);
    return (int)((p1 | (p2 << low)) & 63ULL);
}

// 位运算操作封装
struct Op {
    uint8_t word, off, cross, lowBits;
    uint64_t maskLow, maskHigh;
    int idx, delta;
};

static inline Op make_op(int idx, int delta) {
    Op o; o.idx = idx; o.delta = delta;
    int bit = idx * 6;
    o.word = (uint8_t)(bit >> 6);
    o.off = (uint8_t)(bit & 63);
    if (o.off <= 58) {
        o.cross = 0; o.maskLow = 63ULL << o.off;
        o.lowBits = 0; o.maskHigh = 0;
    } else {
        o.cross = 1; o.lowBits = (uint8_t)(64 - o.off);
        o.maskLow = ((1ULL << o.lowBits) - 1ULL) << o.off;
        o.maskHigh = (1ULL << (6 - o.lowBits)) - 1ULL;
    }
    return o;
}

static inline void apply_add(Key& k, const Op& o) {
    if (!o.cross) {
        uint64_t val = (k.w[o.word] >> o.off) & 63ULL;
        val += o.delta;
        k.w[o.word] = (k.w[o.word] & ~o.maskLow) | (val << o.off);
    } else {
        uint64_t v1 = (k.w[o.word] >> o.off) & ((1ULL << o.lowBits) - 1ULL);
        uint64_t v2 = k.w[o.word + 1] & o.maskHigh;
        uint64_t val = v1 | (v2 << o.lowBits);
        val += o.delta;
        k.w[o.word] = (k.w[o.word] & ~o.maskLow) | ((val & ((1ULL << o.lowBits) - 1ULL)) << o.off);
        k.w[o.word + 1] = (k.w[o.word + 1] & ~o.maskHigh) | ((val >> o.lowBits) & o.maskHigh);
    }
    k.mask |= (1ULL << o.idx);
}

static inline void apply_dec_set_mask(Key& k, const Op& o, int resultVal) {
    if (!o.cross) {
        k.w[o.word] = (k.w[o.word] & ~o.maskLow) | ((uint64_t)resultVal << o.off);
    } else {
        k.w[o.word] = (k.w[o.word] & ~o.maskLow) | ((uint64_t)(resultVal & ((1ULL << o.lowBits) - 1ULL)) << o.off);
        k.w[o.word + 1] = (k.w[o.word + 1] & ~o.maskHigh) | ((uint64_t)(resultVal >> o.lowBits) & o.maskHigh);
    }
    if (resultVal == 0) k.mask &= ~(1ULL << o.idx);
    else k.mask |= (1ULL << o.idx);
}

// 结构性变化
struct StructuralDelta {
    vector<int> coef; 
    vector<Op> ops;   
};

// 静态图边
struct Edge {
    uint32_t to;    
    uint16_t size;  
    uint16_t count; 
    uint8_t type;   // 记录被删除的树类型,用于后续查找 qty
};

// 状态 ID 映射
struct StateMap {
    static const size_t CAP = 1 << 20; 
    static const size_t MASK = CAP - 1;
    struct Entry { Key k; int id; };
    vector<Entry> table;
    vector<bool> occupied;
    
    StateMap() : table(CAP), occupied(CAP, false) {}

    int get_id(const Key& k, int& next_id) {
        size_t idx = key_hash(k) & MASK;
        while (true) {
            if (!occupied[idx]) {
                occupied[idx] = true;
                table[idx] = {k, next_id};
                return next_id++;
            }
            if (table[idx].k == k) return table[idx].id;
            idx = (idx + 1) & MASK;
        }
    }
};

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    int n;
    if (!(cin >> n)) return 0;

    vector<int> parent(n + 1);
    vector<vector<int>> ch(n + 1);
    for (int i = 2; i <= n; i++) {
        cin >> parent[i];
        ch[parent[i]].push_back(i);
    }
    vector<int> c(n + 1);
    for (int i = 1; i <= n; i++) cin >> c[i];

    // 1. DFS 子树大小
    vector<int> sub(n + 1);
    function<void(int)> dfs_sz = [&](int u) {
        sub[u] = 1;
        for (int v : ch[u]) {
            dfs_sz(v);
            sub[u] += sub[v];
        }
    };
    dfs_sz(1);

    // 2. 树同构
    vector<int> type(n + 1);
    vector<int> rep_root; 
    map<vector<int>, int> sig2id;
    function<int(int)> dfs_type = [&](int u) -> int {
        vector<int> sig;
        for (int v : ch[u]) sig.push_back(dfs_type(v));
        sort(sig.begin(), sig.end());
        auto it = sig2id.find(sig);
        if (it == sig2id.end()) {
            int id = (int)sig2id.size();
            sig2id[sig] = id;
            rep_root.push_back(u);
            type[u] = id;
            return id;
        }
        type[u] = it->second;
        return it->second;
    };
    dfs_type(1);
    int T = (int)sig2id.size();

    // 3. 预处理系数幂次
    vector<vector<int>> powC(n + 1, vector<int>(n + 1));
    for (int d = 1; d <= n; d++) {
        powC[d][0] = 1;
        for (int s = 1; s <= n; s++) 
            powC[d][s] = mulmod(powC[d][s - 1], c[d]);
    }

    // 4. 预处理转移逻辑
    vector<Op> type_dec_ops(T);
    for(int i=0; i<T; i++) type_dec_ops[i] = make_op(i, 0);

    vector<vector<StructuralDelta>> transitions(T);
    for (int t = 0; t < T; t++) {
        int r = rep_root[t];
        vector<int> nodes;
        function<void(int)> collect = [&](int u) {
            nodes.push_back(u);
            for (int v : ch[u]) collect(v);
        };
        collect(r);

        map<vector<int>, vector<int>> distinct_outcomes;
        for (int x : nodes) {
            vector<int> delta_counts(T, 0);
            for (int v : ch[x]) delta_counts[type[v]]++;
            int cur = x;
            while(cur != r) {
                int p = parent[cur];
                for (int s : ch[p]) if (s != cur) delta_counts[type[s]]++;
                cur = p;
            }
            if (distinct_outcomes.find(delta_counts) == distinct_outcomes.end())
                distinct_outcomes[delta_counts] = vector<int>(n + 1, 0);
            distinct_outcomes[delta_counts][sub[x]]++;
        }

        for (auto& kv : distinct_outcomes) {
            StructuralDelta sd;
            sd.coef = kv.second;
            for (int i = 0; i < T; i++) {
                if (kv.first[i] > 0) sd.ops.push_back(make_op(i, kv.first[i]));
            }
            transitions[t].push_back(sd);
        }
    }

    // 5. 建立静态状态图 (BFS)
    Key init_key = {0, 0, 0, 0, 0};
    for (int v : ch[1]) apply_add(init_key, make_op(type[v], 1));

    StateMap state_map;
    int num_states = 0;
    int start_node = state_map.get_id(init_key, num_states);

    vector<Key> id_to_key; 
    id_to_key.reserve(700000);
    id_to_key.push_back(init_key);

    vector<vector<Edge>> adj;
    adj.reserve(700000);

    vector<int> bfs_q; bfs_q.reserve(700000);
    bfs_q.push_back(start_node);

    int head = 0;
    while(head < (int)bfs_q.size()){
        int u = bfs_q[head++];
        Key st = id_to_key[u];
        
        if ((int)adj.size() <= u) adj.resize(u + 1);

        uint64_t m = st.mask;
        while (m) {
            int t = __builtin_ctzll(m);
            m &= (m - 1);
            int qty = get6(st, t);
            
            Key base = st;
            apply_dec_set_mask(base, type_dec_ops[t], qty - 1);

            for (const auto& trans : transitions[t]) {
                Key nxt = base;
                for (const auto& op : trans.ops) apply_add(nxt, op);
                
                int v_id = state_map.get_id(nxt, num_states);
                if (v_id == (int)id_to_key.size()) {
                    id_to_key.push_back(nxt);
                    bfs_q.push_back(v_id);
                }
                
                for (int sz = 1; sz <= n; sz++) {
                    if (trans.coef[sz] > 0) {
                        adj[u].push_back({(uint32_t)v_id, (uint16_t)sz, (uint16_t)trans.coef[sz], (uint8_t)t});
                    }
                }
            }
        }
    }

    // 6. DP 阶段:逆向思维优化 (O(N * |E|))
    // 收集非叶子节点以加速 DP
    vector<int> non_leaf_states;
    non_leaf_states.reserve(num_states);
    for (int i = 0; i < num_states; i++) {
        if (!adj[i].empty()) non_leaf_states.push_back(i);
    }

    // dp[u] 表示从状态 u 开始,剩余步骤为 0 时的“路径权值积”(初始为 1)
    // 实际上代表:如果在该状态结束,方案数为 1
    vector<int> dp(num_states, 1);
    vector<int> ndp(num_states, 0);

    vector<int> ways_k(n + 1, 0);
    ways_k[1] = 1; 

    // t 代表倒数第 t 天 (从 1 到 n-1)
    for (int t = 1; t < n; t++) {
        fill(ndp.begin(), ndp.end(), 0);
        
        // 并行计算优化:只遍历非叶子节点
        for (int u : non_leaf_states) {
            const Key& k_u = id_to_key[u]; // 获取当前状态 Key
            long long sum_val = 0;
            
            for (const auto& e : adj[u]) {
                // 获取当前状态下,该类型树的数量 qty
                int qty = get6(k_u, e.type);

                // dp[v] * count_structural * qty * cost_time
                long long term = mulmod(dp[e.to], e.count);
                term = mulmod(term, qty); // 乘上数量系数
                term = mulmod(term, powC[t][e.size]);
                
                sum_val = addmod((int)sum_val, (int)term);
            }
            ndp[u] = (int)sum_val;
        }

        dp = ndp; 
        ways_k[t + 1] = dp[start_node];
    }

    long long total_ans = 0;
    for (int k = 1; k <= n; k++) {
        long long root_fine = powC[k][n]; 
        long long ways = ways_k[k];
        total_ans = (total_ans + mulmod(root_fine, ways)) % MOD;
    }

    cout << total_ans << "\n";
    return 0;
}

Java版本:

/**
 * L3-042 污染大亨 Java版
 */

import java.io.*;
import java.util.*;

public class Main {
    // --- 常数与模运算 ---
    static final int MOD = 998244353;

    static int addmod(int a, int b) {
        a += b;
        if (a >= MOD) a -= MOD;
        return a;
    }

    static int mulmod(long a, long b) {
        return (int) ((a * b) % MOD);
    }

    // --- 状态 Key 定义 ---
    static class Key implements Cloneable {
        long[] w = new long[4]; // 存储每种类型的数量
        long mask;              // 标记哪些类型数量不为 0

        // 深度复制
        public Key copy() {
            Key k = new Key();
            System.arraycopy(this.w, 0, k.w, 0, 4);
            k.mask = this.mask;
            return k;
        }

        @Override
        public boolean equals(Object o) {
            if (this == o) return true;
            if (o == null || getClass() != o.getClass()) return false;
            Key other = (Key) o;
            return mask == other.mask &&
                   w[0] == other.w[0] &&
                   w[1] == other.w[1] &&
                   w[2] == other.w[2] &&
                   w[3] == other.w[3];
        }

        // Hash逻辑
        public long keyHash() {
            long h = mask * 0x9e3779b97f4a7c15L;
            h = (h ^ w[0]) * 0xbf58476d1ce4e5b9L;
            h = (h ^ w[1]) * 0x94d049bb133111ebL;
            h = (h ^ w[2]) * 0xbf58476d1ce4e5b9L;
            h = (h ^ w[3]) * 0x94d049bb133111ebL;
            return h;
        }
    }

    // 获取第 idx 种树类型的当前数量
    static int get6(Key k, int idx) {
        int bit = idx * 6;
        int word = bit >> 6;
        int off = bit & 63;
        
        if (off <= 58) return (int) ((k.w[word] >>> off) & 63L);
        
        int low = 64 - off;
        long p1 = k.w[word] >>> off;
        long p2 = k.w[word + 1] & ((1L << (6 - low)) - 1L);
        return (int) ((p1 | (p2 << low)) & 63L);
    }

    // 位运算操作封装
    static class Op {
        int word, off, cross, lowBits;
        long maskLow, maskHigh;
        int idx, delta;
    }

    static Op make_op(int idx, int delta) {
        Op o = new Op();
        o.idx = idx;
        o.delta = delta;
        int bit = idx * 6;
        o.word = bit >> 6;
        o.off = bit & 63;
        
        if (o.off <= 58) {
            o.cross = 0;
            o.maskLow = 63L << o.off;
            o.lowBits = 0;
            o.maskHigh = 0;
        } else {
            o.cross = 1;
            o.lowBits = 64 - o.off;
            o.maskLow = ((1L << o.lowBits) - 1L) << o.off;
            o.maskHigh = (1L << (6 - o.lowBits)) - 1L;
        }
        return o;
    }

    static void apply_add(Key k, Op o) {
        if (o.cross == 0) {
            long val = (k.w[o.word] >>> o.off) & 63L;
            val += o.delta;
            k.w[o.word] = (k.w[o.word] & ~o.maskLow) | (val << o.off);
        } else {
            long v1 = (k.w[o.word] >>> o.off) & ((1L << o.lowBits) - 1L);
            long v2 = k.w[o.word + 1] & o.maskHigh;
            long val = v1 | (v2 << o.lowBits);
            val += o.delta;
            k.w[o.word] = (k.w[o.word] & ~o.maskLow) | ((val & ((1L << o.lowBits) - 1L)) << o.off);
            k.w[o.word + 1] = (k.w[o.word + 1] & ~o.maskHigh) | ((val >>> o.lowBits) & o.maskHigh);
        }
        k.mask |= (1L << o.idx);
    }

    static void apply_dec_set_mask(Key k, Op o, int resultVal) {
        if (o.cross == 0) {
            k.w[o.word] = (k.w[o.word] & ~o.maskLow) | ((long) resultVal << o.off);
        } else {
            k.w[o.word] = (k.w[o.word] & ~o.maskLow) | ((long) (resultVal & ((1L << o.lowBits) - 1L)) << o.off);
            k.w[o.word + 1] = (k.w[o.word + 1] & ~o.maskHigh) | ((long) (resultVal >>> o.lowBits) & o.maskHigh);
        }
        
        if (resultVal == 0) k.mask &= ~(1L << o.idx);
        else k.mask |= (1L << o.idx);
    }

    // 结构性变化
    static class StructuralDelta {
        int[] coef; 
        List<Op> ops = new ArrayList<>();
    }

    // 静态图边
    static class Edge {
        int to;
        int size;
        int count;
        int type;

        public Edge(int to, int size, int count, int type) {
            this.to = to;
            this.size = size;
            this.count = count;
            this.type = type;
        }
    }

    // 状态 ID 映射
    static class StateMap {
        static final int CAP = 1 << 20;
        static final int MASK = CAP - 1;
        
        static class Entry {
            Key k;
            int id;
            Entry(Key k, int id) { this.k = k; this.id = id; }
        }
        
        Entry[] table = new Entry[CAP];

        int get_id(Key k, int[] next_id_ref) {
            int idx = (int)(k.keyHash() & MASK);
            while (true) {
                if (table[idx] == null) {
                    table[idx] = new Entry(k, next_id_ref[0]);
                    return next_id_ref[0]++;
                }
                if (table[idx].k.equals(k)) {
                    return table[idx].id;
                }
                idx = (idx + 1) & MASK;
            }
        }
    }

    // 全局变量
    static int n;
    static int[] parent;
    static ArrayList<Integer>[] ch;
    static int[] c;
    static int[] sub;
    static int[] type;
    static ArrayList<Integer> rep_root = new ArrayList<>();
    static Map<ArrayList<Integer>, Integer> sig2id = new HashMap<>();

    public static void main(String[] args) throws IOException {
        // 快读
        BufferedReader br = new BufferedReader(new InputStreamReader(System.in));
        StreamTokenizer st = new StreamTokenizer(br);
        
        if (st.nextToken() != StreamTokenizer.TT_EOF) {
            n = (int) st.nval;
        } else {
            return;
        }

        parent = new int[n + 1];
        ch = new ArrayList[n + 1];
        for (int i = 0; i <= n; i++) ch[i] = new ArrayList<>();
        
        for (int i = 2; i <= n; i++) {
            st.nextToken();
            parent[i] = (int) st.nval;
            ch[parent[i]].add(i);
        }
        
        c = new int[n + 1];
        for (int i = 1; i <= n; i++) {
            st.nextToken();
            c[i] = (int) st.nval;
        }

        // 1. DFS 子树大小
        sub = new int[n + 1];
        dfs_sz(1);

        // 2. 树同构
        type = new int[n + 1];
        dfs_type(1);
        int T = sig2id.size();

        // 3. 预处理系数幂次
        int[][] powC = new int[n + 1][n + 1];
        for (int d = 1; d <= n; d++) {
            powC[d][0] = 1;
            for (int s = 1; s <= n; s++) {
                powC[d][s] = mulmod(powC[d][s - 1], c[d]);
            }
        }

        // 4. 预处理转移逻辑
        Op[] type_dec_ops = new Op[T];
        for (int i = 0; i < T; i++) type_dec_ops[i] = make_op(i, 0);

        List<List<StructuralDelta>> transitions = new ArrayList<>();
        for (int i = 0; i < T; i++) transitions.add(new ArrayList<>());

        for (int t = 0; t < T; t++) {
            int r = rep_root.get(t);
            List<Integer> nodes = new ArrayList<>();
            collect_nodes(r, nodes);

            Map<ArrayList<Integer>, int[]> distinct_outcomes = new HashMap<>();
            
            for (int x : nodes) {
                ArrayList<Integer> delta_counts = new ArrayList<>(Collections.nCopies(T, 0));
                
                for (int v : ch[x]) {
                    int typeV = type[v];
                    delta_counts.set(typeV, delta_counts.get(typeV) + 1);
                }
                
                int cur = x;
                while (cur != r) {
                    int p = parent[cur];
                    for (int s : ch[p]) {
                        if (s != cur) {
                            int typeS = type[s];
                            delta_counts.set(typeS, delta_counts.get(typeS) + 1);
                        }
                    }
                    cur = p;
                }
                
                if (!distinct_outcomes.containsKey(delta_counts)) {
                    distinct_outcomes.put(delta_counts, new int[n + 1]);
                }
                distinct_outcomes.get(delta_counts)[sub[x]]++;
            }

            for (Map.Entry<ArrayList<Integer>, int[]> entry : distinct_outcomes.entrySet()) {
                StructuralDelta sd = new StructuralDelta();
                sd.coef = entry.getValue();
                ArrayList<Integer> counts = entry.getKey();
                for (int i = 0; i < T; i++) {
                    int val = counts.get(i);
                    if (val > 0) {
                        sd.ops.add(make_op(i, val));
                    }
                }
                transitions.get(t).add(sd);
            }
        }

        // 5. 建立静态状态图 (BFS)
        Key init_key = new Key();
        for (int v : ch[1]) {
            apply_add(init_key, make_op(type[v], 1));
        }

        StateMap state_map = new StateMap();
        int[] num_states_ref = {0};
        int start_node = state_map.get_id(init_key, num_states_ref);

        List<Key> id_to_key = new ArrayList<>();
        id_to_key.add(init_key); 

        List<List<Edge>> adj = new ArrayList<>();
        adj.add(new ArrayList<>()); 

        int[] bfs_q = new int[700000]; 
        int head = 0, tail = 0;
        bfs_q[tail++] = start_node;

        while (head < tail) {
            int u = bfs_q[head++];
            Key currentKey = id_to_key.get(u);

            while (adj.size() <= u) adj.add(new ArrayList<>());

            long m = currentKey.mask;
            while (m != 0) {
                int t = Long.numberOfTrailingZeros(m);
                m &= ~(1L << t); 
                
                int qty = get6(currentKey, t);

                Key base = currentKey.copy();
                apply_dec_set_mask(base, type_dec_ops[t], qty - 1);

                for (StructuralDelta trans : transitions.get(t)) {
                    Key nxt = base.copy();
                    for (Op op : trans.ops) {
                        apply_add(nxt, op);
                    }

                    int v_id = state_map.get_id(nxt, num_states_ref);
                    if (v_id == id_to_key.size()) {
                        id_to_key.add(nxt);
                        adj.add(new ArrayList<>());
                        bfs_q[tail++] = v_id;
                    }

                    for (int sz = 1; sz <= n; sz++) {
                        if (trans.coef[sz] > 0) {
                            adj.get(u).add(new Edge(v_id, sz, trans.coef[sz], t));
                        }
                    }
                }
            }
        }

        int num_states = num_states_ref[0];

        // 6. DP 阶段
        int[] non_leaf_states = new int[num_states];
        int non_leaf_count = 0;
        for (int i = 0; i < num_states; i++) {
            if (i < adj.size() && !adj.get(i).isEmpty()) {
                non_leaf_states[non_leaf_count++] = i;
            }
        }

        int[] dp = new int[num_states];
        Arrays.fill(dp, 1);
        int[] ndp = new int[num_states];
        
        int[] ways_k = new int[n + 1];
        ways_k[1] = 1;

        for (int t = 1; t < n; t++) {
            Arrays.fill(ndp, 0);

            for (int k = 0; k < non_leaf_count; k++) {
                int u = non_leaf_states[k];
                Key k_u = id_to_key.get(u);
                long sum_val = 0;

                for (Edge e : adj.get(u)) {
                    int qty = get6(k_u, e.type);

                    long term = mulmod(dp[e.to], e.count);
                    term = mulmod(term, qty);
                    term = mulmod(term, powC[t][e.size]);

                    sum_val = addmod((int) sum_val, (int) term);
                }
                ndp[u] = (int) sum_val;
            }

            System.arraycopy(ndp, 0, dp, 0, num_states);
            ways_k[t + 1] = dp[start_node];
        }

        long total_ans = 0;
        for (int k = 1; k <= n; k++) {
            long root_fine = powC[k][n];
            long ways = ways_k[k];
            total_ans = (total_ans + mulmod(root_fine, ways)) % MOD;
        }

        System.out.println(total_ans);
    }

    // --- 辅助 DFS 函数 ---

    static void dfs_sz(int u) {
        sub[u] = 1;
        for (int v : ch[u]) {
            dfs_sz(v);
            sub[u] += sub[v];
        }
    }

    static int dfs_type(int u) {
        ArrayList<Integer> sig = new ArrayList<>();
        for (int v : ch[u]) {
            sig.add(dfs_type(v));
        }
        Collections.sort(sig);
        
        if (!sig2id.containsKey(sig)) {
            int id = sig2id.size();
            sig2id.put(sig, id);
            rep_root.add(u);
            type[u] = id;
            return id;
        } else {
            int id = sig2id.get(sig);
            type[u] = id;
            return id;
        }
    }

    static void collect_nodes(int u, List<Integer> nodes) {
        nodes.add(u);
        for (int v : ch[u]) {
            collect_nodes(v, nodes);
        }
    }
}

顺便 问一下 @Pretty Boy Fox 你盗我代码然后把博文收费是不是多少有点活不起了?

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐