题解:B 小W与伙伴招募(text on 2026.10.6)

\(\textsf{B tj}\) 小W与伙伴招募

题目大意

有 \(n\) 月,每月需要支付 \(c_i\) 的钻石。

对于每个月,有 \(m\) 种购买方式,均为花费 \(a_i\) 的成本购买 \(1\) 钻石,每月至多买 \(b_i\) 个。

若 \(b_i=-1\),则购买无上限。

没有支付的钻石可以留着。

问最小花费。

\(n,m\le 10^5\),保证至少存在一个 \(i\) 使得 \(b_i=-1\)。

思路 \(1\) 贪心

我们很容易写出贪心:

将 \(a,b\) 按 \(a\) 从小到大排序,找到第一个 \(b_i=-1\) 作为上界。

用一个数组 \(h\) 记录当前存货,对于每个月取 \(h\) 最优,然后将下个月的存货加到 \(h\) 里面。

时间复杂度 \(O(nm)\),可以拿 \(\orange{\texttt{60 pts}}\)。

$\orange{\texttt{60 pts}}$ 代码
//60pts O(nm)
#include<bits/stdc++.h>
#define int long long
using namespace std;
const int N = 2e5+10;
int n, m, endr, sum;
struct node { int a, b; } p[N];
int c[N], h[N];
bool cmp(node A, node B){
	return A.a < B.a;
}
signed main(){
	freopen("1539.in","r",stdin);
	freopen("1539.out","w",stdout);
	ios::sync_with_stdio(false);
	cin.tie(0), cout.tie(0);
	cin >> n >> m;
	for(int i = 1; i <= n; i++) cin >> c[i];
	for(int i = 1; i <= m; i++)
		cin >> p[i].a >> p[i].b;
	sort(p+1, p+m+1, cmp);
	for(int i = 1; i <= m; i++){
		if(p[i].b == -1){
			endr = i;
			break;
		}
		h[i] = p[i].b;
	}
	h[endr] = -1;
	for(int i = 1; i <= n; i++){
		int C = c[i], j = 1, ans = 0;
		while(j < endr && h[j] <= C){
			C -= h[j];
			ans += p[j].a * h[j];
			h[j] = p[j].b;
			j++;
		}
		ans += p[j].a * C;
		sum += ans;
		if(h[j] != -1) h[j] += p[j].b - C;
		for(int k = j+1; k < endr; k++)
			h[k] += p[k].b;
	}
	cout << sum;
	return 0;
} 

可是 \(\orange{\texttt{60 pts}}\) 不好吃,我要正解。

考虑对 \(h\) 进行优化。

思路 \(2\) 线段树

众所周知,我最喜欢线段树了,所以我经常挂分

于是我们用线段树维护 \(h\)。

但是我们需要多棵,一棵维护 \(h_i\),用来应付库存;一棵维护 \(a_i\cdot h_i\),用来应付花费。

需要进行以下操作:

  • build():初始化线段树。

    主要将线段树设为初始值。

  • findh(x):快速找到 \(<x\) 的 \(i\),使得 \(\displaystyle\sum_{l=1}^i h_i < x\) 且 \(\displaystyle\sum_{l=1}^{i+1} h_i \ge x\)。

    用于找到 \(c_i\) 中需要全部购买的钻石。

    可以在第一棵树上线段树二分解决。

  • query1(x,y):区间查询第一棵线段树的 \((x,y)\) 区间和。

    用于算出批量购买后,还剩下多少钻石需要购买。

    线段树基本操作。

  • query2(x,y):区间查询第二棵线段树的 \((x,y)\) 区间和。

    用于算出批量购买的成本。

    依旧基本操作。

  • smtcpy(x,y):将线段树的 \((x,y)\) 区间内重置。

    首先了解,第一棵线段树的初始值为 \(\sum b_i\),第二棵为 \(\sum a_ib_i\)。

    所以这个函数用于将线段树的区间重置为初始值,对于那些需要全部购买的钻石 \([1,x]\)。

    于是我继续使用线段树维护 \(a_i, b_i, a_ib_i\)(可能没必要这么多,前缀和即可)

    由于是区间修改,所以需要一个重置的懒标记。

  • update(x):修改 \(x\) 这个点。

    将剩下的那一类的(即 \(x+1\))进行单点修改。

    基本操作。

  • cpyadd(x,y):将线段树内 \((x,y)\) 区间加初始值。

    因为 \(x+2\) 以后的库存没有使用,所以将下一天的库存加上初始值。

    所以再新建一个懒标记下传即可。

    区间重置 smtcpy(x,y) 时,记得将区间加初始值标记删除。

