题解: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)\)

浙公网安备 33010602011771号