算法进阶杂题选讲:从期望DP到树链剖分与图论构造
网格随机切割期望求解
给定一个放置在平面直角坐标系中左下角为 (0,0)、右上角为 (n,m) 的矩形纸片。每次等概率随机选择一条穿过纸片内部且平行于坐标轴的整点直线,沿该直线裁剪并丢弃下侧或右侧部分。求期望多少次操作后剩余面积小于 k,结果对 10^9+7 取模。
解题思路:
可以将所有可能的切割线看作一个排列,以此判断哪些切割是实际有效的。对于横向(或纵向)的第 i 条切割线,其有效的条件为:排在它之前的同类切割线不在 [1, i-1] 范围内,且异类切割线不在 [1, floor(S/i)] 范围内。满足此情况的概率可以表示为 1/i 乘以 1/min(n, S/i)。分别枚举横向和纵向的有效切割概率并累加即可。
代码实现:
#include<bits/stdc++.h>
using namespace std;
typedef long long ll;
const ll MOD = 1e9 + 7;
const int MAXN = 2e6 + 5;
ll inv[MAXN];
ll mod_pow(ll base, ll exp) {
ll res = 1;
while(exp) {
if(exp & 1) res = res * base % MOD;
base = base * base % MOD;
exp >>= 1;
}
return res;
}
ll fast_read() {
ll x = 0, f = 1; char ch = getchar();
while(ch < '0' || ch > '9') { if(ch == '-') f = -1; ch = getchar(); }
while(ch >= '0' && ch <= '9') { x = (x << 1) + (x << 3) + (ch ^ 48); ch = getchar(); }
return x * f;
}
void solve() {
ll n = fast_read(), m = fast_read(), S = fast_read() - 1;
if(n * m <= S) { puts("0"); return; }
ll ans = 1;
for(ll i = 1; i < n; ++i) {
ll j = min(S / i, m);
if(j == m) continue;
ans = (ans + inv[i + j]) % MOD;
}
for(ll i = 1; i < m; ++i) {
ll j = min(S / i, n);
if(j == n) continue;
ans = (ans + inv[i + j]) % MOD;
}
printf("%lld\n", ans);
}
int main() {
inv[1] = 1;
for(ll i = 2; i < MAXN; ++i) inv[i] = (MOD - MOD / i) * inv[MOD % i] % MOD;
int T = fast_read();
while(T--) solve();
return 0;
}
序列划分的凸优化策略
给定长度为 n 的数列和参数 k, s,将数列划分为 k 段,最大化元素和不小于 s 的段数。
解题思路:
直接对合法段数进行 WQS 二分发现不具备凸性。转换思路,二分至少有 mid 个段总和 ≥ s 的答案。定义函数 g(x) 为满足 ≥ s 的段数为 x 时,能够划分出的最多段数。该函数不易直接求解,故修改定义为 g(x) = 总长度 - 最多段数,此时该函数具备凸性,可以使用 WQS 二分进行求解。
代码实现:
#include<bits/stdc++.h>
using namespace std;
typedef long long ll;
const int MAXN = 3e5 + 5;
ll prefix[MAXN], pre_pos[MAXN];
pair<ll, ll> dp[MAXN];
int seq_len, seg_cnt; ll threshold;
ll fast_read() {
ll x = 0, f = 1; char ch = getchar();
while(ch < '0' || ch > '9') { if(ch == '-') f = -1; ch = getchar(); }
while(ch >= '0' && ch <= '9') { x = (x << 1) + (x << 3) + (ch ^ 48); ch = getchar(); }
return x * f;
}
bool validate(ll limit) {
ll lo = 1, hi = seq_len, mid = 0, res = 0;
auto compute = [&]() {
for(int i = 1; i <= seq_len; ++i) {
dp[i] = dp[i - 1];
if(pre_pos[i]) dp[i] = min(dp[i], make_pair(dp[pre_pos[i] - 1].first + (i - pre_pos[i]) - limit, dp[pre_pos[i] - 1].second + 1));
}
};
while(lo <= hi) {
mid = (lo + hi) >> 1; compute();
if(dp[seq_len].second <= limit) { lo = mid + 1; res = mid; }
else hi = mid - 1;
}
compute();
return dp[seq_len].first + res * limit <= seq_len - seg_cnt;
}
int main() {
seq_len = fast_read(); seg_cnt = fast_read(); threshold = fast_read();
for(int i = 1, ptr = 0; i <= seq_len; ++i) {
prefix[i] = fast_read() + prefix[i - 1];
while(prefix[i] - prefix[ptr] >= threshold) ptr++;
pre_pos[i] = ptr;
}
ll lo = 1, hi = seg_cnt, ans = 0;
while(lo <= hi) {
ll mid = (lo + hi) >> 1;
if(validate(mid)) { lo = mid + 1; ans = mid; }
else hi = mid - 1;
}
printf("%lld\n", ans);
return 0;
}
格点三角形的代数构造
多组测试,每次给定 S,构造满足三个顶点均为格点、面积为 S/2,且内部恰好仅包含一个边长为 1 的格点正方形的三角形。
解题思路:
分奇偶讨论。当 S 为偶数时,将底边长度设为 2,端点置于 (0,0) 和 (2,0),第三个顶点置于直线 x-y=0 上,即可保证仅含一个正方形。当 S 为奇数时,由叉积公式 S=|x1y2-x2y1| 推导奇偶性,不妨固定一点为 (3,1),则有 3y2-x2=S。设 y2=x2+1,解得 2x2=S-1,从而直接计算出整数坐标。
代码实现:
#include<bits/stdc++.h>
using namespace std;
int fast_read() {
int x = 0, f = 1; char ch = getchar();
while(ch < '0' || ch > '9') { if(ch == '-') f = -1; ch = getchar(); }
while(ch >= '0' && ch <= '9') { x = (x << 1) + (x << 3) + (ch ^ 48); ch = getchar(); }
return x * f;
}
int main() {
int T = fast_read();
while(T--) {
int S = fast_read();
if(S == 2 || (S < 9 && (S & 1))) { puts("No"); continue; }
puts("Yes");
if((S & 1) == 0) printf("0 0 2 0 %d %d\n", S / 2, S / 2);
else printf("0 0 3 1 %d %d\n", (S - 9) / 2 + 3, (S - 9) / 2 + 4);
}
return 0;
}
树链剖分与带权重心动态查询
给定一棵 n 个节点的树,支持两种操作:对树上某条链或某个子树的节点权值加 w。每次操作后求出当前树的带权重心(若有多个取深度最浅者)。
解题思路:
利用树链剖分与线段树维护区间修改与权值查询。寻找带权重心时,利用关键性质:深度最浅的重心,其子树权值和必定严格大于总权值的一半。若不满足,向父节点移动必然更优。具体实现中,将每个点按其权值展开,通过线段树找到权值中位数对应的节点位置,然后利用倍增法向上跳转,找到首个满足子树权值和大于总权值一半的节点,即为答案。
代码实现:
#include<bits/stdc++.h>
#define int long long
using namespace std;
const int MAXN = 2e5 + 5;
int node_cnt, ops;
vector<int> adj_list[MAXN];
int anc[21][MAXN], sub_sz[MAXN], heavy_son[MAXN], depth[MAXN];
int dfn[MAXN], low[MAXN], time_stamp, chain_top[MAXN], rev[MAXN];
void dfs_first(int u, int f) {
anc[0][u] = f; sub_sz[u] = 1; depth[u] = depth[f] + 1;
for(int i = 1; i <= 20; ++i) anc[i][u] = anc[i-1][anc[i-1][u]];
for(int v : adj_list[u]) {
if(v != f) {
dfs_first(v, u); sub_sz[u] += sub_sz[v];
heavy_son[u] = sub_sz[v] > sub_sz[heavy_son[u]] ? v : heavy_son[u];
}
}
}
void dfs_second(int u, int t) {
chain_top[u] = t; dfn[u] = ++time_stamp; rev[time_stamp] = u; low[u] = sub_sz[u] + dfn[u] - 1;
if(!heavy_son[u]) return; dfs_second(heavy_son[u], t);
for(int v : adj_list[u]) if(v != anc[0][u] && v != heavy_son[u]) dfs_second(v, v);
}
struct SegTree {
struct Node { int sum, l, r, tag; } tree[MAXN << 2];
Node merge(Node a, Node b) { return {a.sum + b.sum, a.l, b.r}; }
void apply_tag(int p, int val) { tree[p].tag += val; tree[p].sum += (tree[p].r - tree[p].l + 1) * val; }
void push_down(int p) { apply_tag(p<<1, tree[p].tag); apply_tag(p<<1|1, tree[p].tag); tree[p].tag = 0; }
void build(int l, int r, int p) {
tree[p] = {0, l, r, 0};
if(l == r) return; int mid = (l + r) >> 1;
build(l, mid, p<<1); build(mid+1, r, p<<1|1);
}
void update(int l, int r, int s, int t, int p, int val) {
if(s <= l && r <= t) { apply_tag(p, val); return; }
int mid = (l + r) >> 1; push_down(p);
if(s <= mid) update(l, mid, s, t, p<<1, val);
if(t > mid) update(mid+1, r, s, t, p<<1|1, val);
tree[p] = merge(tree[p<<1], tree[p<<1|1]);
}
int query(int l, int r, int s, int t, int p) {
if(s <= l && r <= t) return tree[p].sum;
int mid = (l + r) >> 1; push_down(p); int res = 0;
if(s <= mid) res += query(l, mid, s, t, p<<1);
if(t > mid) res += query(mid+1, r, s, t, p<<1|1);
return res;
}
int find_kth(int l, int r, int p, int k) {
if(l == r) return l;
int mid = (l + r) >> 1; push_down(p);
if(tree[p<<1].sum >= k) return find_kth(l, mid, p<<1, k);
return find_kth(mid+1, r, p<<1|1, k - tree[p<<1].sum);
}
} seg;
signed main() {
scanf("%lld", &node_cnt); int u, v, opt;
for(int i = 1; i < node_cnt; ++i) { scanf("%lld%lld", &u, &v); adj_list[u].push_back(v); adj_list[v].push_back(u); }
dfs_first(1, 0); dfs_second(1, 1); seg.build(1, node_cnt, 1);
scanf("%lld", &ops);
while(ops--) {
scanf("%lld%lld", &opt, &u);
if(opt == 1) seg.update(1, node_cnt, dfn[u], low[u], 1, 1);
else {
scanf("%lld", &v);
while(chain_top[u] != chain_top[v]) {
if(depth[chain_top[u]] < depth[chain_top[v]]) swap(u, v);
seg.update(1, node_cnt, dfn[chain_top[u]], dfn[u], 1, 1);
u = anc[0][chain_top[u]];
}
if(dfn[u] > dfn[v]) swap(u, v);
seg.update(1, node_cnt, dfn[u], dfn[v], 1, 1);
}
int target = seg.tree[1].sum / 2 + 1;
u = rev[seg.find_kth(1, node_cnt, 1, target)];
if(seg.query(1, node_cnt, dfn[u], low[u], 1) >= target) { printf("%lld\n", u); continue; }
for(int i = 20; i >= 0; --i) if(anc[i][u] && seg.query(1, node_cnt, dfn[anc[i][u]], low[anc[i][u]], 1) < target) u = anc[i][u];
printf("%lld\n", anc[0][u]);
}
return 0;
}
二分图模型下的连锁引爆问题
在一个 n×m 的网格上有 k 个宝石和 b 个炸弹。炸弹可选择横向或纵向引爆,销毁该行或列的宝石,并触发该行或列上的其他炸弹。选择每个炸弹方向并引爆一个,求链式反应后最多炸掉的宝石数。
解题思路:
若某行存在炸弹横向引爆,该行其余炸弹必纵向引爆。将行与列作为二分图中的节点,炸弹视为连接行列的边。引爆过程等价于遍历图中的节点与边。若连通块含环,则所有点均可遍历;若为树形结构,必然存在一个叶子节点无法走到。据此分类统计连通块内宝石数的最大值。
代码实现:
#include<bits/stdc++.h>
using namespace std;
const int MAXN = 1e4 + 5;
int rows, cols;
char grid[MAXN / 3][MAXN / 3];
int parent[MAXN], comp_sz[MAXN], edge_cnt[MAXN], degree[MAXN];
int gem_cnt[3][MAXN];
int find_set(int x) { return parent[x] == x ? x : parent[x] = find_set(parent[x]); }
int fast_read() {
int x = 0, f = 1; char ch = getchar();
while(ch < '0' || ch > '9') { if(ch == '-') f = -1; ch = getchar(); }
while(ch >= '0' && ch <= '9') { x = (x << 1) + (x << 3) + (ch ^ 48); ch = getchar(); }
return x * f;
}
int main() {
rows = fast_read(); cols = fast_read(); fast_read(); fast_read();
for(int i = 1; i <= rows + cols; ++i) { parent[i] = i; comp_sz[i] = 1; edge_cnt[i] = 0; }
for(int i = 1; i <= rows; ++i) {
scanf("%s", grid[i] + 1);
for(int j = 1; j <= cols; ++j) {
if(grid[i][j] == 'b') {
degree[i]++; degree[j + rows]++;
int u = find_set(i), v = find_set(j + rows);
if(u != v) { parent[u] = v; comp_sz[v] += comp_sz[u]; edge_cnt[v] += edge_cnt[u]; }
edge_cnt[v]++;
}
}
}
for(int i = 1; i <= rows; ++i) {
for(int j = 1; j <= cols; ++j) {
if(grid[i][j] == 'k') {
int u = find_set(i), v = find_set(j + rows);
if(u == v) gem_cnt[0][u]++;
else { gem_cnt[1][u]++; gem_cnt[1][v]++; gem_cnt[2][i]++; gem_cnt[2][j + rows]++; }
}
}
}
int ans = 0;
for(int i = 1; i <= rows + cols; ++i)
if(find_set(i) == i && edge_cnt[i] >= comp_sz[i]) ans = max(ans, gem_cnt[0][i] + gem_cnt[1][i]);
for(int i = 1; i <= rows + cols; ++i)
if(degree[i] == 1) { int f = find_set(i); ans = max(ans, gem_cnt[0][f] + gem_cnt[1][f] - gem_cnt[2][i]); }
printf("%d\n", ans);
return 0;
}
基于逆序对的单峰序列构造
定义序列权值为使其变为单峰或单谷的最少相邻交换次数。给定无重复元素的数组,求每个前缀的权值。
解题思路:
单峰与单谷情况对称,通过取反映射即可相互转化。以单谷为例,假设已知每个元素在谷左右的归属,答案即为谷左侧元素取反后的逆序对数。初始假设全在右侧,从小到大考虑元素,利用树状数组维护前缀信息,通过倍增法确定每个元素划归左侧的最优决策点,动态计算并更新最小逆序对数。
代码实现:
#include<bits/stdc++.h>
#define ll long long
using namespace std;
const int MAXN = 2e5 + 5;
int seq_len, arr[MAXN];
vector<int> coords;
struct Fenwick {
int bit[MAXN]; void init() { memset(bit, 0, sizeof(bit)); }
int low_bit(int x) { return x & (-x); }
int query(int x) { int s = 0; for(; x; x -= low_bit(x)) s += bit[x]; return s; }
void update(int x, int v) { for(; x <= seq_len; x += low_bit(x)) bit[x] += v; }
} ft;
int left_inv[MAXN], pos[MAXN];
ll res[MAXN];
vector<int> ops_at[MAXN];
int fast_read() {
int x = 0, f = 1; char ch = getchar();
while(ch < '0' || ch > '9') { if(ch == '-') f = -1; ch = getchar(); }
while(ch >= '0' && ch <= '9') { x = (x << 1) + (x << 3) + (ch ^ 48); ch = getchar(); }
return x * f;
}
void process() {
ft.init();
for(int i = 1; i <= seq_len; ++i) { left_inv[i] = ft.query(arr[i] - 1); ft.update(arr[i], 1); pos[arr[i]] = i; }
ft.init();
for(int i = 1; i <= seq_len; ++i) ops_at[i].clear();
for(int i = 1; i <= seq_len; ++i) {
int p = pos[i], cur = p;
for(int j = 20; j >= 0; --j) {
int nxt = cur + (1 << j);
if(nxt <= seq_len && ft.query(nxt) - left_inv[p] <= left_inv[p]) cur = nxt;
}
ops_at[cur + 1].push_back(i); ft.update(p, 1);
}
ft.init(); ll cur_res = 0;
for(int i = 1; i <= seq_len; ++i) {
for(int j : ops_at[i]) ft.update(j, -1);
cur_res += ft.query(seq_len) - ft.query(arr[i]);
res[i] = min(res[i], cur_res); ft.update(arr[i], 1);
}
}
int main() {
seq_len = fast_read();
for(int i = 1; i <= seq_len; ++i) coords.push_back(arr[i] = fast_read());
sort(coords.begin(), coords.end());
coords.erase(unique(coords.begin(), coords.end()), coords.end());
for(int i = 1; i <= seq_len; ++i) arr[i] = lower_bound(coords.begin(), coords.end(), arr[i]) - coords.begin() + 1;
memset(res, 0x3f, sizeof(res));
process();
for(int i = 1; i <= seq_len; ++i) arr[i] = seq_len - arr[i] + 1;
process();
for(int i = 1; i <= seq_len; ++i) printf("%lld\n", res[i]);
return 0;
}