整体思想是先缩点,建一张新图,新图记录每个缩点的最大最小值
然后由1开始跑一遍bfs,记录那些点能从1走到(用了并查集)
然后再从n反着跑一边bfs,算出答案,到点1停止
#include<bits/stdc++.h>
#include<vector>
#include<queue>
using namespace std;
int n,m,tot,cnt,num,top,ans;
struct hi{
int x,y;
}a[1000005];
struct newprice{
int s,l;
}w[100005];
int p[100005];
int dfn[100005],low[100005],s[100005],f[100005];
int be[100005],id[100005],ed[100005];
int to[2000005],la[100005],ne[2000005];
int minn[100005],vis[100005];
vector<int> v1[100005];
vector<int> v2[100005];
queue<int> q;
int fa[100005];
void add(int x,int y){
to[++tot]=y;
ne[tot]=la[x];
la[x]=tot;
}
void work(int x,int num){
f[x]=0;
be[x]=num;//属于哪个新点集
w[num].s=min(w[num].s,p[x]);
w[num].l=max(w[num].l,p[x]);
}
void tarjan(int x){
dfn[x]=low[x]=++cnt;
s[++top]=x;
f[x]=1;
for(int i=la[x];i;i=ne[i]){
int y=to[i];
if(dfn[y]==0){
tarjan(y);
low[x]=min(low[x],low[y]);
}
else if(f[y]==1) low[x]=min(low[x],low[y]);
}
if(dfn[x]==low[x]){
work(x,++num);
while(s[top]!=x){
work(s[top],num);
top--;
}
top--;
//cout<<num<<" "<<w[num].l<<" "<<w[num].s<<"\n";
}
}
int find(int x){//并查集
if(fa[x]==x) return fa[x];
return fa[x]=find(fa[x]);
}
void bfs(){
q.push(be[1]);
vis[be[1]]=1;
for(int i=1;i<=num;i++) minn[i]=w[i].s;
while(!q.empty()){
int x=q.front();
q.pop();
for(int i=0;i<v1[x].size();i++){
int y=v1[x][i];
vis[y]=1;
fa[y]=find(fa[x]);
minn[y]=min(minn[y],minn[x]);
id[y]--;
if(id[y]==0){
q.push(y);
if(be[n]==y) break;
}
}
}
q.push(be[n]);
while(!q.empty()){
int x=q.front();
q.pop();
for(int i=0;i<v2[x].size();i++){
int y=v2[x][i];
if(fa[y]!=be[1]||vis[y]==0) continue;
ans=max(ans,w[x].l-minn[y]);
ans=max(ans,w[x].l-w[x].s);
//cout<<x<<" "<<w[x].l<<" "<<y<<" "<<minn[y]<<" "<<ans<<"\n";
ed[y]--;
if(ed[y]==0) q.push(y);
}
if(x==be[1]) return ;
}
}
int main(){
scanf("%d %d",&n,&m);
for(int i=1;i<=n;i++) scanf("%d",&p[i]);
for(int i=1;i<=n;i++){
w[i].s=200;
w[i].l=0;
}
for(int i=1;i<=m;i++){
int z;
scanf("%d %d %d",&a[i].x,&a[i].y,&z);
add(a[i].x,a[i].y);
if(z==2) add(a[i].y,a[i].x);
}
for(int i=1;i<=n;i++){
if(dfn[i]==0) tarjan(i);
}
for(int i=1;i<=m;i++){//建新图
int x=a[i].x,y=a[i].y;
if(be[x]!=be[y]){
v1[be[x]].push_back(be[y]);
v2[be[y]].push_back(be[x]);
ed[be[x]]++; id[be[y]]++;
//cout<<be[x]<<" "<<be[y]<<"\n";
}
}
for(int i=1;i<=num;i++) fa[i]=i;
bfs();
printf("%d",ans);
return 0;
}
谢谢大家!!!