普及月赛T4样例没过,蒟蒻求助
  • 板块学术版
  • 楼主Expert_Dream
  • 当前回复0
  • 已保存回复0
  • 发布时间2023/8/26 22:06
  • 上次更新2023/11/3 01:00:17
查看原帖
普及月赛T4样例没过,蒟蒻求助
768530
Expert_Dream楼主2023/8/26 22:06

认为时间复杂度应该只能水过50分,但请dalao看看到底哪错了,(思路有点乱

#include <bits/stdc++.h>
#define ll long long
#define ull unsigned long long
#define endl '\n'
using namespace std;
namespace Syxqwq {
	inline int read() {
		int x = 0, s = 1;
		char c = getchar();
		while (c > '9' || c < '0') {
			if (c == '-') s = -1;
			c = getchar();
		}
		while (c >= '0' && c <= '9') {
			x = (x << 1) + (x << 3) + (c - '0');
			c = getchar();
		}
		return x * s;
	}
	void Write(int x) {
		if (x < 0) {
			putchar('-');
			x = -x;
		}
		if (x > 9) Write(x / 10);
		putchar(x % 10 + '0');
	}
	inline void write(int x, char c) {
		Write(x), putchar(c);
	}
}
using namespace Syxqwq; 
const int mod = 998244353;
const int N = 2e5+5;
const int inf = 0x3f3f3f3f;
int n,q;
struct node{
	int to,v;
};
struct data{
	int u,v,w;
}a[N];
vector<node> mp[N];
int dep[N];
int cost[N];
void dfs(int x,int fa){
	dep[x] = 1;
	for(auto it:mp[x]){
		if(it.to==fa)continue;
		cost[it.to] = (cost[x] + it.v) % mod;
		dfs(it.to,x);
		dep[x] += dep[it.to];
	}
}
ll cnt,sum,sum2;
void dfs2(int x,int kk){
	for(auto it:mp[x]){
		if(cost[it.to] <= cost[x]){
			sum2 = (sum2-cost[it.to]+cost[kk]-cost[it.to]) % mod;
			sum2 = (sum2 - (cost[it.to] * (dep[it.to]-dep[x])) + ((cost[kk]-cost[it.to]) * (dep[it.to]-dep[x]))) %mod;
			dfs2(it.to,kk);
			break;
		}
		
	}
}
int main(){
	n=read();
	q=read();
	for(int i = 1;i <= n-1;i++){
		int u,v,w;
		u=read();
		v=read();
		w=read();
		mp[u].push_back({v,w});
		mp[v].push_back({u,w});
		a[i].u=u;a[i].v = v;a[i].w = w;
	}
	dfs(1,0);
	/*
	预处理出:深度(cost)即每一个结点到根节点的距离和子树数量(dep)即以i为根的子树大小
	*/
	
	for(int i = 1;i <= n-1;i++){	
		if(cost[a[i].v] < cost[a[i].u])	swap(a[i].v,a[i].u);
		cnt = (cnt + ((dep[a[i].v] * (n-dep[a[i].v]))%mod) * a[i].w)%mod;//计算出各个点之间的距离
		//通过 枚举每一条边,知道它两边的数量,乘法原理,得出每一条边的贡献(原本n个点)
	}
	for(int i = 1;i <= n;i++){
		sum = (sum+cost[i]) % mod;//计算出所有边距离根节点的长度之和
	}
	for(int i = 1;i <= q;i++){
		int k,w;
		k=read();
		w=read();
		sum2 = sum;//重新初始化,记录的是w到每一个点的距离之和
		dfs2(k,k);
		sum2 = (sum2 + n*w) % mod;
		sum2 = (sum2+cnt)%mod;
		write(sum2,endl);
	}
	

	return 0;
}

2023/8/26 22:06
加载中...