80pts求调
查看原帖
80pts求调
254491
橙橙like海绵楼主2023/2/3 16:07
#include<bits/stdc++.h>
#define ll long long 
using namespace std;
const int N=5e5+10;
const int M=100;
const int inf=0x3f3f3f3f;
int qfs=0;
int n,ans=inf,rt,k;
int maxx[N],sz[N];
int d[N],a[N],b[N],c[N],tot;
int h[N],cnt;
bool vis[N];
struct node{
    int nxt,to,val;
}e[N<<2];
inline ll read(){
    ll s=0,w=1;
    char ch=getchar();
    while(ch<'0'||ch>'9'){if(ch=='-') w=-1;ch=getchar();}
    while(ch>='0'&&ch<='9'){s=s*10+ch-'0';ch=getchar();}
    return s*w;
}
void add(int u,int v,int w){
    e[++cnt].nxt=h[u];
    e[cnt].to=v;
    e[cnt].val=w;
    h[u]=cnt;
}
void get_rt(int u,int fa,int t){
    sz[u]=1,maxx[u]=0;
    for(int i=h[u];i;i=e[i].nxt){
        int v=e[i].to;
        if(v==fa||vis[v]) continue;
        get_rt(v,u,t);
        sz[u]+=sz[v];
        maxx[u]=max(maxx[u],sz[v]);
    }
    maxx[u]=max(maxx[u],t-sz[u]);
    if(!rt||maxx[u]<maxx[rt]) rt=u;
}
void get_dis(int u,int fa,int dis,int num,int from){
    //puts("in get_dis");
    a[++tot]=u;
    b[u]=from;
    d[u]=dis;
    c[u]=num;
    for(int i=h[u];i;i=e[i].nxt){
        int v=e[i].to;
        if(v==fa||vis[v]) continue;
        get_dis(v,u,dis+e[i].val,num+1,from);
    }
}
bool cmp(int x,int y){
    if(d[x]==d[y]) return c[x]<c[y];
    else return d[x]<d[y];
} 
void get_sum(int u){
    //puts("in get_sum");
    tot=0;
    a[++tot]=u;
    b[u]=u;
    d[u]=0;
    c[u]=0;
    for(int i=h[u];i;i=e[i].nxt){
        int v=e[i].to;
        if(vis[v]) continue;
        get_dis(v,u,e[i].val,1,v);
    }
    //puts("get_dis ok");
    sort(a+1,a+tot+1,cmp);
    int l=1,r=tot;
    while(l<r){
        //puts("in while");
        if(d[a[l]]+d[a[r]]>k) r--;
        else if(d[a[l]]+d[a[r]]<k) l++;
        else if(b[a[l]]==b[a[r]]){
            if(d[a[r]]==d[a[r-1]]) r--;
            else l++;
        }
        else{
            /*qfs++;
            if(qfs>=M) break;*/
            ans=min(ans,c[a[l]]+c[a[r]]);
            break;
        } 
    }
}
void solve(int u){
    vis[u]=1;
    get_sum(u);
    //puts("get_sum ok");
    for(int i=h[u];i;i=e[i].nxt){
        int v=e[i].to;
        if(vis[v]) continue;
        rt=0;   
        get_rt(v,0,sz[v]);                
        solve(rt);
    }
}
int main(){
    n=read();k=read();
    int u,v,w;
    for(int i=1;i<n;i++){
        u=read()+1;v=read()+1;w=read();
        add(u,v,w);add(v,u,w);
    }
    maxx[0]=n;
    get_rt(1,0,n);
    //puts("get_rt ok");
    solve(rt);
    if(ans==inf) puts("-1");
    else printf("%d",ans);
    return 0;
}
2023/2/3 16:07
加载中...