re#3求调
查看原帖
re#3求调
968912
ajsdlkasd楼主2023/7/30 13:34
#include <bits/stdc++.h>

const int N = 800010,M = 2*N;

using namespace std;

typedef long long LL;

int h[N],hs[N],e[M],ne[M],w[N],f[N];
int dfn[N],low[N],val[N],stk[N],id[N];
bool is_bar[N],in_stk[N];
int n,m,stkTop = 0,idx = 0,timestamp = 0,scc_cnt = 0;

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

void tarjan(int u){
    dfn[u] = low[u] = ++timestamp;
    in_stk[u] = true;
    stk[stkTop++] = u;
    for(int i = h[u];i!=-1;i = ne[i]){
        int j = e[i];
        if(dfn[j] == 0){
            tarjan(j);
            low[u] = min(low[u],low[j]);
        }else if(in_stk[j])
            low[u] = min(low[u],dfn[j]);
    }
    if(dfn[u] == low[u]){
        int y = -1;
        ++scc_cnt;
        do{
            y = stk[--stkTop];
            id[y] = scc_cnt;
            in_stk[y] = false;
            val[scc_cnt]+=w[y];
        }while(y!=u);
    }
}

int main(){
    scanf("%d%d",&n,&m);
    memset(h,-1,sizeof h);
    memset(hs,-1,sizeof hs);
    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++){
        int a;
        scanf("%d",&a);
        w[i] = a;
    }
    for(int i = 1;i<=n;i++)
        if(dfn[i] == 0)
            tarjan(i);
            
    int s,p;
    scanf("%d%d",&s,&p);
    for(int i = 1;i<=p;i++){
        int a;
        scanf("%d",&a);
        is_bar[id[a]] = true;
    }
    
    unordered_set<LL> set;
    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];
            LL hash = (LL)a*N+b;
            if(set.find(hash)!=set.end()) continue;
            if(a!=b){
                set.insert(hash);
                add(hs,a,b);
            }
        }
    
    int q[N];
    int hh = 0,tt = -1;
    bool visit[N];
    q[++tt] = id[s];
    f[id[s]] = val[id[s]];
    visit[id[s]] = true;
    while(tt>=hh){
        int t = q[hh++];
        for(int i = hs[t];i!=-1;i = ne[i]){
            int j = e[i];
            f[j] = max(f[j],f[t]+val[j]);
            if(!visit[j]) q[++tt] = j;
        }
    }
    int ans = 0;
    for(int i = 1;i<=scc_cnt;i++)
        if(is_bar[i]) ans = max(ans,f[i]);
    printf("%d",ans);
    return 0;
}
2023/7/30 13:34
加载中...