线段树进阶

更新中

Problem 1

题目

给定两个长度为 \(n\) 序列 \(a,b\),要求支持如下四个操作:

  1. 给定区间 \([l,r]\) 和一个整数 \(k\),对于所有 \(i \in [l,r]\),令 \(a_i \gets a_i + kb_i\)。
  2. 给定区间 \([l,r]\) 和两个整数 \(k, x\),对于所有 \(i \in [l,r]\),令 \(b_i \gets kb_i + x\)。
  3. 给定区间 \([l,r]\),询问 \(\textstyle \sum_{i=l}^{r} a_i \bmod 998244353\)。
  4. 给定区间 \([l,r]\),询问 \(\textstyle \sum_{i=l}^{r} b_i \bmod 998244353\)。

思路

观察两个修改操作,考虑对每个单个位置 \((a,b)\) 叠加修改,我们设计出一个统一的变换操作。于是我们很自然地得到标记 \(\{k_A, k_B, x_B\}\),操作一可看做 \(\{k, 0, 0\}\),操作二可看做 \(\{0, k, x\}\),每次操作将 \((a,b)\) 变为 \((a + k_Ab, k_Bb + x_B)\)。

对于多个修改操作叠加在一起,我们并不知道两种操作的修改顺序,但是最终结果依然等效于做一次操作一再做一次操作二,只不过参数改变了。

但是随后我们又发现,如果真实顺序是先一步操作二后一步操作一,\(a\) 的最终结果漏掉了一个常数项。例如,假设对 \((a, b)\) 进行的操作为 2 l r k1 x1 和 1 l r k2,第一步得到 \((a, k_1b+x_1)\),第二步期望得到 \((a + k_2k_1b + k_2x_1,k_1b+x_1)\),但由于我们的标记没有给 \(a\) 留常数项,所以我们会漏掉 \(k_2x_1\) 这一项。

解决办法便是修改我们的标记,变成 \(\{k_A, x_A, k_B, x_B\}\),操作一看做 \(\{k, 0, 0, 0\}\),操作二看做 \(\{0, 0, k, x\}\),每次操作将 \((a,b)\) 变为 \((a + k_Ab + x_A, k_Bb + x_B)\)。

下面我们考虑怎样合并两个标记。设 \(t_1 = \{ k_{A1}, x_{A1}, k_{B1}, x_{B1} \}\),\(t_2 = \{ k_{A2}, x_{A2}, k_{B2}, x_{B2} \}\)。

对 \((a, b)\) 进行操作一,得到 \((a_1, b_1) = (a + k_{A1}b + x_{A1}, k_{B1}b + x_{B1})\)。

对 \((a_1, b_1)\) 进行操作二,得到 \((a_2, b_2) = (a_1 + k_{A2}b_1 + x_{A2}, k_{B2}b_1 + x_{B2})\)。

带入 \((a_1, b_1)\),得

\[ \begin{align*} a_2 &= a_1 + k_{A2}b_1 + x_{A2} \\ &= a + k_{A1}b + x_{A1} + k_{A2}(k_{B1}b + x_{B1}) + x_{A2} \\ &= a + (k_{A1} + k_{A2}k_{B1})b + (x_{A1} + k_{A2}x_{B1} + x_{A2}) \\ b_2 &= k_{B2}b_1 + x_{B2} \\ &= k_{B2}(k_{B1}b + x_{B1}) + x_{B2} \\ &= k_{B2}k_{B1}b + (k_{B2}x_{B1} + x_{B2}) \end{align*} \]

我们希望找到一个新的标记 \(t = \{ k_A, x_A, k_B, x_B \}\) 使得 \((a,b)\) 能一步变为 \((a_2, b_2)\)。由 \(a_2,b_2\) 表达式可知,\(t = \{ k_{A1} + k_{A2}k_{B1}, x_{A1} + k_{A2}x_{B1} + x_{A2}, k_{B2}k_{B1}, k_{B2}x_{B1} + x_{B2} \}\)。

接下来我们考虑实现,线段树维护 \(sum_a\),\(sum_b\) 和 \(siz\),对其打上一个标记 \(\{ k_A, x_A, k_B, x_B\}\),得到 \(sum_a + k_A \times sum_b + x_A \times siz\) 和 \(k_B \times sum_b + x_B \times siz\)。(为了增强可读性,此处添加 \(\times\) 号。)

重载运算符会使代码好写一点。

代码

#include <bits/stdc++.h>
using namespace std;

const int N = 5e5 + 5;
const int mod = 998244353;

#define ll long long
#define root 1, 1, n
#define ls rt << 1
#define rs rt << 1 | 1
#define lson ls, l, mid
#define rson rs, mid + 1, r

int n, q;
ll a[N], b[N];

