[2449일] F. Momoyo and the Network (18차 도전, 실패) python3

SparklingJustForYou·2026년 9월 30일

코드포스 공부기록

목록 보기
16/23

> F. Momoyo and the Network
time limit per test3 seconds
memory limit per test256 megabytes
Where Is That Bustling Marketplace Now— Unconnected Marketeers
The Underground Great Line Network is a grand transit system connecting all corners of Gensokyo. Momoyo noticed that the network's layout resembled a tree∗ structure. She couldn't help but imagine the most effective way to dismantle that tree.
Given a tree with n nodes where node i has weight ai, select a simple path of exactly k edges and remove all edges on it. This splits the tree into k+1 connected components, each with weight equal to the sum of its nodes' weights. You need to maximize the minimum component weight, or output −1 if no simple path of exactly k edges exists.
∗A tree is a connected graph without cycles.
Input
Each test contains multiple test cases. The first line contains the number of test cases t (1≤t≤104). The description of the test cases follows.
The first line of each test case contains two integers n and k (1≤k≤n−1, 2≤n≤2⋅105).
The second line contains n integers, where the i-th integer represents ai (1≤ai≤109).
The next n−1 lines each contain two integers u and v, representing an edge of the tree.
It is guaranteed that the sum of n over all test cases does not exceed 2⋅105.
Output
> For each test case, output the maximum possible minimum component weight, or −1 if no such path exists.

c#

