P3826 [NOI2017] 蔬菜 题解

前置:堆、反悔贪心

题目链接:https://www.luogu.com.cn/problem/P3826

题意

有 n 种蔬菜,每天最多卖 m 个单位。

第 i 种蔬菜有四个参数:

  • a_i:每卖出一个单位获得 a_i 收益;
  • s_i:第一次卖这种蔬菜时,额外获得 s_i 收益;
  • c_i:初始库存;
  • x_i:每天结束时会坏掉 x_i 个单位,直到卖完或坏完为止。

每种蔬菜会按天分批变质,越早不卖,后面能卖的数量就越少。

现在有 k 次询问,每次给出天数 p_j,问在前 p_j 天内,合理安排每天卖哪些蔬菜,最多能获得多少收益。

思路

只要让蔬菜在坏掉前卖掉就行,所以我们按照坏掉的时间排序下来,把每种蔬菜放到它最晚还能被卖掉的那一天。

对第 i 种蔬菜:

  • 库存是 c[i];
  • 每天坏 x[i] 个;
  • 如果 x[i] == 0,它不会坏,所以放到第 p 天;
  • 否则,它最晚能撑到 \(\left\lceil \frac{c_i}{x_i} \right\rceil\) 天。

那我们可以开一个数组存一下最晚在第 i 天还能被卖掉的蔬菜,这样第一步就搞定了。

对应代码:

cpp 复制代码
for (int i = 1; i <= n; i ++) {
    if (!x[i]) {
        d[p].push_back(i);
    } else {
        d[min(p, (c[i] + x[i] - 1) / x[i])].push_back(i);
    }
}

价值分为销量和单个物品的价值两部分,我们分别处理。

我们可以先算第 p 天的安排,因为之前的可以直接倒推。

怎么安排销量呢?

每一天,把所有最晚只能撑到这一天的蔬菜以价值为关键字加入大根堆。

怎么算价值呢?

根据题意,第一次卖某种蔬菜时,收益是 \(a_i + s_i\),之后每次再卖,收益只有 \(a_i\)。

所以我们可以加一个 vis 标记:

  • 如果这种蔬菜还没卖过,就先卖一个,收益是 a + s;
  • 如果已经卖过,就尽量多卖普通单位。

对于已经卖过的我们可以加入栈中,方便后面直接取出,因为同一种蔬菜我们可以等后面再卖。

那我们后面卖的时候还要计算每种蔬菜还剩下多少。这个很好算,用总量减去卖掉的再减去坏掉的即可。也就是c_i - \\text{used}_i - (i-1)x_i其中 used 表示已经卖过了的。

如果本轮还没卖完,我们就把剩余部分放回栈 ,之后再放回大根堆,继续参与后面更早天数的选择。

那我们此时已经算出了 ans[p],也就是最多 p 天时的最大收益。

然后我们直接往前倒推就行。

用 sm 表示总共卖了多少个蔬菜,如果只卖 i 天,最多只能卖 \(i \times m\)。

  • 如果 sm > i * m,说明多卖了,必须退掉一些;
  • 退掉的时候,优先退收益最小的。
    那我们把之前的大根堆清空,拿来算最小负收益即可。由于是负数,我们用大根堆可以取得减少价值的最小值(因为是负数所以相当于选择最趋近于零的那个,也就是最小值)。
cpp 复制代码
{-s[i] - a[i], i}   // 只卖了一个,退掉就失去 a + s
{-a[i], i}          // 卖了多个,退掉一个只失去 a

每次从堆里取出收益最小的单位退掉:

  • 如果这种蔬菜卖了超过 1 个,就退普通单位;
  • 如果只剩 1 个,再退就会失去第一次销售的额外收益 s。

这样就能从 ans[p] 推出:

\ans\[p-1, ansp-2, \dots, ans1 \]

时间复杂度\(O((n + n \times m)logn)\)

参考代码

cpp 复制代码
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const ll maxn = 1e5 + 10;

inline int read() {
    int x = 0;
    bool t = false;
    char ch = getchar();
    while ((ch < '0' || ch > '9') && ch != '-') {
        ch = getchar();
    }
    if (ch == '-') {
        t = true;
        ch = getchar();
    }
    while (ch <= '9' && ch >= '0') {
        x = x * 10 + ch - 48;
        ch = getchar();
    }
    return t ? -x : x;
}

struct node {
    int v, i;
} stk[maxn];

bool operator<(node a, node b) {
    return a.v < b.v;
}

int p = 1e5, tp, sm;
int n, m, qry, a[maxn], s[maxn], c[maxn], x[maxn], used[maxn];
ll ans[maxn];
bool vis[maxn];
vector<int> d[maxn];
priority_queue<node> q;

int main() {
    n = read();
    m = read();
    qry = read();
    for (int i = 1; i <= n; i ++) {
        a[i] = read();
        s[i] = read();
        c[i] = read();
        x[i] = read();
    }
    
    for (int i = 1; i <= n; i ++) {
        if (!x[i]) {
            d[p].push_back(i);
        }
        else {
            d[min(p, (c[i] + x[i] - 1) / x[i])].push_back(i);
        }
    }
    
    for (int i = p; i; i --) {
        for (int j = 0, l = d[i].size(); j < l; ++ j) {
            q.push((node){a[d[i][j]] + s[d[i][j]], d[i][j]});
        }
        
        if (q.empty()) {
            continue;
        }
        
        for (int j = m; j && !q.empty(); ) {
            node u = q.top();
            q.pop();
            if (!vis[u.i]) {
                vis[u.i] = true;
                ans[p] += u.v;
                used[u.i] += 1;
                j --;
                if (c[u.i] > 1) {
                    q.push((node){a[u.i], u.i});
                }
            }
            else {
                int rest = min(j, c[u.i] - used[u.i] - (i - 1) * x[u.i]);
                ans[p] += 1ll * rest * u.v;
                used[u.i] += rest;
                j -= rest;
                if (used[u.i] != c[u.i]) {
                    stk[++ tp] = (node){a[u.i], u.i};
                }
            }
        }
        while (tp) {
            q.push(stk[tp --]);
        }
    }
    
    while (!q.empty()) {
        q.pop();
    }
    
    for (int i = 1; i <= n; i ++) {
        sm += used[i];
    }
    
    for (int i = 1; i <= n; i ++) {
        if (used[i] == 1) {
            q.push((node){-s[i] - a[i], i});
        }
        else if (used[i]) {
            q.push((node){-a[i], i});
        }
    }
    
    for (int i = p - 1; i; i --) {
        ans[i] = ans[i + 1];
        while (sm > i * m && !q.empty()) {
            node u = q.top();
            q.pop();
            u.v *= -1;
            if (used[u.i] > 1) {
                int rest = min(sm - i * m, used[u.i] - 1);
                used[u.i] -= rest;
                sm -= rest;
                ans[i] -= 1ll * rest * u.v;
                if (used[u.i] == 1) {
                    q.push((node){-a[u.i] - s[u.i], u.i});
                }
                else {
                    q.push((node){-a[u.i], u.i});
                }
            }
            else {
                sm --;
                used[u.i] --;
                ans[i] -= u.v;
            }
        }
    }
    
    while (qry --) {
        printf("%lld\n", ans[read()]);
    }
    
    return 0;
}