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 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 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}