萌新 TLE 求助
查看原帖
萌新 TLE 求助
533854
CodingShark楼主2023/3/21 17:11

rt,TLE 72pts,代码是矩阵乘法:

#include <bits/stdc++.h>
using namespace std;
#define inf 0x3f3f3f3f3f3f3f3f
typedef long long ll;
typedef vector<vector<ll>> matrix;
const int N = 2e5 + 5;
struct Edge {
	int u, v, next;
} e[N << 1];
int n, q, k, pos, head[N], val[N], dep[N], st[N][21], b[N];
matrix up[N][21], dw[N][21];
void addEdge(int u, int v) {
	e[++pos] = {u, v, head[u]};
	head[u] = pos;
}

matrix mul(matrix a, matrix b) {
	if (!a.size() || !b.size()) return matrix{};
	int m = a.size(), n = b.size(), p = b[0].size();
	matrix res(m, vector<ll>(p, inf));
	for (int i = 0; i < m; i++)
		for (int j = 0; j < p; j++)
			for (int k = 0; k < n; k++)
				res[i][j] = min(res[i][j], a[i][k] + b[k][j]);
	return res;
}

matrix qpow(matrix a, ll y) {
	int n = a.size();
	matrix base = a, res(n, vector<ll>(n, inf));
	for (int i = 0; i < n; i++) res[i][i] = 0;
	while (y) {
		if (y & 1) res = mul(res, base);
		base = mul(base, base), y >>= 1;
	}
	return res;
}

matrix init(int u) {
	if (k == 1) {
		return matrix{{val[u]}};
	} else if (k == 2) {
		return matrix{{val[u], 0}, {val[u], inf}};
	} else {
		return matrix{{val[u], 0, inf}, {val[u], b[u], 0}, {val[u], inf, inf}};	
	}
	assert(false);
}

void dfs(int u, int fa = 0) {
	dep[u] = dep[fa] + 1; st[u][0] = fa;
	for (int i = head[u]; i; i = e[i].next) {
		int v = e[i].v;
		if (v == fa) continue;
		b[v] = val[u];
		dfs(v, u);
		b[u] = min(b[u], val[v]);
	}
	up[u][0] = dw[u][0] = init(u);
}

int lca(int u, int v) {
	if (dep[u] < dep[v]) swap(u, v);
	while (dep[u] != dep[v])
		u = st[u][__lg(dep[u] - dep[v])];
	if (u == v) return u;
	for (int i = __lg(dep[u]); i >= 0; i--)
		if (st[u][i] != st[v][i])
			u = st[u][i], v = st[v][i];
	return st[u][0];
}

matrix query_up(int u, int v) {
	matrix res(k, vector<ll>(k, inf));
	for (int i = 0; i < k; i++) res[i][i] = 0;
	while (dep[u] != dep[v]) {
		int g = __lg(dep[u] - dep[v]);
		res = mul(res, up[u][g]);
		u = st[u][g];
	}
	return res;
}

 matrix query_dw(int u, int v) {
	matrix res(k, vector<ll>(k, inf));
	for (int i = 0; i < k; i++) res[i][i] = 0;
	while (dep[u] != dep[v]) {
		int g = __lg(dep[u] - dep[v]);
		res = mul(dw[u][g], res);
		u = st[u][g];
	}
	return res;
 }

matrix get() {
	if (k == 1) {
		return matrix{{0}};
	} else if (k == 2) {
		return matrix{{inf, 0}};
	} else {
		return matrix{{inf, inf, 0}};
	}
	assert(false);
}

ll query(int u, int v) {
	matrix mat = get();
	int f = lca(u, v);
	return mul(mul(mul(mat, query_up(u, f)), init(f)), query_dw(v, f))[0][0];
}

int main() {
	scanf("%d%d%d", &n, &q, &k);
	for (int i = 1; i <= n; i++) scanf("%d", val + i);
	for (int i = 1; i < n; i++) {
		int u, v;
		scanf("%d%d", &u, &v);
		addEdge(u, v), addEdge(v, u);
	}
	b[1] = inf, dfs(1);
	for (int i = 1; i <= __lg(n); i++) {
		for (int u = 1; u <= n; u++) {
			st[u][i] = st[st[u][i - 1]][i - 1];
			up[u][i] = mul(up[u][i - 1], up[st[u][i - 1]][i - 1]);
			dw[u][i] = mul(dw[st[u][i - 1]][i - 1], dw[u][i - 1]);
		}
	}
	while (q--) {
		int u, v;
		scanf("%d%d", &u, &v);
		printf("%lld\n", query(u, v));
	}
    return 0;
}
2023/3/21 17:11
加载中...