using System;
using System.Collections.Generic;
using System.IO;
using System.Text;
using System.Threading;
class Program
{
    static BufferedStream sr = new BufferedStream(Console.OpenStandardInput());
    static StreamWriter sw = new StreamWriter(new BufferedStream(Console.OpenStandardOutput()), Encoding.Default);
    static int[] arr=new int [200001];
    static List<int>[] edge = new List<int>[200001];
    static long[] edge_v = new long[200001], edge_min = new long[200001], edge_len = new long[200001];
    static long total = 0,k=0;
    struct Pair:IComparable<Pair>
    {
        public long Key, Value;
        public Pair(long Key,long Value)
        {
            this.Key = Key;
            this.Value = Value;
        }
        public int CompareTo(Pair other)
        {
            return other.Key.CompareTo(this.Key);
        }
    }
    static Pair[][] temps = new Pair[200001][];
    static int[][] temp = new int[200001][];
    static int[] temp_cnt = new int[200001];
    static void Set_Edge(int idx,int parent)
    {
        temp_cnt[idx] = 0;
        edge_v[idx]   = arr[idx];
        edge_min[idx] = 0;
        var list = edge[idx];
        int temp_i = 0;
        for(int i=0;i< list.Count;i++)
        {
            int next=list[i];
            if (next==parent) continue;
            Set_Edge(next,idx);
            edge_v[idx]  += edge_v[next];
            temp_cnt[idx]++;
            temp[idx][temp_i++] = next;
        }
        var t_list = temp[idx];
        for (int i=0;i< temp_i;i++)
        {
            int next = t_list[i];
            edge_min[next] = Math.Min(edge_v[next], edge_v[idx] - edge_v[next]);
        }
    }
    static bool DFS(int idx,long mid)
    {
        edge_len[idx] = 0;
        int temps_i = 0;
        var list = temp[idx];
        int size = temp_cnt[idx];
        for (int i = 0; i < size;i++)
        {
            int next=list[i];
            if (DFS(next, mid)) return true;
            if (mid <= edge_v[next])
            {
                temps[idx][temps_i++]=new Pair( edge_v[next], edge_len[next]);
            }
            if (mid<=edge_min[next])
            {
                edge_len[idx] = Math.Max(1 + edge_len[next], edge_len[idx] )  ;
            }
            if(edge_len[next] +1>=k&&total- edge_v[next] >=mid&& edge_v[next] >=mid)
            {
                return true;
            }
        }
        if (temps_i >1) Array.Sort(temps[idx], 0, temps_i);
        int ii = -1;
        long l = -1, r = -1;
        //아 무조건 l,r범위가 커지도록 하려면 작은쪽에서 커지도록 해야함;;
        var t_list = temps[idx];
        for (int j=0;j<temps_i;j++)
        {
            var i = t_list[j];
            while (ii + 1 < temps_i && total - t_list[ii + 1].Key- i.Key >= mid)
            {
                ii++;
                if (l <= t_list[ii].Value)
                {
                    r = l;
                    l = t_list[ii].Value;
                }
                else if( r<= t_list[ii].Value) r = t_list[ii].Value;
            }
            if (ii ==-1) continue;
            if (total - 2 * i.Key >= mid && l == i.Value)
            {
                if (i.Value + r + 2 >= k) return true;
            }
            else
                if (l + i.Value+ 2 >= k) return true;

        }
        return false;
    }
    static int ReadInt()
    {
        int c = sr.ReadByte();
        while (c <= 32) { if (c == -1) return -1; c = sr.ReadByte(); }
        bool neg = false;
        if (c == '-') { neg = true; c = sr.ReadByte(); }
        int val = 0;
        while (c > 32)
        {
            val = val * 10 + (c - '0');
            c = sr.ReadByte();
        }
        return neg ? -val : val;
    }
    static void Solve()
    {
        int t = ReadInt();
        for (int i = 0; i < 200001; i++)
        {
            edge[i] = new List<int>();
        }
        for (int i = 0; i < t; i++)
        {
            int n = ReadInt();
            k = ReadInt();
            total = 0;
            for (int j = 0; j < n; j++)
            {
                arr[j] = ReadInt();
                total += arr[j];
                edge[j].Clear();
            }
            for (int j = 0; j < n - 1; j++)
            {
                int u = ReadInt() - 1;
                int v = ReadInt() - 1;
                edge[u].Add(v);
                edge[v].Add(u);
            }
            for (int j = 0; j < n; j++)
            {
                temp[j] = new int[edge[j].Count];
                temps[j] = new Pair[edge[j].Count];
            }
            Set_Edge(0, -1);
            long l = 0, r = total, answer = -1;
            while (l <= r)
            {
                long mid = (l + r) >> 1;
                if (DFS(0, mid))
                {
                    l = mid + 1;
                    answer = mid;
                }
                else
                {
                    r = mid - 1;
                }
            }
            sw.WriteLine(answer);
        }
        sw.Flush();
    }
    static void Main(string[] args)
    {
        Thread thread = new Thread(Solve, 1024 * 1024 * 64);
        thread.Start();
        thread.Join();
    }
}

cpp코드

