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;
}