蜜汁缩点DP10pts
  • 板块P2656 采蘑菇
  • 楼主ZYK_luogu
  • 当前回复3
  • 已保存回复3
  • 发布时间2023/10/3 03:26
  • 上次更新2023/11/2 16:25:11
查看原帖
蜜汁缩点DP10pts
742157
ZYK_luogu楼主2023/10/3 03:26

Tarjan + DP, 建一个新图来用于求DP,但是只有10pts,不知道哪里错了

#include <iostream>
#include <cstdio>
#include <cstring>
using namespace std;
const int N = 8 * 1e4 + 10, M = 2 * 1e5 + 10;
struct Edge {
	int u, v, w, delta, sum;
} edges1[M], edges2[M];
int n, m, root, f[N];
int head1[N], nxt1[M], idx1 = 0;
int head2[N], nxt2[M], idx2 = 0;
int dfn[N], low[N], time_cnt = 0;
int stk[N], in_stk[N], top = 0;
int scc_siz[N], scc_sum[N], id[N], scc_cnt = 0;
inline void add1(int u, int v, int w = 0, int delta = 0, int sum = 0) {
	Edge &e = edges1[idx1];
	e.u = u, e.v = v, e.w = w, e.delta = delta, e.sum = sum;
	nxt1[idx1] = head1[u], head1[u] = idx1 ++;
}
inline void add2(int u, int v, int w) {
	Edge &e = edges2[idx2];
	e.u = u, e.v = v, e.w = w;
	nxt2[idx2] = head2[u], head2[u] = idx2 ++;
}
void tarjan(int u) {
	dfn[u] = low[u] = ++ time_cnt;
	stk[++ top] = u, in_stk[u] = 1;
	for(int i = head1[u]; ~i; i = nxt1[i]) {
		int v = edges1[i].v;
		if(!dfn[v]) {
			tarjan(v);
			low[u] = min(low[u], low[v]);
		} else if(in_stk[v])
			low[u] = min(low[u], dfn[v]);
	}
	if(low[u] == dfn[u]) {
		scc_cnt ++;
		int x;
		do {
			x = stk[top --];
			in_stk[x] = 0;
			id[x] = scc_cnt;
			scc_siz[scc_cnt] ++;
		} while (u != x);
	}
}
inline void build() {
	for(int i = 1; i <= n; i ++) 
		for(int j = head1[i]; ~j; j = nxt1[j]) {
			int u = i, v = edges1[j].v;
			if(id[u] == id[v]) 
				scc_sum[id[u]] += edges1[j].sum;
			else 
				add2(id[v], id[u], edges1[j].w);
		}
}
void dp(int u) {
	if(f[u]) return;
	int mx = 0;
	for(int i = head2[u]; ~i; i = nxt2[i]) {
		int v = edges2[i].v, w = edges2[i].w;
		dp(v);
		mx = max(mx, f[v] + w);
	}
	// printf("f[%d] = %d\n", u, f[u]);
	f[u] = scc_sum[u] + mx;
	return;
}
int main() {
	// freopen("in", "r", stdin);
	// freopen("out", "w", stdout);
	memset(head1, -1, sizeof head1);
	memset(head2, -1, sizeof head2);
	scanf("%d%d", &n, &m);
	for(int i = 1; i <= m; i ++) {
		int u, v, w, delta;
		scanf("%d %d %d 0.%d", &u, &v, &w, &delta);
		int w_ = w, tot = w;
		while(w_) w_ = w_ * delta / 10, tot += w_;
		add1(u, v, w, delta, tot);
	}
	scanf("%d", &root);
	for(int i = 1; i <= n; i ++)
		if(!dfn[i]) tarjan(i);
	build();
	dp(id[root]);
	// for(int i = 1; i <= scc_cnt; i ++)
	// 	printf("scc %d : scc_siz = %d, scc_sum = %d\n", i, scc_siz[i], scc_sum[i]);
	// for(int i = 0; i < idx2; i ++)
	// 	printf("scc %d -> scc %d = %d\n", edges2[i].u, edges2[i].v, edges2[i].w);
	printf("%d", f[id[root]]);
	// fclose(stdin);
	// fclose(stdout);
	return 0;
}
2023/10/3 03:26
加载中...