struct T {
	ll ka, xa, kb, xb;
	T () {
		ka = 0;
		xa = 0;
		kb = 1;
		xb = 0;
	}
	void initt(ll x, ll y, ll z, ll w) {
		ka = x % mod;
		xa = y % mod;
		kb = z % mod;
		xb = w % mod;
	}
};
struct node {
	ll sa, sb;
	int siz;
	node () {
		sa = 0;
		sb = 0;
		siz = 0;
	}
	void init(int i) {
		sa = a[i];
		sb = b[i];
		siz = 1;
	}
};

node w[N << 2];
T tag[N << 2];

node operator + (const node &l, const node &r) {
	node res;
	res.sa = (l.sa + r.sa) % mod;
	res.sb = (l.sb + r.sb) % mod;
	res.siz = l.siz + r.siz;
	return res;
}
T operator * (const T &t1, const T &t2) {
	T res;
	res.ka = (t1.ka + t2.ka * t1.kb) % mod;
	res.xa = (t1.xa + t2.ka * t1.xb + t2.xa) % mod;
	res.kb = t1.kb * t2.kb % mod;
	res.xb = (t2.kb * t1.xb + t2.xb) % mod;
	return res;
}

void color(int rt, int l, int r, const T &t) {
	w[rt].sa = (w[rt].sa + w[rt].sb * t.ka + t.xa * w[rt].siz) % mod;
	w[rt].sb = (t.kb * w[rt].sb + t.xb * w[rt].siz) % mod;
	tag[rt] = tag[rt] * t;
}

void pushdown(int rt, int l, int r) {
	int mid = (l + r) >> 1;
	color(lson, tag[rt]);
	color(rson, tag[rt]);
	tag[rt] = T();
}

void build(int rt, int l, int r) {
	if (l == r) {
		w[rt].init(l);
		return ;
	}
	int mid = (l + r) >> 1;
	build(lson);
	build(rson);
	w[rt] = w[ls] + w[rs];
}

void modify(int rt, int l, int r, int ql, int qr, const T &t) {
	if (ql <= l && r <= qr) {
		color(rt, l, r, t);
		return ;
	}
	pushdown(rt, l, r);
	int mid = (l + r) >> 1;
	if (ql <= mid) modify(lson, ql, qr, t);
	if (qr > mid) modify(rson, ql, qr, t);
	w[rt] = w[ls] + w[rs];
}

node query(int rt, int l, int r, int ql, int qr) {
	if (ql <= l && r <= qr) {
		return w[rt];
	}
	pushdown(rt, l, r);
	int mid = (l + r) >> 1;
	if (ql <= mid) {
		if (qr > mid) return query(lson, ql, qr) + query(rson, ql, qr);
		else return query(lson, ql, qr);
	}
	else return query(rson, ql, qr);
}

int main() {
	cin.tie(0)->sync_with_stdio(0);
	
	cin >> n >> q;
	for (int i = 1; i <= n; i++) cin >> a[i];
	for (int i = 1; i <= n; i++) cin >> b[i];
	build(root);
	while (q--) {
		int opt, l, r;
		cin >> opt >> l >> r;
		if (opt == 1) {
			int k;
			cin >> k;
			T t;
			t.initt(k, 0, 1, 0);
			modify(root, l, r, t);
		}
		else if (opt == 2) {
			int k, x;
			cin >> k >> x;
			T t;
			t.initt(0, 0, k, x);
			modify(root, l, r, t);
		}
		else if (opt == 3) {
			cout << query(root, l, r).sa << '\n';
		}
		else {
			cout << query(root, l, r).sb << '\n';
		}
	}
	
	return 0;
}

Problem 2

题目

CF438D

给定长度为 \(n\) 序列 \(a\),要求支持如下三个操作:

  1. 给定区间 \([l,r]\),询问 \(\textstyle \sum_{i=l}^{r} a_i\)。
  2. 给定区间 \([l,r]\) 和整数 \(x\),对于所有 \(i \in [l,r]\),令 \(a_i \gets a_i \bmod x\)。
  3. 给定两个整数 \(k,x\),令 \(a_k \gets x\)。

思路

取模操作有些难搞。标记和合并都不好做。于是考虑优化暴力。

首先可以想到可行性剪枝,如果区间最大值小于模数,取模没意义,直接返回。