主函数逻辑

前面基本和贪心思路一样。

设第一个 \(b_i=-1\) 的 \(i\) 为 \(\text{endr}\)。

注意特判 \(\text{endr}=1\) 的情况。

循环:

build(1, 1, endr-1); //初始化
for(int i = 1; i <= n; i++){
	int C = c[i], ans = 0;
	if(C >= tr[1].h){ //购买有限的全部库存还不够/刚好 c[i]
		ans = (C - tr[1].h) * P[endr].a + tr[1].pah; //计算成本
		tr[1].cpytag = 1;
		tr[1].h = tr[1].pb;
		tr[1].pah = tr[1].pab;
		//由于全部购买,所以全部需要重置
	} else {
		int x = findh(1, C, 1, endr-1); //找到需要全部购买的库存
		C -= query1(1, 1, x, 1, endr-1); //剩下多少钻石需要购买
		ans = query2(1, 1, x, 1, endr-1) + P[x+1].a * C; //批量购买的成本
		smtcpy(1, 1, x, 1, endr-1); //重置批量购买的
		if(x + 1 < endr)
			update(1, x+1, 1, endr-1, P[x+1].b-C);
		cpyadd(1, x+2, endr-1, 1, endr-1);
		//可以与上面的函数解释一起理解
	}
	sum += ans; //累加答案
}

构造线段树

struct tree {
	int h, pah, pa, pb, pab, cpytag, addtag;
} tr[4*N];
  • h:维护 \(\sum h_i\)。
  • pah:维护 \(\sum a_i\cdot h_i\)。
  • pa pb pab:维护 \(\sum a_i,\sum b_i,\sum a_ib_i\)。
  • cpytag:维护重置懒标记,可以为 \(0/1\)。
  • addtag:维护增加初始值懒标记。

\(\textsf{code}\)

