WA#6,和答案差1求看!
查看原帖
WA#6,和答案差1求看!
813227
Hunter19019楼主2023/3/21 23:22
#include<iostream>
#include<algorithm>
#include<cmath>
#include<vector>
#include<cstring>
#include<queue>
#define x first
#define y second
using namespace std;
typedef long long ll;
const int N = 2e+5+20;
typedef pair<int,int> PII;
vector<PII>g[N];
vector<int> path;
bool vis[N];
int dist[N],maxd;
int n,k,ans;
void dp(int m)
{
    vis[m] = true;
    for(auto u:g[m])
    {
        if(vis[u.x]) continue;
        int t = u.x;
        dp(t);
        ans = max(ans,dist[m]+dist[t]+u.y);
        dist[m] = max(dist[m],dist[t]+u.y);
    }
}
int bfs(int s)
{
    queue<int> q;
    q.push(s);
    int v = 0;
    while(!q.empty())
    {
        int t = q.front();
        q.pop();
        vis[t] = true;
        for(auto m:g[t])
        {
            if(vis[m.x])continue;
            vis[m.x] = true;
            int u = m.x;
            //vis[u] = true;
            dist[u] = dist[t]+m.y;
            if(dist[u] > maxd)
            {
                maxd = dist[u];
                v = u;
            }
            q.push(u);

        }
    }
    return v;
}
bool dfs(int a,int b)
{
    vis[a] = true;
    for(auto t : g[a])
    {
        int u = t.x;
        if(vis[u]) continue;
        if(u == b || dfs(u,b))
        {
            path.push_back(a);
            return true;
        }
    }
    return false;
}
int main()
{
    int a,b;
    scanf("%d%d",&n,&k);
    for (int i = 0; i < n - 1; ++i) {
        scanf("%d%d",&a,&b);
        g[a].push_back({b,1});
        g[b].push_back({a,1});
    }
    int u = bfs(1),s;

    memset(dist,0,sizeof dist);
    memset(vis, false,sizeof vis);

    //fill(dist,dist+n+1,0);
    //fill(vis,vis+n+1,false);
    maxd = 0;
    s = bfs(u);
    memset(dist,0,sizeof dist);
    memset(vis, false,sizeof vis);
    path.push_back(s);
    dfs(u,s);

    for(int i = 0; i < path.size()-1; i++)
    {
        int a = path[i],b=path[i+1];
        for(auto it = g[a].begin(); it != g[a].end(); ++it)
        {
            if((*it).x==b)
            {
                (*it).y=-1;
                break;
            }
        }
        for(auto it = g[b].begin(); it != g[b].end(); ++it)
        {
            if((*it).x==a)
            {
                (*it).y=-1;
                break;
            }
        }
    }
    g[s].push_back({u,-1});
    g[u].push_back({s,-1});
    fill(dist,dist+n+1,0);
    fill(vis,vis+n+1,false);
    //memset(dist,0,sizeof dist);
    //memset(vis, false,sizeof vis);
    if(k == 2)
    {
        dp(1);
        printf("%d",n*2-ans-maxd);
    }
    else printf("%d",n*2-1-maxd);
   // printf("\n%d %d",ans,maxd);
    return 0;
}

和答案差了1,看了下是dp出了问题,不知道为何

2023/3/21 23:22
加载中...