Official

D - Automat Editorial by sheyasutaka


達成可能な組合せ

商品の組合せ \(A_{i_1}, \dots, A_{i_n}, B_{j_1}, \dots, B_{j_m}\) は,以下を満たすとき,またその時に限り購入可能です.

  • 合計金額 \(\displaystyle \sum_{1 \leq k \leq n} A_{i_k} + \sum_{1 \leq k \leq m} B_{j_k}\) が \(X+KY\) 以下
  • 必要な \(K\) ドル紙幣の枚数 \(\displaystyle \sum_{1 \leq k \leq m} \lceil B_{j_k} / K \rceil\) が \(Y\) 以下

必要性は明らかです.ドリンク \(m\) 個を先にすべて買い,そのあとデザート \(n\) 個をすべて買うとすることで,上を満たす組合せは購入可能であることが分かります.

解法

\(B\) から \(m\) 個とるとします.このとき,\(B\) からは価格が小さいほうから \(m\) 個取るのが最適です.同様に,\(A\) も価格が小さいほうからとるとしてよいです.

\(B\) からとる個数を固定したとき,\(A\) にかけることができる金額が決まるので,二分探索によって \(A\) からとれる個数が求まります.

したがって,\(B\) からとる個数を全探索し,それぞれ二分探索によって \(A\) からとる個数を求めることで,\(O((N+M) \log (N+M))\) 時間で答えが求まります.

なお,\(B\) からとる個数が小さいほど,\(A\) にかけられる金額は広義単調増加するので,尺取り法によってソート以外の処理を \(O(N+M)\) 時間で行うこともできます.

余談

この問題設定においては,「価格が小さいほうから商品を見ていって,その商品を追加した組合せが購入可能なら追加し,不可能ならスルーする」として商品の組合せを構成したとき,それが最大値を達成することが証明できます.

実装例 (C++)

#include <iostream>
using std::cin;
using std::cout;
using std::cerr;
using std::endl;
#include <vector>
using std::vector;
using std::pair;
#include <map>
using std::map;
using std::max;
using std::min;


#ifdef DEBUG
const int debug = 1;
#else
const int debug = 0;
#endif

using ll = int64_t;
using P = pair<ll, ll>;

const ll FOD = 998244353;
#include <atcoder/modint>
using mint = atcoder::modint998244353;
#include <atcoder/convolution>
using atcoder::convolution;

ll n, m, k;
ll x, y;
vector<ll> a, b;

inline void output (const ll x) {
	cout << x << "\n";
}

vector<ll> accum (const vector<ll> &a) {
	ll sz = (ll)a.size();
	vector<ll> b(sz+1, 0);
	b[0] = 0;
	for (ll i = 0; i < sz; i++) {
		b[i+1] = b[i] + a[i];
	}
	return b;
}

void solve() {
	sort(a.begin(), a.end());
	sort(b.begin(), b.end());

	vector<ll> bk(m);
	for (ll i = 0; i < m; i++) {
		bk[i] = (b[i] - 1) / k + 1;
	}

	vector<ll> ac = accum(a);
	vector<ll> bc = accum(b);
	vector<ll> bkc = accum(bk);

	ll ans = 0;
	for (ll ri = 0; ri <= m; ri++) {
		ll s = (x+y*k) - bc[ri];
		ll sk = y - bkc[ri];
		if (s < 0 || sk < 0) break;

		ll li = ([&]() -> ll {
			ll ok = 0, ng = n+1;
			while (ok + 1 < ng) {
				ll med = (ok + ng) / 2;
				if (ac[med] <= s) {
					ok = med;
				} else {
					ng = med;
				}
			}
			return ok;
		})();

		ans = max(ans, li + ri);
	}

	output(ans);


	return;
}

int main (void) {
	std::cin.tie(nullptr);
	std::ios_base::sync_with_stdio(false);

	cin >> n >> m >> k;
	cin >> x >> y;
	a.resize(n);
	b.resize(m);
	for (ll i = 0; i < n; i++) {
		cin >> a[i];
	}
	for (ll i = 0; i < m; i++) {
		cin >> b[i];
	}

	
	solve();

	return 0;
}

posted:
last update: