Skip to main content

competitive/graph/
steiner_tree.rs

1use super::{
2    AdditiveOperation, BitDpExt, Bounded, Graph, Monoid, PartialIgnoredOrd, ShortestPathSemiRing,
3    UnionFind, VertexMap, Zero,
4    shortest_path::{NoParent, OptionSp, RecordParent, StandardSp},
5};
6use std::{
7    cmp::Reverse, collections::BinaryHeap, iter::repeat_with, marker::PhantomData, ops::Add,
8};
9
10pub enum SteinerTreeParent<V, L> {
11    None,
12    Split(usize),
13    Edge(V, L),
14}
15
16pub trait SteinerTreeParentPolicy<G: Graph> {
17    type State;
18    type Label;
19    fn init(graph: &G) -> Self::State;
20    fn label(label: &G::Label) -> Self::Label;
21    fn save_split(graph: &G, state: &mut Self::State, vertex: G::Vertex, subset: usize);
22    fn save_parent(
23        graph: &G,
24        state: &mut Self::State,
25        from: G::Vertex,
26        to: G::Vertex,
27        label: Self::Label,
28    );
29}
30
31impl<G: Graph> SteinerTreeParentPolicy<G> for NoParent {
32    type State = ();
33    type Label = ();
34    fn init(_graph: &G) {}
35    fn label(_label: &G::Label) {}
36    fn save_split(_graph: &G, _state: &mut (), _vertex: G::Vertex, _subset: usize) {}
37    fn save_parent(_graph: &G, _state: &mut (), _from: G::Vertex, _to: G::Vertex, _label: ()) {}
38}
39
40impl<G> SteinerTreeParentPolicy<G> for RecordParent
41where
42    G: Graph<Label: Clone>
43        + VertexMap<SteinerTreeParent<<G as Graph>::Vertex, <G as Graph>::Label>>,
44{
45    type State = <G as VertexMap<SteinerTreeParent<G::Vertex, G::Label>>>::Vmap;
46    type Label = G::Label;
47    fn init(graph: &G) -> Self::State {
48        graph.construct_vmap(|| SteinerTreeParent::None)
49    }
50    fn label(label: &G::Label) -> Self::Label {
51        label.clone()
52    }
53    fn save_split(graph: &G, state: &mut Self::State, vertex: G::Vertex, subset: usize) {
54        *graph.vmap_get_mut(state, vertex) = SteinerTreeParent::Split(subset);
55    }
56    fn save_parent(
57        graph: &G,
58        state: &mut Self::State,
59        from: G::Vertex,
60        to: G::Vertex,
61        label: Self::Label,
62    ) {
63        *graph.vmap_get_mut(state, to) = SteinerTreeParent::Edge(from, label);
64    }
65}
66
67pub trait SteinerTreeExt: Graph {
68    fn steiner_tree(&self) -> SteinerTreeBuilder<'_, Self>
69    where
70        Self: Sized,
71    {
72        SteinerTreeBuilder {
73            graph: self,
74            _marker: PhantomData,
75        }
76    }
77}
78impl<G> SteinerTreeExt for G where G: Graph + ?Sized {}
79
80pub struct SteinerTreeBuilder<'g, G, S = (), P = NoParent>
81where
82    G: Graph,
83    P: SteinerTreeParentPolicy<G>,
84{
85    graph: &'g G,
86    _marker: PhantomData<fn() -> (S, P)>,
87}
88
89impl<'g, G, S, P> SteinerTreeBuilder<'g, G, S, P>
90where
91    G: Graph,
92    P: SteinerTreeParentPolicy<G>,
93{
94    pub fn with_sp<T>(self) -> SteinerTreeBuilder<'g, G, T, P>
95    where
96        T: ShortestPathSemiRing,
97    {
98        SteinerTreeBuilder {
99            graph: self.graph,
100            _marker: PhantomData,
101        }
102    }
103
104    pub fn with_standard_sp<M>(self) -> SteinerTreeBuilder<'g, G, StandardSp<M>, P>
105    where
106        M: Monoid<T: Bounded + Ord>,
107    {
108        self.with_sp()
109    }
110
111    pub fn with_standard_sp_additive<T>(
112        self,
113    ) -> SteinerTreeBuilder<'g, G, StandardSp<AdditiveOperation<T>>, P>
114    where
115        T: Clone + Zero + Add<Output = T> + Bounded + Ord,
116    {
117        self.with_sp()
118    }
119
120    pub fn with_option_sp<M>(self) -> SteinerTreeBuilder<'g, G, OptionSp<M>, P>
121    where
122        M: Monoid<T: Ord>,
123    {
124        self.with_sp()
125    }
126
127    pub fn with_option_sp_additive<T>(
128        self,
129    ) -> SteinerTreeBuilder<'g, G, OptionSp<AdditiveOperation<T>>, P>
130    where
131        T: Clone + Zero + Add<Output = T> + Ord,
132    {
133        self.with_sp()
134    }
135}
136
137impl<'g, G, S> SteinerTreeBuilder<'g, G, S>
138where
139    G: Graph,
140{
141    pub fn with_parent(self) -> SteinerTreeBuilder<'g, G, S, RecordParent>
142    where
143        RecordParent: SteinerTreeParentPolicy<G>,
144    {
145        SteinerTreeBuilder {
146            graph: self.graph,
147            _marker: PhantomData,
148        }
149    }
150}
151
152impl<'g, G, S, P> SteinerTreeBuilder<'g, G, S, P>
153where
154    G: Graph + VertexMap<S::T>,
155    S: ShortestPathSemiRing,
156    P: SteinerTreeParentPolicy<G>,
157{
158    /// Requires commutative multiplication and nonnegative edge weights.
159    pub fn solve<M, I>(&self, terminals: I, weight: M) -> SteinerTreeOutput<'g, S, G, P>
160    where
161        M: Fn(G::Label) -> S::T,
162        I: ExactSizeIterator<Item = G::Vertex>,
163    {
164        let graph = self.graph;
165        let tsize = terminals.len();
166        let states = if tsize == 0 { 0 } else { 1 << tsize };
167        let mut dp: Vec<_> = repeat_with(|| graph.construct_vmap(S::inf))
168            .take(states)
169            .collect();
170        let mut parent: Vec<_> = repeat_with(|| P::init(graph)).take(states).collect();
171        for (i, t) in terminals.enumerate() {
172            *graph.vmap_get_mut(&mut dp[1 << i], t) = S::source();
173        }
174        let inf = S::inf();
175        for bit in 1..states {
176            let (prev, current) = dp.split_at_mut(bit);
177            let dp = &mut current[0];
178            for sub in bit.subsets().skip(1).take_while(|&sub| sub > bit ^ sub) {
179                let left = &prev[sub];
180                let right = &prev[bit ^ sub];
181                for u in graph.vertices() {
182                    let left = graph.vmap_get(left, u);
183                    let right = graph.vmap_get(right, u);
184                    if left != &inf && right != &inf {
185                        let cost = S::mul(left, right);
186                        if S::add_assign(graph.vmap_get_mut(dp, u), &cost) {
187                            P::save_split(graph, &mut parent[bit], u, sub);
188                        }
189                    }
190                }
191            }
192            let mut heap: BinaryHeap<_> = graph
193                .vertices()
194                .filter_map(|u| {
195                    let d = graph.vmap_get(dp, u);
196                    (d != &inf).then(|| PartialIgnoredOrd(Reverse(d.clone()), u))
197                })
198                .collect();
199            while let Some(PartialIgnoredOrd(Reverse(d), u)) = heap.pop() {
200                if graph.vmap_get(dp, u) != &d {
201                    continue;
202                }
203                for neighbor in graph.neighbors(u) {
204                    let v = neighbor.to;
205                    let label = P::label(&neighbor.label);
206                    let nd = S::mul(&d, &weight(neighbor.label));
207                    if S::add_assign(graph.vmap_get_mut(dp, v), &nd) {
208                        P::save_parent(graph, &mut parent[bit], u, v, label);
209                        heap.push(PartialIgnoredOrd(Reverse(nd), v));
210                    }
211                }
212            }
213        }
214        SteinerTreeOutput { graph, dp, parent }
215    }
216}
217
218pub struct SteinerTreeOutput<'g, S, G, P = NoParent>
219where
220    G: Graph + VertexMap<S::T>,
221    S: ShortestPathSemiRing,
222    P: SteinerTreeParentPolicy<G>,
223{
224    graph: &'g G,
225    dp: Vec<<G as VertexMap<S::T>>::Vmap>,
226    parent: Vec<P::State>,
227}
228
229impl<S, G, P> SteinerTreeOutput<'_, S, G, P>
230where
231    G: Graph + VertexMap<S::T>,
232    S: ShortestPathSemiRing,
233    P: SteinerTreeParentPolicy<G>,
234{
235    pub fn minimum_from_source(&self, source: G::Vertex) -> S::T {
236        match self.dp.last() {
237            Some(dp) => self.graph.vmap_get(dp, source).clone(),
238            None => S::source(),
239        }
240    }
241}
242
243impl<S, G> SteinerTreeOutput<'_, S, G, RecordParent>
244where
245    G: Graph<Label: Clone>
246        + VertexMap<S::T>
247        + VertexMap<usize>
248        + VertexMap<SteinerTreeParent<<G as Graph>::Vertex, <G as Graph>::Label>>,
249    S: ShortestPathSemiRing,
250{
251    /// For undirected graphs with nonnegative additive weights.
252    /// Returns `None` when the terminals cannot be connected to `source`.
253    pub fn edges_from_source(&self, source: G::Vertex) -> Option<Vec<G::Label>> {
254        if self.dp.is_empty() {
255            return Some(vec![]);
256        }
257        if self.minimum_from_source(source) == S::inf() {
258            return None;
259        }
260        let graph = self.graph;
261        let mut index = graph.construct_vmap(|| 0usize);
262        for (i, u) in graph.vertices().enumerate() {
263            *graph.vmap_get_mut(&mut index, u) = i;
264        }
265        let mut uf = UnionFind::new(graph.vsize());
266        let mut edges = vec![];
267        let mut stack = vec![(self.dp.len() - 1, source)];
268        while let Some((bit, u)) = stack.pop() {
269            match graph.vmap_get(&self.parent[bit], u) {
270                SteinerTreeParent::None => {}
271                &SteinerTreeParent::Split(sub) => {
272                    stack.push((sub, u));
273                    stack.push((bit ^ sub, u));
274                }
275                SteinerTreeParent::Edge(v, label) => {
276                    if uf.unite(*graph.vmap_get(&index, u), *graph.vmap_get(&index, *v)) {
277                        edges.push(label.clone());
278                    }
279                    stack.push((bit, *v));
280                }
281            }
282        }
283        Some(edges)
284    }
285}
286
287#[cfg(test)]
288mod tests {
289    use super::*;
290    use crate::{algebra::AdditiveOperation, graph::UndirectedSparseGraph, tools::Xorshift};
291
292    #[test]
293    fn test_steiner_tree() {
294        let check = |n: usize, edges: &[(usize, usize, u64)], terminals: &[usize]| {
295            let graph = UndirectedSparseGraph::from_edges(
296                n,
297                edges.iter().map(|&(u, v, _)| (u, v)).collect(),
298            );
299            let plain = graph
300                .steiner_tree()
301                .with_standard_sp::<AdditiveOperation<_>>()
302                .solve(terminals.iter().copied(), |eid| edges[eid].2);
303            let recorded = graph
304                .steiner_tree()
305                .with_standard_sp_additive()
306                .with_parent()
307                .solve(terminals.iter().copied(), |eid| edges[eid].2);
308            let optional = graph
309                .steiner_tree()
310                .with_parent()
311                .with_option_sp::<AdditiveOperation<_>>()
312                .solve(terminals.iter().copied(), |eid| Some(edges[eid].2));
313            for source in 0..n {
314                let mut expected = None;
315                for mask in 0..1usize << edges.len() {
316                    let mut reachable = vec![false; n];
317                    reachable[source] = true;
318                    let mut stack = vec![source];
319                    while let Some(u) = stack.pop() {
320                        for (eid, &(a, b, _)) in edges.iter().enumerate() {
321                            if mask >> eid & 1 == 0 {
322                                continue;
323                            }
324                            let v = if a == u {
325                                b
326                            } else if b == u {
327                                a
328                            } else {
329                                continue;
330                            };
331                            if !reachable[v] {
332                                reachable[v] = true;
333                                stack.push(v);
334                            }
335                        }
336                    }
337                    if terminals.iter().all(|&t| reachable[t]) {
338                        let cost: u64 = edges
339                            .iter()
340                            .enumerate()
341                            .filter(|&(eid, _)| mask >> eid & 1 != 0)
342                            .map(|(_, e)| e.2)
343                            .sum();
344                        expected = Some(expected.map_or(cost, |best: u64| best.min(cost)));
345                    }
346                }
347                assert_eq!(
348                    plain.minimum_from_source(source),
349                    expected.unwrap_or(u64::MAX)
350                );
351                assert_eq!(
352                    recorded.minimum_from_source(source),
353                    expected.unwrap_or(u64::MAX)
354                );
355                assert_eq!(optional.minimum_from_source(source), expected);
356                for restored in [
357                    recorded.edges_from_source(source),
358                    optional.edges_from_source(source),
359                ] {
360                    assert_eq!(restored.is_some(), expected.is_some());
361                    if let Some(restored) = restored {
362                        let mut reachable = vec![false; n];
363                        reachable[source] = true;
364                        let mut stack = vec![source];
365                        while let Some(u) = stack.pop() {
366                            for &eid in &restored {
367                                let (a, b, _) = edges[eid];
368                                let v = if a == u {
369                                    b
370                                } else if b == u {
371                                    a
372                                } else {
373                                    continue;
374                                };
375                                if !reachable[v] {
376                                    reachable[v] = true;
377                                    stack.push(v);
378                                }
379                            }
380                        }
381                        assert!(terminals.iter().all(|&t| reachable[t]));
382                        assert_eq!(
383                            restored.iter().map(|&eid| edges[eid].2).sum::<u64>(),
384                            expected.unwrap()
385                        );
386                        assert_eq!(restored.len() + 1, reachable.iter().filter(|&&v| v).count());
387                    }
388                }
389            }
390        };
391        for n in 1..=4 {
392            let pairs: Vec<_> = (0..n)
393                .flat_map(|u| (u + 1..n).map(move |v| (u, v)))
394                .collect();
395            for code in 0..3usize.pow(pairs.len() as u32) {
396                let edges: Vec<_> = pairs
397                    .iter()
398                    .enumerate()
399                    .filter_map(|(i, &(u, v))| {
400                        let value = code / 3usize.pow(i as u32) % 3;
401                        (value != 0).then_some((u, v, value.saturating_sub(1) as u64))
402                    })
403                    .collect();
404                for subset in 0..1usize << n {
405                    let terminals: Vec<_> = (0..n).filter(|&u| subset >> u & 1 != 0).collect();
406                    check(n, &edges, &terminals);
407                }
408            }
409        }
410        let mut rng = Xorshift::default();
411        for n in 1..=7 {
412            for case in 0..16 {
413                let mut edges = match case % 4 {
414                    0 => vec![],
415                    1 => (1..n).map(|u| (u - 1, u, rng.rand(3))).collect(),
416                    2 => (1..n).map(|u| (0, u, rng.rand(3))).collect(),
417                    _ => (0..6)
418                        .map(|_| {
419                            (
420                                rng.rand(n as u64) as usize,
421                                rng.rand(n as u64) as usize,
422                                rng.rand(3),
423                            )
424                        })
425                        .collect(),
426                };
427                if case % 4 == 3 {
428                    edges.push(edges[0]);
429                    edges.push((0, 0, 0));
430                }
431                if case >= 8 {
432                    for e in &mut edges {
433                        e.2 *= 1_000_000_000;
434                    }
435                }
436                let subset = if case % 4 == 0 {
437                    0
438                } else if case % 4 == 1 {
439                    (1 << n) - 1
440                } else {
441                    rng.rand(1 << n)
442                };
443                let mut terminals: Vec<_> = (0..n).filter(|&u| subset >> u & 1 != 0).collect();
444                if let Some(&t) = terminals.first() {
445                    terminals.push(t);
446                }
447                check(n, &edges, &terminals);
448            }
449        }
450    }
451}