题解:P15582 [KTSC 2026] 平衡序列 / Balanced Sequence

我咋写的那么蠢。

显然平衡序列的长度只能是 \(2^k-1\),因此平衡区间的总个数是 \(\mathcal{O}(n\log{n})\) 的。进一步观察性质,发现一个区间不会包含另一个等长度的区间的中点。考察一个位置 \(x\),对于一种长度 \(len\),包含 \(x\) 的长度为 \(len\) 的平衡区间中,中点 \(\leq x\)\(>x\) 的区间都至多只有 \(1\) 个。

因此我们可以对于每种长度分别维护所有平衡区间,单点修改时,从小到大枚举长度 \(2^{k+1}-1\),分类讨论包含 \(x\) 的平衡区间的中点 \(mid\)\(x\) 的大小关系:

  • \(mid<x\):此时 \(mid\) 一定是 \(a_{x-2^k+1\sim x-1}\) 的最大值。
  • \(mid=x\)
  • \(mid>x\):此时 \(mid\) 一定是 \(a_{x+1\sim x+2^k-1}\) 的最大值。

用线段树找出最大值位置,确定 \(mid\) 后根据定义直接判定是否合法即可。

使用哈希表维护每一层的中点位置,时间复杂度为 \(\mathcal{O}((n+q)\log^2{n})\)

代码
#include <bits/stdc++.h>
#include <ext/pb_ds/assoc_container.hpp>
#include <ext/pb_ds/tree_policy.hpp>

using namespace std;
using namespace __gnu_pbds;

using ll = long long;
using i128 = __int128;
using ui = unsigned int;
using ull = unsigned long long;
using u128 = unsigned __int128;
using ld = long double;
using pii = pair<int, int>;
const int N = 1e5 + 5;

template<typename T> inline T lowbit(T x) { return x & -x; }
template<typename T> inline void chk_min(T &x, T y) { x = y < x ? y : x; }
template<typename T> inline void chk_max(T &x, T y) { x = x < y ? y : x; }

int n, ans, mxd, a[N];
gp_hash_table<int, bool> s[17];

struct SegTree {
#define ls(p) (p << 1)
#define rs(p) (p << 1 | 1)
	struct Node {
		int mx, mxp;
		friend Node operator+(const Node &lhs, const Node &rhs) {
			return lhs.mx > rhs.mx ? lhs : rhs;
		}
	} nd[N << 1];
	void push_up(int p) { nd[p] = nd[ls(p)] + nd[rs(p)]; }
	void build() {
		for (int i = 1; i <= n; ++i) nd[i + n - 1] = {a[i], i};
		for (int i = n - 1; i; --i) push_up(i);
	}
	void upd(int x, int v) {
		nd[x + n - 1] = {v, x};
		x += n - 1;
		while (x > 1) push_up(x >>= 1);
	}
	Node query(int l, int r) {
		Node res = {-1, 0};
		for (l += n - 1, r += n - 1; l <= r; l >>= 1, r >>= 1) {
			if (l & 1) res = res + nd[l++];
			if (~r & 1) res = nd[r--] + res;
		}
		return res;
	}
#undef ls
#undef rs
} sgt;

bool check(int x, int dep) {
	int mid = 1 << dep;
	if (x - mid + 1 <= 0 || x + mid - 1 > n) return 0;
	int mxl = sgt.query(x - mid + 1, x - 1).mx;
	int mxr = sgt.query(x + 1, x + mid - 1).mx;
	int midl = x - (mid >> 1);
	int midr = x + (mid >> 1);
	return s[dep - 1].find(midl) != s[dep - 1].end() && s[dep - 1].find(midr) != s[dep - 1].end() && a[x] > max(mxl, mxr);
}

void add(int x, int dep) {
	int mid = 1 << dep;
	ans -= s[dep].size();
	if (x > mid) {
		int p = sgt.query(x - mid + 1, x - 1).mxp;
		if (check(p, dep)) s[dep][p] = 1;
		else if (s[dep].find(p) != s[dep].end()) s[dep].erase(p);
	}
	if (check(x, dep)) s[dep][x] = 1;
	else if (s[dep].find(x) != s[dep].end()) s[dep].erase(x);
	if (x + mid <= n) {
		int p = sgt.query(x + 1, x + mid - 1).mxp;
		if (check(p, dep)) s[dep][p] = 1;
		else if (s[dep].find(p) != s[dep].end()) s[dep].erase(p);
	}
	ans += s[dep].size();
}

ll initialize(int N, vector<int> A) {
	n = N, mxd = __lg(n + 1) - 1;
	copy(A.begin(), A.end(), a + 1);
	sgt.build();
	ans += n;
	for (int i = 1; i <= n; ++i) s[0][i] = 1;
	for (int d = 1; d <= mxd; ++d)
		for (int i = 1; i <= n; ++i)
			add(i, d);
	return ans;
}

ll update_sequence(int p, int v) {
	++p;
	if (a[p] == v) return ans;
	a[p] = v, sgt.upd(p, v);
	for (int d = 1; d <= mxd; ++d) add(p, d);
	return ans;
}
posted @ 2026-03-09 11:46  P2441M  阅读(16)  评论(0)    收藏  举报