#include<bits/stdc++.h>
#define int long long
using namespace std;
const int N = 2e5+10;
int n, m, endr, sum;
struct node { int a, b; } P[N];
struct tree {
	int h, pah, pa, pb, pab, cpytag, addtag;
} tr[4*N];
int c[N];
bool cmp(node A, node B){
	return A.a < B.a;
}
int lp(int x){
	return x << 1;
}
int rp(int x){
	return x << 1 | 1;
}
void pushup(int p){
	tr[p].h = tr[lp(p)].h + tr[rp(p)].h;
	tr[p].pah = tr[lp(p)].pah + tr[rp(p)].pah;
	tr[p].pa = tr[lp(p)].pa + tr[rp(p)].pa;
	tr[p].pb = tr[lp(p)].pb + tr[rp(p)].pb;
	tr[p].pab = tr[lp(p)].pab + tr[rp(p)].pab;
}
void pushdown(int p){
	if(tr[p].cpytag){
		tr[lp(p)].h = tr[lp(p)].pb;
		tr[lp(p)].pah = tr[lp(p)].pab;
		tr[rp(p)].h = tr[rp(p)].pb;
		tr[rp(p)].pah = tr[rp(p)].pab;
		tr[lp(p)].cpytag = 1;
		tr[rp(p)].cpytag = 1;
		tr[lp(p)].addtag = 0;
		tr[rp(p)].addtag = 0;
		tr[p].cpytag = 0;
	}
	if(tr[p].addtag){
		tr[lp(p)].h += tr[p].addtag * tr[lp(p)].pb;
		tr[lp(p)].pah += tr[p].addtag * tr[lp(p)].pab;
		tr[rp(p)].h += tr[p].addtag * tr[rp(p)].pb;
		tr[rp(p)].pah += tr[p].addtag * tr[rp(p)].pab;
		tr[lp(p)].addtag += tr[p].addtag;
		tr[rp(p)].addtag += tr[p].addtag;
		tr[p].addtag = 0;
	}
}
void build(int p, int l, int r){
	if(l == r){
		tr[p].h = P[l].b;
		tr[p].pah = P[l].a * P[l].b;
		tr[p].pa = P[l].a;
		tr[p].pb = P[l].b;
		tr[p].pab = P[l].a * P[l].b;
		return;
	}
	int mid = (l + r) >> 1;
	build(lp(p), l, mid);
	build(rp(p), mid+1, r);
	pushup(p);
}
int findh(int p, int x, int l, int r){
	if(l == r) return l-1;
	pushdown(p);
	int mid = (l + r) >> 1;
	int lef = tr[lp(p)].h;
	if(lef >= x) return findh(lp(p), x, l, mid);
	else return findh(rp(p), x-lef, mid+1, r);
}
int query1(int p, int x, int y, int l, int r){
	if(x > y) return 0;
	if(x <= l && r <= y) return tr[p].h;
	pushdown(p);
	int mid = (l + r) >> 1;
	int ans = 0;
	if(x <= mid) ans += query1(lp(p), x, y, l, mid);
	if(y > mid) ans += query1(rp(p), x, y, mid+1, r);
	return ans;
}
int query2(int p, int x, int y, int l, int r){
	if(x > y) return 0;
	if(x <= l && r <= y) return tr[p].pah;
	pushdown(p);
	int mid = (l + r) >> 1;
	int ans = 0;
	if(x <= mid) ans += query2(lp(p), x, y, l, mid);
	if(y > mid) ans += query2(rp(p), x, y, mid+1, r);
	return ans;
}
void smtcpy(int p, int x, int y, int l, int r){
	if(x > y) return;
	if(x <= l && r <= y){
		tr[p].h = tr[p].pb;
		tr[p].pah = tr[p].pab;
		tr[p].cpytag = 1;
		tr[p].addtag = 0;
		return;
	}
	pushdown(p);
	int mid = (l + r) >> 1;
	if(x <= mid) smtcpy(lp(p), x, y, l, mid);
	if(y > mid) smtcpy(rp(p), x, y, mid+1, r);
	pushup(p);
}
void update(int p, int x, int l, int r, int k){
	if(l == r){
		tr[p].h += k;
		tr[p].pah += k * tr[p].pa;
		return;
	}
	pushdown(p);
	int mid = (l + r) >> 1;
	if(x <= mid) update(lp(p), x, l, mid, k);
	else update(rp(p), x, mid+1, r, k);
	pushup(p);
}
void cpyadd(int p, int x, int y, int l, int r){
	if(x > y) return;
	if(x <= l && r <= y){
		tr[p].h += tr[p].pb;
		tr[p].pah += tr[p].pab;
		tr[p].addtag++;
		return;
	}
	pushdown(p);
	int mid = (l + r) >> 1;
	if(x <= mid) cpyadd(lp(p), x, y, l, mid);
	if(y > mid) cpyadd(rp(p), x, y, mid+1, r);
	pushup(p);
}
signed main(){
	freopen("1539.in","r",stdin);
	freopen("1539.out","w",stdout);
	ios::sync_with_stdio(false);
	cin.tie(0), cout.tie(0);
	cin >> n >> m;
	for(int i = 1; i <= n; i++) cin >> c[i];
	for(int i = 1; i <= m; i++)
		cin >> P[i].a >> P[i].b;
	sort(P+1, P+m+1, cmp);
	for(int i = 1; i <= m; i++)
		if(P[i].b == -1){
			endr = i;
			break;
		}
	if(endr == 1){
		for(int i = 1; i <= n; i++) sum += c[i];
		cout << sum * P[1].a;
		return 0;
	}
	build(1, 1, endr-1);
	for(int i = 1; i <= n; i++){
		int C = c[i], ans = 0;
		if(C >= tr[1].h){
			ans = (C - tr[1].h) * P[endr].a + tr[1].pah;
			tr[1].cpytag = 1;
			tr[1].h = tr[1].pb;
			tr[1].pah = tr[1].pab;
		} else {
			int x = findh(1, C, 1, endr-1);
			C -= query1(1, 1, x, 1, endr-1);
			ans = query2(1, 1, x, 1, endr-1) + P[x+1].a * C;
			smtcpy(1, 1, x, 1, endr-1);
			if(x + 1 < endr)
				update(1, x+1, 1, endr-1, P[x+1].b-C);
			cpyadd(1, x+2, endr-1, 1, endr-1);
		}
		sum += ans;
	}
	cout << sum;
	return 0;
}

提示

需要判断线段树的参数是否合法,防止死循环。

复杂度 \(O(m+n\log m)\)

posted @ 2026-10-06 13:44  qjy123  阅读(60)  评论(0)    收藏  举报