萌新代码正确性求助
查看原帖
萌新代码正确性求助
225883
MiRaciss楼主2023/2/16 16:59

rt,代码没过样例,现在不知道自己打的是不是对的,用的是这篇题解的思路,希望有大佬帮忙看一下正确性。

#include<bits/stdc++.h>
using namespace std;
#define ll long long

struct Matrix{
	ll a[3][3];
	void clear(){
		for(int i=0;i<3;i++) for(int j=0;j<3;j++) a[i][j]=0x7fffff;
	}
	void Init(){
		clear();
		for(int i=0;i<3;i++) a[i][i]=0;	
	}
}A[200005],Ma[200005][22],Re[200005][22];
Matrix operator *(Matrix a,Matrix b){
	Matrix ans;ans.clear();
	for(int i=0;i<3;i++) for(int j=0;j<3;j++) for(int k=0;k<3;k++) ans.a[i][j]=min(ans.a[i][j],a.a[i][k]+b.a[k][j]);
	return ans;
}

vector<int> v[200005];
int a[200005];
int num[200005];
int n,m,k;

bool cmp(int x,int y){
	return a[x]<a[y];
}

void Build(){
	for(int i=1;i<=n;i++) A[i].clear();
	if(k==1){
		for(int i=1;i<=n;i++) A[i].a[0][0]=a[i],A[i].a[1][1]=A[i].a[2][2]=0;
	}
	if(k==2){
	/*	f[i][0]=min(f[i-1][0],f[i-1][1])+val;
		f[i][1]=f[i-1][0];
		f i-1,f i-2-> f i,f i-1  */
		for(int i=1;i<=n;i++) A[i].a[0][0]=A[i].a[0][1]=a[i],A[i].a[1][0]=A[i].a[2][2]=0;		
	}
	if(k==3){
	/* 	f[i][0]=min(f[i-1][0],f[i-1][1],f[i-1][2])+val;
		f[i][1]=min(f[i-1][0],f[i-1][1]+num,f[i-1][2]+val)
		f[i][2]=f[i-1][1]		*/
		for(int i=1;i<=n;i++) 
			A[i].a[0][0]=A[i].a[0][1]=A[i].a[0][2]=a[i],
			A[i].a[1][0]=0,A[i].a[1][1]=a[v[i][0]],A[i].a[1][2]=a[i],
			A[i].a[2][1]=0; 
	}
}

int fa[200005][22];
int dep[200005];

void DFS(int x,int f){
	fa[x][0]=f;Ma[x][0]=Re[x][0]=A[x];dep[x]=dep[f]+1;
	for(int i=1;i<=20;i++) 
		fa[x][i]=fa[fa[x][i-1]][i-1],
		Ma[x][i]=Ma[x][i-1]*Ma[fa[x][i-1]][i-1],
		Re[x][i]=Re[fa[x][i-1]][i-1]*Re[x][i-1];
	for(auto y:v[x]) if(y^f)
		DFS(y,x);
}

int Lca(int x,int y){
	if(dep[x]<dep[y]) swap(x,y);
	for(int i=20;i>=0;i--) if(dep[fa[x][i]]>dep[y]) x=fa[x][i];
	if(x==y) return x;
	for(int i=20;i>=0;i--) if(fa[x][i]!=fa[y][i]) x=fa[x][i],y=fa[y][i];
	return fa[x][0];
}

int Find(int x,int y){
	int lca=Lca(x,y);
	Matrix now1,now2;now1.clear(),now2.clear();
	now1.a[0][0]=a[x],now2.a[0][0]=a[y];
	for(int i=20;i>=0;i--) if(dep[fa[x][i]]<dep[lca]) now1=now1*Ma[x][i],x=fa[x][i];
	now1=now1*Ma[x][0];
	for(int i=20;i>=0;i--) if(dep[fa[y][i]]<dep[lca]) now2=now2*Re[y][i],y=fa[y][i];
	now2=now2*Re[y][0];
	now1=now1*now2;
	return now1.a[0][0];
}

int main(){
	cin>>n>>m>>k;
	for(int i=1;i<=n;i++) scanf("%d",&a[i]);
	for(int i=1,x,y;i<n;i++) scanf("%d%d",&x,&y),v[x].push_back(y),v[y].push_back(x); 
	for(int i=1;i<=n;i++) sort(v[i].begin(),v[i].end(),cmp);
	Build();DFS(1,0);
	for(int i=1,x,y;i<=m;i++){
		scanf("%d%d",&x,&y);
		printf("%lld\n",Find(x,y));
	}
	return 0; 
}
2023/2/16 16:59
加载中...