#include<iostream>
#include<vector>
#include<algorithm>
using namespace std;
long long arr[200001];
int new_edge_cnt[200001];
vector<int> edge[200001];
long long edge_v[200001];
long long edge_min[200001];
long long edge_len[200001];
vector<int> new_edge[200001];
vector<pair<long long, long long>> edges[200001];
long long n, k, t,edge_total=0;
bool dfs_find(int idx,long long mid)
{
    edge_len[idx] = 0;
    edges[idx].clear();
    for (int i : new_edge[idx]) 
    {
        if (dfs_find(i, mid)) return true;
        if (mid <= edge_v[i])
        {
            edges[idx].push_back({ edge_v[i], edge_len[i] });
        }
        if (edge_min[i] >= mid) {
            edge_len[idx] = max(edge_len[i] + 1, edge_len[idx]);
        }
        if (edge_len[i] + 1 >= k && edge_total - edge_v[i] >= mid && edge_v[i] >= mid) {
            //자를 수 있음.
            return true;
        }
    }
    sort(edges[idx].begin(), edges[idx].end());
    int ii= -1;
    int l =-1, r = -1;
    int cnt = edges[idx].size();
    for (int j = cnt- 1; j >= 0; j--) {
        while (ii + 1 <cnt && edge_total- edges[idx][j].first - edges[idx][ii + 1].first >= mid) {
            ii++;
            if (edges[idx][ii].second > l) {//l값이 크고 r값이 작음
                r =l;
                l = edges[idx][ii].second;
            }
            else if (edges[idx][ii].second > r) {
                r = edges[idx][ii].second;
            }
            //이 과정이 잘 이해가 안갔는데 결국 k 값 이상되는 길이를 찾는 과정임.
        }
        if (ii == -1) {
            continue;
        }
        if (edge_total - 2 * edges[idx][j].first >= mid &&l== edges[idx][j].second) {
            //어 지금 제일 큰 길이와 동일하면 r을 써야함
            if (r + edges[idx][j].second + 2 >= k) {
                return true;
            }
        }
        else {
            if (l + edges[idx][j].second + 2 >= k) {
                return true;
            }
        }
    }
    return false;
}
void dfs(int idx,int parent)
{
    edge_v[idx] = arr[idx];
    edge_min[idx] = 0;
    new_edge[idx].clear();
    for (int i : edge[idx])
    {
        if (i==parent)continue;
        dfs(i,idx);
        edge_v[idx] += edge_v[i];
        new_edge[idx].push_back(i);
    }
    for (int i: new_edge[idx])
    {
        edge_min[i] = min(edge_v[idx] - edge_v[i], edge_v[i]);
    }
}
int main()
{
	ios_base::sync_with_stdio(false);
	cin.tie(NULL);
	cout.tie(NULL);
	int u,v; 
	cin >> t;
	for (int a = 0; a < t; a++)
	{
		cin >> n >> k;
        edge_total = 0;
		for (int i = 0; i < n; i++)
		{
			cin >> arr[i];
            edge_total += arr[i];
            edge[i].clear();
		}
		for (int i = 0; i < n-1; i++)
		{
			cin >> u >> v;
            u--;
            v--;
            edge[u].push_back(v);
            edge[v].push_back(u);
		}
        dfs(0,-1);
        long long l = 0, r = edge_total;
        long long result = -1;
        while(l<=r)
        {
            long long mid = (l + r) >> 1;
          if(  dfs_find(0, mid))
          {
              l=mid+1;
              result = mid;
          }
          else
          {
              r=mid-1;
          }
        }
        cout << result << "\n";
	}
	return 0;
}

python3

import sys
input = sys.stdin.readline
data = sys.stdin.buffer.read()
ndata = len(data)
ii = 0
def read_int():
    global ii, ndata, data
    while ii < ndata and data[ii] <= 32:
        ii += 1
    if ii >= ndata:
        return 0
    sign = 1
    if data[ii] == 45:  # ASCII 45 == '-'
        sign = -1
        ii += 1
    num = 0
    while ii < ndata and data[ii] > 32:
        num = num * 10 + (data[ii] - 48)
        ii += 1
    return num * sign
t=read_int()
for i in range(t):
    n=read_int()
    edge= [[] for _ in range(n) ] 
    weight=[0]*n
    k=read_int()
    for j in range(n):
        weight[j]=  read_int()
    for j in range(n-1):
        u=read_int()-1
        v=read_int()-1
        edge [u] .append(v);
        edge [v] .append(u);
    st= [0] 
    bottom_up_edge= [] 
    parent= [-1] *n
    while st:
        current=st.pop()
        bottom_up_edge.append(current)
        for x in edge[current] :
            if x== parent [current] :
                continue
            parent [x] =current
            st.append(x)
    bottom_up_edge.reverse()

#보니까 일단 bfs형식으로 트리를 미리 만들어주고
#만들 때부터 정렬까지 진행해버리기
# 부모 노드에 자식 노드 가중치를 더해서 값을 세팅해줌.
    tree = weight[:]   
    for node in bottom_up_edge:
        if node != 0:
            tree[parent[node]] += tree[node]

오늘은 조금만..

0개의 댓글