85分优美代码求助
查看原帖
85分优美代码求助
657442
wusihao1931楼主2022/10/4 12:28
#include <stack>
#include <cstdio>
#include <cstring>
#include <iostream>

using namespace std;

const int N = 200010, M = 500010 * 2;

int n, m;
int w[N];
int f[N];
int id[N];
int din[N];
int scc_cnt;
int dp[N][2];
stack<int> stk;
bool in_stk[N];
int max_cnt[N], value[N];
int low[N], dfn[N], timestamp;
int h[N], hr[N], e[M], ne[M], idx;

void add(int h[], int a, int b)
{
    e[idx] = b;
    ne[idx] = h[a];
    h[a] = idx;
    idx ++ ;
}

void tarjan(int u)
{
    dfn[u] = low[u] = ++ timestamp;
    stk.push(u), in_stk[u] = true;
    
    for (int i = h[u]; i != -1; i = ne[i])
    {
        int j = e[i];
        if (!dfn[j])
        {
            tarjan(j);
            low[u] = min(low[j], low[u]);
        }
        else if (in_stk[j]) low[u] = min(low[u], dfn[j]);
    }
    
    if (dfn[u] == low[u])
    {
        int y;
        scc_cnt ++ ;
        do
        {
            y = stk.top();
            stk.pop();
            in_stk[y] = false;
            id[y] = scc_cnt;
            value[scc_cnt] += w[y];
            max_cnt[scc_cnt] = max(max_cnt[scc_cnt], w[y]);
        } while (y != u);
    }
}

int main()
{
    cin >> n >> m;
    memset(h, -1, sizeof h);
    memset(hr, -1, sizeof hr);
    
    for (int i = 1; i <= n; i ++ ) scanf("%d", &w[i]);
    for (int i = 1; i <= m; i ++ )
    {
        int a, b;
        scanf("%d %d", &a, &b);
        add(h, a, b);
    }
    
    for (int i = 1; i <= n; i ++ )
        if (!dfn[i]) tarjan(i);
    
    for (int i = 1; i <= n; i ++ )
        for (int j = h[i]; j != -1; j = ne[j])
        {
            int k = e[j];
            int a = id[i], b = id[k];
            if (a != b) add(hr, a, b), din[b] ++ ;
        }

    for (int i = 1; i <= scc_cnt; i ++ )
    {
        if (!din[i])
        {
            dp[i][0] = value[i];
            dp[i][1] = max_cnt[i];
        }
    }

    for (int i = scc_cnt; i ; i -- )
    {
        for (int j = hr[i]; j != -1; j = ne[j])
        {
            int k = e[j];
            if (dp[k][0] < dp[i][0] + value[k])
            {
                dp[k][0] = dp[i][0] + value[k];
                dp[k][1] = max(dp[k][1], max(dp[i][1], max_cnt[k]));
            }
            
            if (dp[k][0] == dp[i][0] + value[k])
            {
                dp[k][1] = max(dp[k][1], max(max_cnt[k], dp[i][1]));
            }
        }
    }
    
    int ans = 1;
    for (int i = 2; i <= scc_cnt; i ++ )
        if (dp[ans][0] < dp[i][0] || (dp[ans][0] == dp[i][0] && dp[ans][1] < dp[i][1]))
            ans = i;
    
    cout << dp[ans][0] << " " << dp[ans][1] << endl;
    
    return 0;
}
2022/10/4 12:28
加载中...