Skip to main content

competitive/algorithm/
combinations.rs

1use std::{collections::BTreeSet, mem::swap};
2
3pub trait SliceCombinationsExt<T> {
4    fn for_each_product<F>(&self, r: usize, f: F)
5    where
6        F: FnMut(&[T]);
7    fn for_each_permutations<F>(&self, r: usize, f: F)
8    where
9        F: FnMut(&[T]);
10    fn for_each_combinations<F>(&self, r: usize, f: F)
11    where
12        F: FnMut(&[T]);
13    fn for_each_combinations_with_replacement<F>(&self, r: usize, f: F)
14    where
15        F: FnMut(&[T]);
16    fn next_permutation(&mut self) -> bool
17    where
18        T: Ord;
19    fn prev_permutation(&mut self) -> bool
20    where
21        T: Ord;
22    fn next_combination(&mut self, r: usize) -> bool
23    where
24        T: Ord;
25    fn prev_combination(&mut self, r: usize) -> bool
26    where
27        T: Ord;
28
29    fn apply_permutation(&mut self, permutation: &[usize]);
30}
31
32impl<T> SliceCombinationsExt<T> for [T]
33where
34    T: Clone,
35{
36    /// choose `r` elements from `n` independently
37    ///
38    /// # Example
39    ///
40    /// ```
41    /// # use competitive::algorithm::SliceCombinationsExt;
42    /// let n = vec![1, 2, 3, 4];
43    /// let mut p = Vec::new();
44    /// let mut q = Vec::new();
45    /// n.for_each_product(2, |v| p.push(v.to_vec()));
46    /// for x in n.iter().cloned() {
47    ///     for y in n.iter().cloned() {
48    ///         q.push(vec![x, y]);
49    ///     }
50    /// }
51    /// assert_eq!(p, q);
52    /// ```
53    fn for_each_product<F>(&self, r: usize, mut f: F)
54    where
55        F: FnMut(&[T]),
56    {
57        fn product_inner<T, F>(n: &[T], mut r: usize, buf: &mut Vec<T>, f: &mut F)
58        where
59            T: Clone,
60            F: FnMut(&[T]),
61        {
62            if r == 0 {
63                f(buf.as_slice());
64            } else {
65                r -= 1;
66                for a in n.iter().cloned() {
67                    buf.push(a);
68                    product_inner(n, r, buf, f);
69                    buf.pop();
70                }
71            }
72        }
73
74        let mut v = Vec::with_capacity(r);
75        product_inner(self, r, &mut v, &mut f);
76    }
77
78    /// choose `r` elements from `n` independently
79    ///
80    /// # Example
81    ///
82    /// ```
83    /// # use competitive::algorithm::SliceCombinationsExt;
84    /// let n = vec![1, 2, 3, 4];
85    /// let mut p = Vec::new();
86    /// let mut q = Vec::new();
87    /// n.for_each_product(2, |v| p.push(v.to_vec()));
88    /// for x in n.iter().cloned() {
89    ///     for y in n.iter().cloned() {
90    ///         q.push(vec![x, y]);
91    ///     }
92    /// }
93    /// assert_eq!(p, q);
94    /// ```
95    fn for_each_permutations<F>(&self, r: usize, mut f: F)
96    where
97        F: FnMut(&[T]),
98    {
99        fn permutations_inner<T, F>(
100            n: &[T],
101            mut r: usize,
102            rem: &mut BTreeSet<usize>,
103            buf: &mut Vec<T>,
104            f: &mut F,
105        ) where
106            T: Clone,
107            F: FnMut(&[T]),
108        {
109            if r == 0 {
110                f(buf.as_slice());
111            } else {
112                r -= 1;
113                for i in rem.iter().cloned().collect::<Vec<_>>() {
114                    buf.push(n[i].clone());
115                    rem.remove(&i);
116                    permutations_inner(n, r, rem, buf, f);
117                    rem.insert(i);
118                    buf.pop();
119                }
120            }
121        }
122
123        if r <= self.len() {
124            let mut v = Vec::with_capacity(r);
125            let mut rem: BTreeSet<usize> = (0..self.len()).collect();
126            permutations_inner(self, r, &mut rem, &mut v, &mut f);
127        }
128    }
129
130    /// choose distinct `r` elements from `n` in any order
131    ///
132    /// # Example
133    ///
134    /// ```
135    /// # use competitive::algorithm::SliceCombinationsExt;
136    /// let n = vec![1, 2, 3, 4];
137    /// let mut p = Vec::new();
138    /// let mut q = Vec::new();
139    /// n.for_each_permutations(2, |v| p.push(v.to_vec()));
140    /// for (i, x) in n.iter().cloned().enumerate() {
141    ///     for (j, y) in n.iter().cloned().enumerate() {
142    ///         if i != j {
143    ///             q.push(vec![x, y]);
144    ///         }
145    ///     }
146    /// }
147    /// assert_eq!(p, q);
148    /// ```
149    fn for_each_combinations<F>(&self, r: usize, mut f: F)
150    where
151        F: FnMut(&[T]),
152    {
153        fn combinations_inner<T, F>(
154            n: &[T],
155            mut r: usize,
156            start: usize,
157            buf: &mut Vec<T>,
158            f: &mut F,
159        ) where
160            T: Clone,
161            F: FnMut(&[T]),
162        {
163            if r == 0 {
164                f(buf.as_slice());
165            } else {
166                r -= 1;
167                for i in start..n.len() - r {
168                    buf.push(n[i].clone());
169                    combinations_inner(n, r, i + 1, buf, f);
170                    buf.pop();
171                }
172            }
173        }
174
175        if r <= self.len() {
176            let mut v = Vec::with_capacity(r);
177            combinations_inner(self, r, 0, &mut v, &mut f);
178        }
179    }
180
181    /// choose `r` elements from `n` in sorted order
182    ///
183    /// # Example
184    ///
185    /// ```
186    /// # use competitive::algorithm::SliceCombinationsExt;
187    /// let n = vec![1, 2, 3, 4];
188    /// let mut p = Vec::new();
189    /// let mut q = Vec::new();
190    /// n.for_each_combinations_with_replacement(2, |v| p.push(v.to_vec()));
191    /// for (i, x) in n.iter().cloned().enumerate() {
192    ///     for y in n[i..].iter().cloned() {
193    ///         q.push(vec![x, y]);
194    ///     }
195    /// }
196    /// assert_eq!(p, q);
197    /// ```
198    fn for_each_combinations_with_replacement<F>(&self, r: usize, mut f: F)
199    where
200        F: FnMut(&[T]),
201    {
202        fn combinations_with_replacement_inner<T, F>(
203            n: &[T],
204            mut r: usize,
205            start: usize,
206            buf: &mut Vec<T>,
207            f: &mut F,
208        ) where
209            T: Clone,
210            F: FnMut(&[T]),
211        {
212            if r == 0 {
213                f(buf.as_slice());
214            } else {
215                r -= 1;
216                for i in start..n.len() {
217                    buf.push(n[i].clone());
218                    combinations_with_replacement_inner(n, r, i, buf, f);
219                    buf.pop();
220                }
221            }
222        }
223
224        let mut v = Vec::with_capacity(r);
225        combinations_with_replacement_inner(self, r, 0, &mut v, &mut f);
226    }
227
228    /// Permute the elements into next permutation in lexicographical order.
229    /// Return whether such a next permutation exists.
230    fn next_permutation(&mut self) -> bool
231    where
232        T: Ord,
233    {
234        if self.len() < 2 {
235            return false;
236        }
237        let mut target = self.len() - 2;
238        while target > 0 && self[target] > self[target + 1] {
239            target -= 1;
240        }
241        if target == 0 && self[target] > self[target + 1] {
242            return false;
243        }
244        let mut next = self.len() - 1;
245        while next > target && self[next] < self[target] {
246            next -= 1;
247        }
248        self.swap(next, target);
249        self[target + 1..].reverse();
250        true
251    }
252
253    /// Permute the elements into previous permutation in lexicographical order.
254    /// Return whether such a previous permutation exists.
255    fn prev_permutation(&mut self) -> bool
256    where
257        T: Ord,
258    {
259        if self.len() < 2 {
260            return false;
261        }
262        let mut target = self.len() - 2;
263        while target > 0 && self[target] < self[target + 1] {
264            target -= 1;
265        }
266        if target == 0 && self[target] < self[target + 1] {
267            return false;
268        }
269        self[target + 1..].reverse();
270        let mut next = self.len() - 1;
271        while next > target && self[next - 1] < self[target] {
272            next -= 1;
273        }
274        self.swap(target, next);
275        true
276    }
277
278    /// Permute the elements into next combination choosing r elements in lexicographical order.
279    /// Return whether such a next combination exists.
280    fn next_combination(&mut self, r: usize) -> bool
281    where
282        T: Ord,
283    {
284        assert!(r <= self.len());
285        let (a, b) = self.split_at_mut(r);
286        next_combination_inner(a, b)
287    }
288
289    /// Permute the elements into previous combination choosing r elements in lexicographical order.
290    /// Return whether such a previous combination exists.
291    fn prev_combination(&mut self, r: usize) -> bool
292    where
293        T: Ord,
294    {
295        assert!(r <= self.len());
296        let (a, b) = self.split_at_mut(r);
297        next_combination_inner(b, a)
298    }
299
300    /// Apply a permutation to the elements.
301    /// self[i] <- self[p[i]] for each i
302    fn apply_permutation(&mut self, p: &[usize]) {
303        assert_eq!(self.len(), p.len());
304        let mut visited = vec![false; self.len()];
305        for mut current in 0..self.len() {
306            if visited[current] {
307                continue;
308            }
309            loop {
310                visited[current] = true;
311                let next = p[current];
312                if visited[next] {
313                    break;
314                }
315                self.swap(current, next);
316                current = next;
317            }
318        }
319    }
320}
321
322fn rotate_distinct<'a, T>(mut a: &'a mut [T], mut b: &'a mut [T]) {
323    while !a.is_empty() && !b.is_empty() {
324        if a.len() >= b.len() {
325            let (l, r) = a.split_at_mut(b.len());
326            l.swap_with_slice(b);
327            a = r;
328        } else {
329            let (l, r) = b.split_at_mut(a.len());
330            l.swap_with_slice(a);
331            a = l;
332            b = r;
333        }
334    }
335}
336
337fn next_combination_inner<T>(a: &mut [T], b: &mut [T]) -> bool
338where
339    T: Ord,
340{
341    if a.is_empty() || b.is_empty() {
342        return false;
343    }
344    let mut target = a.len() - 1;
345    let last_elem = b.last().unwrap();
346    while target > 0 && &a[target] >= last_elem {
347        target -= 1;
348    }
349    if target == 0 && &a[target] >= last_elem {
350        rotate_distinct(a, b);
351        return false;
352    }
353    let mut next = 0;
354    while a[target] >= b[next] {
355        next += 1;
356    }
357    swap(&mut a[target], &mut b[next]);
358    rotate_distinct(&mut a[target + 1..], &mut b[next + 1..]);
359    true
360}
361
362#[cfg(test)]
363mod tests {
364    use super::*;
365    use crate::tools::Xorshift;
366
367    #[test]
368    fn test_enumeration() {
369        let mut rng = Xorshift::default();
370        for (n, r) in (1usize..=6).flat_map(|n| (0..=6).map(move |r| (n, r))) {
371            let values: Vec<_> = rng.random_iter(-10..=10).take(n).collect();
372            let mut product = Vec::new();
373            let mut permutations = Vec::new();
374            let mut combinations = Vec::new();
375            let mut replacement = Vec::new();
376            for mut code in 0..n.pow(r as u32) {
377                let mut indices = vec![0; r];
378                for i in indices.iter_mut().rev() {
379                    *i = code % n;
380                    code /= n;
381                }
382                let row: Vec<_> = indices.iter().map(|&i| values[i]).collect();
383                product.push(row.clone());
384                if (0..r).all(|i| !indices[..i].contains(&indices[i])) {
385                    permutations.push(row.clone());
386                }
387                if indices.windows(2).all(|w| w[0] < w[1]) {
388                    combinations.push(row.clone());
389                }
390                if indices.is_sorted() {
391                    replacement.push(row);
392                }
393            }
394            let mut actual = Vec::new();
395            values.for_each_product(r, |row| actual.push(row.to_vec()));
396            assert_eq!(actual, product);
397            actual.clear();
398            values.for_each_permutations(r, |row| actual.push(row.to_vec()));
399            assert_eq!(actual, permutations);
400            actual.clear();
401            values.for_each_combinations(r, |row| actual.push(row.to_vec()));
402            assert_eq!(actual, combinations);
403            actual.clear();
404            values.for_each_combinations_with_replacement(r, |row| actual.push(row.to_vec()));
405            assert_eq!(actual, replacement);
406        }
407    }
408
409    #[test]
410    fn test_next_prev() {
411        let mut rng = Xorshift::default();
412        for n in 1..=7usize {
413            let mut values: Vec<_> = (0..n)
414                .map(|i| i as i32 * 100 + rng.random(0..100))
415                .collect();
416            values.sort();
417            values.dedup();
418            let n = values.len();
419            let mut permutations = Vec::new();
420            values.for_each_permutations(n, |row| permutations.push(row.to_vec()));
421            let mut p = values.clone();
422            for (i, expected) in permutations.iter().enumerate() {
423                assert_eq!(&p, expected);
424                if i + 1 < permutations.len() {
425                    assert!(p.next_permutation());
426                    assert!(p.prev_permutation());
427                    assert_eq!(&p, expected);
428                }
429                assert_eq!(p.next_permutation(), i + 1 < permutations.len());
430            }
431            for r in 0..=n {
432                let mut combinations = Vec::new();
433                values.for_each_combinations(r, |row| combinations.push(row.to_vec()));
434                p = values.clone();
435                for (i, expected) in combinations.iter().enumerate() {
436                    assert_eq!(&p[..r], expected);
437                    if i + 1 < combinations.len() {
438                        assert!(p.next_combination(r));
439                        assert!(p.prev_combination(r));
440                        assert_eq!(&p[..r], expected);
441                    }
442                    assert_eq!(p.next_combination(r), i + 1 < combinations.len());
443                }
444            }
445        }
446    }
447
448    #[test]
449    fn test_apply_permutation() {
450        let mut rng = Xorshift::default();
451        for _ in 0..100 {
452            let n = rng.random(1..100);
453            let a: Vec<_> = rng.random_iter(0..1_000).take(n).collect();
454            let mut p: Vec<usize> = (0..n).collect();
455            rng.shuffle(&mut p);
456            let expected: Vec<_> = p.iter().map(|&i| a[i]).collect();
457            let mut result = a.to_vec();
458            result.apply_permutation(&p);
459            assert_eq!(expected, result);
460        }
461    }
462}