use std::cmp::Ordering;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct Objective {
pub billed: usize,
pub per_rule: usize,
pub spread: usize,
}
impl Objective {
fn key(self) -> (usize, usize, usize) {
(self.billed, self.per_rule, self.spread)
}
}
impl Ord for Objective {
fn cmp(&self, other: &Self) -> Ordering {
self.key().cmp(&other.key())
}
}
impl PartialOrd for Objective {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
#[derive(Debug, Clone)]
pub struct Model {
items: Vec<Vec<usize>>,
file_weight: Vec<usize>,
}
const EXACT_MAX_FILES: usize = 128;
const EXACT_NODE_BUDGET: u64 = 300_000;
impl Model {
pub fn new(mut items: Vec<Vec<usize>>, file_weight: Vec<usize>) -> Self {
for it in &mut items {
it.sort_unstable();
it.dedup();
}
Model { items, file_weight }
}
fn union_weight(&self, batch: &[usize]) -> usize {
let mut seen = vec![false; self.file_weight.len()];
let mut total = 0;
for &i in batch {
for &f in &self.items[i] {
if !seen[f] {
seen[f] = true;
total += self.file_weight[f];
}
}
}
total
}
pub fn objective(&self, batches: &[Vec<usize>]) -> Objective {
let mut billed = 0;
let mut per_rule = 0;
let mut spread = 0;
for b in batches {
let w = self.union_weight(b);
billed += w;
per_rule += b.len() * w;
spread += b.len() * b.len();
}
Objective {
billed,
per_rule,
spread,
}
}
pub fn assign(&self, batch_count: usize, batch_size: usize) -> Vec<Vec<usize>> {
let n = self.items.len();
if n == 0 {
return Vec::new();
}
if batch_count <= 1 {
return vec![(0..n).collect()];
}
if let Some(batches) = self.assign_exact(batch_count, batch_size) {
return batches;
}
self.assign_heuristic(batch_count, batch_size)
}
fn assign_exact(&self, batch_count: usize, batch_size: usize) -> Option<Vec<Vec<usize>>> {
let n = self.items.len();
if self.file_weight.len() > EXACT_MAX_FILES {
return None;
}
let masks: Vec<u128> = self
.items
.iter()
.map(|it| it.iter().fold(0u128, |m, &f| m | (1u128 << f)))
.collect();
let mut best: Option<(Objective, Vec<Vec<usize>>)> = None;
let mut batches: Vec<Vec<usize>> = Vec::with_capacity(batch_count);
let mut unions: Vec<u128> = Vec::with_capacity(batch_count);
let mut billed = 0usize;
let mut nodes = 0u64;
self.search(
0,
n,
batch_count,
batch_size,
&masks,
&mut batches,
&mut unions,
&mut billed,
&mut best,
&mut nodes,
)?;
best.map(|(_, b)| canonicalize(b))
}
#[allow(clippy::too_many_arguments)]
fn search(
&self,
i: usize,
n: usize,
batch_count: usize,
batch_size: usize,
masks: &[u128],
batches: &mut Vec<Vec<usize>>,
unions: &mut Vec<u128>,
billed: &mut usize,
best: &mut Option<(Objective, Vec<Vec<usize>>)>,
nodes: &mut u64,
) -> Option<()> {
*nodes += 1;
if *nodes > EXACT_NODE_BUDGET {
return None;
}
if let Some((obj, _)) = best {
if *billed > obj.billed {
return Some(());
}
}
if i == n {
if batches.len() == batch_count {
let obj = self.objective(batches);
if best.as_ref().is_none_or(|(b, _)| obj < *b) {
*best = Some((obj, batches.clone()));
}
}
return Some(());
}
let remaining = n - i;
let needed = batch_count - batches.len();
if remaining < needed {
return Some(());
}
for b in 0..batches.len() {
if batches[b].len() >= batch_size {
continue;
}
let added = weight_of(masks[i] & !unions[b], &self.file_weight);
batches[b].push(i);
unions[b] |= masks[i];
*billed += added;
self.search(
i + 1,
n,
batch_count,
batch_size,
masks,
batches,
unions,
billed,
best,
nodes,
)?;
*billed -= added;
batches[b].pop();
unions[b] = batches[b].iter().fold(0u128, |m, &x| m | masks[x]);
}
if batches.len() < batch_count {
let added = weight_of(masks[i], &self.file_weight);
batches.push(vec![i]);
unions.push(masks[i]);
*billed += added;
self.search(
i + 1,
n,
batch_count,
batch_size,
masks,
batches,
unions,
billed,
best,
nodes,
)?;
*billed -= added;
batches.pop();
unions.pop();
}
Some(())
}
fn assign_heuristic(&self, batch_count: usize, batch_size: usize) -> Vec<Vec<usize>> {
let n = self.items.len();
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a, &b| {
self.union_weight(&[b])
.cmp(&self.union_weight(&[a]))
.then(a.cmp(&b))
});
let mut batches: Vec<Vec<usize>> = vec![Vec::new(); batch_count];
for &i in &order {
let mut best: Option<(usize, usize, usize, usize)> = None;
for (bi, batch) in batches.iter().enumerate() {
if batch.len() >= batch_size {
continue;
}
let marginal = {
let mut with = batch.clone();
with.push(i);
self.union_weight(&with) - self.union_weight(batch)
};
let key = (marginal, self.union_weight(batch), batch.len(), bi);
if best.is_none_or(|k| key < k) {
best = Some(key);
}
}
batches[best.expect("a non-full batch always exists").3].push(i);
}
self.repair_empty(&mut batches);
self.local_search(&mut batches, batch_size);
canonicalize(batches)
}
fn repair_empty(&self, batches: &mut [Vec<usize>]) {
while let Some(empty) = batches.iter().position(|b| b.is_empty()) {
let largest = batches
.iter()
.enumerate()
.filter(|(_, b)| b.len() > 1)
.max_by_key(|(_, b)| b.len())
.map(|(i, _)| i);
let Some(src) = largest else { break };
let moved = batches[src].pop().expect("largest batch is non-empty");
batches[empty].push(moved);
}
}
fn local_search(&self, batches: &mut [Vec<usize>], batch_size: usize) {
for _ in 0..1_000_000 {
let base = self.objective(batches);
if self.try_move(batches, batch_size, base) || self.try_swap(batches, base) {
continue;
}
break;
}
}
fn try_move(&self, batches: &mut [Vec<usize>], batch_size: usize, base: Objective) -> bool {
for from in 0..batches.len() {
if batches[from].len() <= 1 {
continue; }
for pos in 0..batches[from].len() {
for to in 0..batches.len() {
if to == from || batches[to].len() >= batch_size {
continue;
}
let item = batches[from].remove(pos);
batches[to].push(item);
if self.objective(batches) < base {
return true;
}
batches[to].pop();
batches[from].insert(pos, item);
}
}
}
false
}
fn try_swap(&self, batches: &mut [Vec<usize>], base: Objective) -> bool {
for a in 0..batches.len() {
for b in (a + 1)..batches.len() {
for pa in 0..batches[a].len() {
for pb in 0..batches[b].len() {
let ia = batches[a][pa];
let ib = batches[b][pb];
batches[a][pa] = ib;
batches[b][pb] = ia;
if self.objective(batches) < base {
return true;
}
batches[a][pa] = ia;
batches[b][pb] = ib;
}
}
}
}
false
}
}
fn weight_of(mask: u128, weight: &[usize]) -> usize {
let mut m = mask;
let mut total = 0;
while m != 0 {
let bit = m.trailing_zeros() as usize;
total += weight[bit];
m &= m - 1;
}
total
}
fn canonicalize(mut batches: Vec<Vec<usize>>) -> Vec<Vec<usize>> {
for b in &mut batches {
b.sort_unstable();
}
batches.retain(|b| !b.is_empty());
batches.sort_by(|a, b| a.first().cmp(&b.first()));
batches
}
#[cfg(test)]
mod tests {
use super::*;
fn brute_optimum(model: &Model, batch_count: usize, batch_size: usize) -> Objective {
let n = model.items.len();
let mut best: Option<Objective> = None;
let mut batches: Vec<Vec<usize>> = Vec::new();
fn rec(
i: usize,
n: usize,
bc: usize,
bs: usize,
model: &Model,
batches: &mut Vec<Vec<usize>>,
best: &mut Option<Objective>,
) {
if i == n {
if batches.len() == bc {
let obj = model.objective(batches);
if best.is_none_or(|b| obj < b) {
*best = Some(obj);
}
}
return;
}
for b in 0..batches.len() {
if batches[b].len() < bs {
batches[b].push(i);
rec(i + 1, n, bc, bs, model, batches, best);
batches[b].pop();
}
}
if batches.len() < bc {
batches.push(vec![i]);
rec(i + 1, n, bc, bs, model, batches, best);
batches.pop();
}
}
rec(
0,
n,
batch_count,
batch_size,
model,
&mut batches,
&mut best,
);
best.expect("at least one valid partition exists")
}
fn batch_count(n: usize, bs: usize) -> usize {
n.div_ceil(bs)
}
fn unit(items: &[&[usize]], num_files: usize) -> Model {
Model::new(
items.iter().map(|s| s.to_vec()).collect(),
vec![1; num_files],
)
}
fn weighted(items: &[&[usize]], weights: &[usize]) -> Model {
Model::new(items.iter().map(|s| s.to_vec()).collect(), weights.to_vec())
}
fn assert_optimal(model: &Model, bs: usize) {
let n = model.items.len();
let bc = batch_count(n, bs);
let got = model.assign(bc, bs);
assert_eq!(got.len(), bc, "batch count: {got:?}");
let mut seen: Vec<usize> = got.iter().flatten().copied().collect();
seen.sort_unstable();
assert_eq!(seen, (0..n).collect::<Vec<_>>(), "coverage: {got:?}");
assert!(
got.iter().all(|b| !b.is_empty() && b.len() <= bs),
"cap: {got:?}"
);
let achieved = model.objective(&got);
let optimum = brute_optimum(model, bc, bs);
assert_eq!(achieved, optimum, "assign not optimal: {got:?}");
}
#[test]
fn single_batch_is_a_trivial_no_op() {
let m = unit(&[&[0], &[1], &[2]], 3);
assert_eq!(m.assign(1, 20), vec![vec![0, 1, 2]]);
}
#[test]
fn one_rule_per_batch_when_cap_is_one() {
let m = unit(&[&[0], &[1], &[2]], 3);
assert_optimal(&m, 1);
}
#[test]
fn shared_files_are_grouped_over_the_interleaved_order() {
let m = unit(&[&[0], &[1], &[0], &[1]], 2);
assert_optimal(&m, 2);
let got = m.assign(2, 2);
assert_eq!(m.objective(&got).billed, 2);
}
#[test]
fn wide_rule_goes_in_the_smaller_batch_to_cut_per_rule_exposure() {
let m = unit(&[&[0, 1, 2], &[0], &[0], &[0], &[0]], 3);
assert_optimal(&m, 3);
let got = m.assign(batch_count(5, 3), 3); let wide = got.iter().find(|b| b.contains(&0)).unwrap();
assert_eq!(wide.len(), 2, "wide rule in the 2-rule batch: {got:?}");
}
#[test]
fn heavy_file_weight_dominates_grouping() {
let m = weighted(&[&[0], &[0], &[1], &[2]], &[100, 1, 1]);
assert_optimal(&m, 2);
let got = m.assign(2, 2);
assert_eq!(
m.objective(&got).billed,
102,
"heavy file billed once: {got:?}"
);
}
#[test]
fn all_disjoint_is_balanced_and_optimal() {
let m = unit(&[&[0], &[1], &[2], &[3]], 4);
assert_optimal(&m, 2);
}
#[test]
fn all_identical_scope_is_optimal() {
let m = unit(&[&[0, 1], &[0, 1], &[0, 1], &[0, 1]], 2);
assert_optimal(&m, 2);
}
#[test]
fn objective_is_lexicographic() {
let a = Objective {
billed: 4,
per_rule: 8,
spread: 2,
};
let b = Objective {
billed: 4,
per_rule: 6,
spread: 100,
};
let c = Objective {
billed: 3,
per_rule: 100,
spread: 100,
};
assert!(
b < a,
"lower per_rule wins on a billed tie, regardless of spread"
);
assert!(c < a, "lower billed wins regardless of the rest");
assert!(c < b);
let balanced = Objective {
billed: 2,
per_rule: 21,
spread: 221,
};
let packed = Objective {
billed: 2,
per_rule: 21,
spread: 401,
};
assert!(balanced < packed, "balanced sizes win the token tie");
}
#[test]
fn assignment_is_deterministic() {
let m = unit(&[&[0], &[1], &[0], &[1], &[2]], 3);
let bc = batch_count(5, 2);
assert_eq!(m.assign(bc, 2), m.assign(bc, 2));
}
#[test]
fn assign_is_optimal_across_a_broad_table() {
type Case = (Vec<Vec<usize>>, usize, Option<Vec<usize>>);
let cases: Vec<Case> = vec![
(vec![vec![0], vec![1], vec![2], vec![3]], 4, None),
(
vec![vec![0], vec![1], vec![0], vec![1], vec![0], vec![1]],
2,
None,
),
(
vec![vec![0, 1, 2], vec![0], vec![1], vec![2], vec![0]],
3,
None,
),
(
vec![vec![0, 1], vec![1, 2], vec![2, 3], vec![0, 3]],
4,
None,
),
(
vec![vec![0], vec![0, 1], vec![0, 1, 2], vec![0, 1, 2, 3]],
4,
None,
),
(
vec![vec![0], vec![0], vec![1], vec![2], vec![1]],
3,
Some(vec![50, 5, 1]),
),
(
vec![
vec![0, 1, 2, 3],
vec![0, 1, 2, 3],
vec![4],
vec![0, 1, 2, 3],
],
5,
None,
),
(
vec![
vec![0, 1],
vec![0],
vec![2],
vec![0, 3],
vec![2, 3],
vec![4],
],
5,
Some(vec![10, 1, 8, 2, 1]),
),
(
vec![
vec![0],
vec![1],
vec![2],
vec![0],
vec![1],
vec![2],
vec![0],
],
3,
None,
),
];
for (items, num_files, weights) in cases {
let n = items.len();
let model = match weights {
Some(w) => Model::new(items.clone(), w),
None => Model::new(items.clone(), vec![1; num_files]),
};
for bs in 1..=n {
let bc = batch_count(n, bs);
if bc < 1 {
continue;
}
let got = model.assign(bc, bs);
let achieved = model.objective(&got);
let optimum = brute_optimum(&model, bc, bs);
assert_eq!(
achieved, optimum,
"suboptimal for items={items:?} bs={bs} -> {got:?} ({achieved:?} vs {optimum:?})"
);
}
}
}
#[test]
fn heuristic_fallback_is_valid_and_not_worse_than_balanced() {
let n = 40;
let items: Vec<Vec<usize>> = (0..n).map(|i| vec![i]).collect();
let model = Model::new(items, vec![1; n]);
let bs = 3;
let bc = batch_count(n, bs);
let got = model.assign(bc, bs);
assert_eq!(got.len(), bc);
let mut seen: Vec<usize> = got.iter().flatten().copied().collect();
seen.sort_unstable();
assert_eq!(seen, (0..n).collect::<Vec<_>>());
assert!(got.iter().all(|b| !b.is_empty() && b.len() <= bs));
assert_eq!(model.objective(&got).billed, n);
}
}