然后注意到,一个数 \(t\) 对 \(x\) 取模,其中 \(x < t\),那么结果 \(t'\) 一定小于 \(\frac{t}{2}\)。

考虑怎样证明。假设 \(\frac{t}{2} < x < t\),则 \(t \bmod x = t - x\)。由于 \(x > \frac{t}{2}\),所以 \(t - x < t - \frac{t}{2} = \frac{t}{2}\)。假设 \(x \le \frac{t}{2}\),根据取模的基本性质,余数一定小于除数,所以 \(t \bmod x < x\),即 \(t \bmod x < \frac{t}{2}\)。

基于这个原理,可以发现每个数至多取模 \(\log t\) 次,直接暴力就行。

代码

#include <bits/stdc++.h>
using namespace std;

#define ll long long

const int N = 1e5 + 5;

#define root 1, 1, n
#define ls rt << 1
#define rs rt << 1 | 1
#define lson ls, l, mid
#define rson rs, mid + 1, r

struct node {
	ll sum;
	int mx;
	node () {
		sum = 0ll;
		mx = 0;
	}
	void init(int v) {
		sum = v * 1ll;
		mx = v;
	}
};

node operator + (const node &x, const node &y) {
	node res;
	res.sum = x.sum + y.sum;
	res.mx = max(x.mx, y.mx);
	return res;
}

node w[N << 2];
int n, m;
int a[N];

void build(int rt, int l, int r) {
	if (l == r) {
		w[rt].init(a[l]);
		return ;
	}
	int mid = (l + r) >> 1;
	build(lson);
	build(rson);
	w[rt] = w[ls] + w[rs];
}

void modify(int rt, int l, int r, int x, int v) {
	if (l == r) {
		w[rt].init(v);
		return ;
	}
	int mid = (l + r) >> 1;
	if (x <= mid) modify(lson, x, v);
	else modify(rson, x, v);
	w[rt] = w[ls] + w[rs];
}

void qm(int rt, int l, int r, int ql, int qr, int v) {
	if (w[rt].mx < v) return ;
	if (l == r) {
		w[rt].sum %= v;
		w[rt].mx %= v;
		return ;
	}
	int mid = (l + r) >> 1;
	if (ql <= mid) qm(lson, ql, qr, v);
	if (qr > mid) qm(rson, ql, qr, v);
	w[rt] = w[ls] + w[rs];
}

node query(int rt, int l, int r, int ql, int qr) {
	if (ql <= l && r <= qr) {
		return w[rt];
	}
	int mid = (l + r) >> 1;
	if (ql <= mid) {
		if (qr > mid) return query(lson, ql,qr) + query(rson, ql, qr);
		else return query(lson, ql, qr);
	}
	return query(rson, ql, qr);
}

int main() {
//	freopen(".in", "r", stdin);
//	freopen(".out", "w", stdout);

	cin.tie(0)->sync_with_stdio(0);

	cin >> n >> m;
	for (int i = 1; i <= n; i++) {
		cin >> a[i];
	} 
	build(root);
	while (m--) {
		int opt, l, r;
		cin >> opt >> l >> r;
		if (opt == 1) {
			cout << query(root, l, r).sum << '\n';
		}
		else if (opt == 2) {
			int x;
			cin >> x;
			qm(root, l, r, x);
		}
		else {
			modify(root, l, r);
		}
	}

	return 0;
}

Problem 3

题目

HDU5828

给定长度为 \(n\) 序列 \(a\),要求支持如下三个操作:

  1. 给定区间 \([l,r]\) 和一个整数 \(x\),对于所有 \(i \in [l,r]\),令 \(a_i \gets a_i + x\)。
  2. 给定区间 \([l,r]\),对于所有 \(i \in [l,r]\),令 \(a_i \gets \lfloor \sqrt{a_i} \rfloor\)。
  3. 给定区间 \([l,r]\),询问 \(\textstyle \sum_{i=l}^{r} a_i\)。

思路

直接考虑优化暴力。维护最大值和最小值。

如果区间 \([l,r]\) 内极差为 \(0\),即区间内所有数均相等,此时我们没有必要一个一个地修改,统一赋值为 \(\lfloor \sqrt{a_l} \rfloor\) 即可。

如果极差为 \(1\) 时,其实也没有必要逐个修改,我们分两种情况讨论。

  • 如果区间最大值是平方数,不妨设最大值为 \(t^2\),最小值为 \(t^2 - 1\),此时 \(\lfloor \sqrt{t^2} \rfloor = t\),\(\lfloor \sqrt{t^2 - 1} \rfloor = t - 1\),可以发现最大值和最小值都减小了 \(t^2 - t\),看作区间减操作。
  • 如果区间最大值不是平方数,那更简单了,此时区间开根后均相同,区间赋值即可。此时最小值一定不是平方数,或者即使是,开根后也相等。

剩下的暴力计算。

为了分析这道题的复杂度,我们定义关键点为极差大于 \(1\) 的节点。每次区间加操作,至多产生 \(O(\log n)\) 个关键点,受影响的只是边界上被部分覆盖的节点。对于区间开根操作,由上文分析可知,我们要对关键点暴力递归,直到该点变成非关键点(极差小于等于 \(1\))。设值域为 \(V\),则至多进行 \(O(\log \log V)\) 次操作即可将该点变为非关键点。所以总的时间复杂度为 \(O((n+m\log n)\log \log V)\)。

代码

#include <bits/stdc++.h>
using namespace std;

#define ll long long

const int N = 1e5 + 5;

#define root 1, 1, n
#define ls rt << 1
#define rs rt << 1 | 1
#define lson ls, l, mid
#define rson rs, mid + 1, r

struct node {
	ll sum;
	int siz;
	int tag, fg;
	int mn, mx;
	node () {
		sum = 0ll;
		siz = 0;
		tag = 0;
		mn = 0;
		mx = 0;
		fg = 0;
	}
	void init(int v) {
		sum = 1ll * v;
		siz = 1;
		mx = v;
		mn = v;
	    tag = 0;
    	fg = 0;
	}
};

node operator + (const node &x, const node &y) {
	node res;
	res.sum = x.sum + y.sum;
	res.siz = x.siz + y.siz;
	res.mx = max(x.mx, y.mx);
	res.mn = min(x.mn, y.mn);
	return res;
}

node w[N << 2];
int n, m;
int a[N];

void build(int rt, int l, int r) {
	if (l == r) {
		w[rt].init(a[l]);
		return ;
	}
	int mid = (l + r) >> 1;
	build(lson);
	build(rson);
	w[rt] = w[ls] + w[rs];
}

void colorfg(int rt, int v) {
	w[rt].fg = v;
	w[rt].tag = 0;
	w[rt].sum = 1ll * v * w[rt].siz;
	w[rt].mx = v;
	w[rt].mn = v;
}

void coloradd(int rt, int v) {
	if (w[rt].fg) {
		w[rt].fg += v;
        w[rt].sum = 1ll * w[rt].siz * w[rt].fg;
        w[rt].mx = w[rt].fg;
        w[rt].mn = w[rt].fg;
        w[rt].tag = 0;
        return ;
	}
	w[rt].sum += 1ll * w[rt].siz * v;
	w[rt].mx += v;
	w[rt].mn += v;
	w[rt].tag += v;
}

void pushdown(int rt) {
	if (w[rt].fg) {
		colorfg(ls, w[rt].fg);
		colorfg(rs, w[rt].fg);
		w[rt].fg = 0;
	}
	else {
		coloradd(ls, w[rt].tag);
		coloradd(rs, w[rt].tag);
		w[rt].tag = 0;
	}
}

void add(int rt, int l, int r, int ql, int qr, int v) {
	if (ql <= l && r <= qr) {
		coloradd(rt, v);
		return ;
	}
	pushdown(rt);
	int mid = (l + r) >> 1;
	if (ql <= mid) add(lson, ql, qr, v);
	if (qr > mid) add(rson, ql, qr, v);
	w[rt] = w[rs] + w[ls];
}

void modify(int rt, int l, int r, int ql, int qr) {
	if (l == r) {
		int x = (int)(sqrt(w[rt].mx));
		w[rt].init(x);
		return ;
	}
	if (ql <= l && r <= qr) {
		pushdown(rt);
		int x = sqrt(w[rt].mx);
		if (w[rt].mx == w[rt].mn) {
			colorfg(rt, x);
			return ;
		}
		else if (w[rt].mx == w[rt].mn + 1) {
			if (x * x == w[rt].mx) coloradd(rt, x - w[rt].mx);
			else colorfg(rt, x);
			return ;
		}
	}
	pushdown(rt);
	int mid = (l + r) >> 1;
	if (ql <= mid) modify(lson, ql, qr);
	if (qr > mid) modify(rson, ql, qr);
	w[rt] = w[ls] + w[rs];
}

node query(int rt, int l, int r, int ql, int qr) {
	if (ql <= l && r <= qr) {
		return w[rt];
	}
	pushdown(rt);
	int mid = (l + r) >> 1;
	if (ql <= mid) {
		if (qr > mid) return query(lson, ql, qr) + query(rson, ql, qr);
		else return query(lson, ql, qr);
	}
	return query(rson, ql, qr);
}

void solve() {
	cin >> n >> m;
	for (int i  = 1; i <= n; i++) cin >> a[i]; 
	build(root);
	while (m--) {
		int opt, l, r;
		cin >> opt >> l >> r;
		if (opt == 1) {
			int x;
			cin >> x;
			add(root, l, r, x);
		}
		else if (opt == 2) {
			modify(root, l, r);
		}
		else {
			cout << query(root, l, r).sum << '\n';
		}
	}
}

int main() {
//	freopen(".in", "r", stdin);
//	freopen(".out", "w", stdout);

	cin.tie(0)->sync_with_stdio(0);

	int t;
	cin >> t;
	while (t--) {
		solve();
	}

	return 0;
}
posted @ 2026-08-08 15:45  chaqjs  阅读(7)  评论(0)    收藏  举报