PTA团体程序设计天梯赛 L3-042 污染大亨 (30/30)(满分)(C++ Java 双版本)
九条可怜今天玩一款关于环境、污染和交罚款的德式桌游。
游戏的地图是一条河的河道。因为河流会分叉,所以其河道以有根树的形式呈现。在河道边上有 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]。
- 在第一轮,所有小镇都没有受到污染,于是可怜可以任选一个小镇新建工厂。如果她选择了小镇 3 ,那么她将对这一个小镇造成污染,收到一张数值为 1 的罚单。
- 在第二轮,目前只有小镇 3 受到了污染,于是可怜可以在小镇 1,2,4 中选择下一个工厂的地址。如果她选择了小镇 2,那么她将对小镇 2,3,4 同时造成污染,收到三张数值为 2 的罚单。
- 在第三轮,目前只有小镇 1 还没有受到污染,于是可怜只能选择在这一小镇新建下一个工厂。此时,她将同时污染所有小镇,收到四张数值为 3 的罚单。
- 此时,所有小镇都已经受到了污染,游戏结束。可怜一共需要缴纳的罚款为 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 棵, ...}。
预处理:静态图构建
这题的 是随天数变化的,这意味着我们不能直接做一个简单的计数。但“从一棵
Type A 的树里删掉一个点,会变成哪些新树”这个逻辑是固定的。
-
做法:把每种
Type的树拎出来,枚举删掉它里面的每一个点 x。 -
产物:删掉 x 后,这棵树会消失,同时释放出 x 的子节点组成的树,以及 x 到原根路径上所有“旁系”子树。
-
记录:把这个转移存下来。为了 DP 跑得快,我们要把这些状态(森林的组成)通过 BFS 全部跑出来,存进一个静态邻接表里。
倒推 DP:处理变动的系数 
题目要求的是所有方案的罚款乘积之和。
设总共玩了 k 天。第 k 天肯定选的是根节点(因为选了根节点游戏就结束了)。
由于 每天不一样,我们需要外层枚举总天数 k,内层跑 DP。
dp[day][state]表示第 day 天 ,森林状态为 state 时的所有方案贡献和
转移方程:
dp[day][next_state] += dp[day-1][state] * 数量 * 结构方案数 *
这里 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 你盗我代码然后把博文收费是不是多少有点活不起了?
更多推荐




所有评论(0)