Skip to main content

library_checker/tree/
point_set_tree_path_composite_sum.rs

1use competitive::prelude::*;
2use competitive::{
3    algebra::{Associative, Magma, Unital},
4    graph::TreeGraphScanner,
5    num::{One, Zero, mint_basic::MInt998244353 as M},
6    tree::MonoidCluster,
7};
8
9#[derive(Clone, Copy, Debug, PartialEq, Eq)]
10struct Point {
11    sum: M,
12    cnt: u32,
13}
14
15#[derive(Clone, Copy, Debug, PartialEq, Eq)]
16struct Path {
17    a: M,
18    b: M,
19    sum: M,
20    cnt: u32,
21}
22
23struct PointMonoid;
24impl Magma for PointMonoid {
25    type T = Point;
26    fn operate(x: &Self::T, y: &Self::T) -> Self::T {
27        Point {
28            sum: x.sum + y.sum,
29            cnt: x.cnt + y.cnt,
30        }
31    }
32}
33impl Unital for PointMonoid {
34    fn unit() -> Self::T {
35        Point {
36            sum: M::zero(),
37            cnt: 0,
38        }
39    }
40}
41impl Associative for PointMonoid {}
42
43struct PathMonoid;
44impl Magma for PathMonoid {
45    type T = Path;
46    fn operate(x: &Self::T, y: &Self::T) -> Self::T {
47        Path {
48            a: x.a * y.a,
49            b: x.b + x.a * y.b,
50            sum: x.sum + x.a * y.sum + x.b * M::new_unchecked(y.cnt),
51            cnt: x.cnt + y.cnt,
52        }
53    }
54}
55impl Unital for PathMonoid {
56    fn unit() -> Self::T {
57        Path {
58            a: M::one(),
59            b: M::zero(),
60            sum: M::zero(),
61            cnt: 0,
62        }
63    }
64}
65impl Associative for PathMonoid {}
66
67#[derive(Clone, Copy, Debug, PartialEq, Eq)]
68struct PathPair {
69    forward: Path,
70    reverse: Path,
71}
72
73struct PathPairMonoid;
74impl Magma for PathPairMonoid {
75    type T = PathPair;
76    fn operate(x: &Self::T, y: &Self::T) -> Self::T {
77        PathPair {
78            forward: PathMonoid::operate(&x.forward, &y.forward),
79            reverse: PathMonoid::operate(&y.reverse, &x.reverse),
80        }
81    }
82}
83impl Unital for PathPairMonoid {
84    fn unit() -> Self::T {
85        PathPair {
86            forward: PathMonoid::unit(),
87            reverse: PathMonoid::unit(),
88        }
89    }
90}
91impl Associative for PathPairMonoid {}
92
93struct Dp;
94
95impl MonoidCluster for Dp {
96    type Vertex = M;
97    type Edge = (M, M);
98    type PointMonoid = PointMonoid;
99    type PathMonoid = PathPairMonoid;
100
101    fn add_vertex(point: &Point, vertex: &M, parent_edge: Option<&(M, M)>) -> PathPair {
102        let cnt = point.cnt + 1;
103        let subtotal = point.sum + *vertex;
104        let (a, b) = parent_edge.copied().unwrap_or((M::one(), M::zero()));
105        PathPair {
106            forward: Path {
107                a,
108                b,
109                sum: a * subtotal + b * M::new_unchecked(cnt),
110                cnt,
111            },
112            reverse: Path {
113                a,
114                b,
115                sum: subtotal,
116                cnt,
117            },
118        }
119    }
120
121    fn add_edge(path: &PathPair) -> Point {
122        Point {
123            sum: path.forward.sum,
124            cnt: path.forward.cnt,
125        }
126    }
127}
128
129competitive::define_enum_scan! {
130    enum Query: usize {
131        0 => SetVertex { v: usize, x: M, r: usize }
132        1 => SetEdge { e: usize, a: M, b: M, r: usize }
133    }
134}
135
136#[verify::library_checker("point_set_tree_path_composite_sum")]
137pub fn point_set_tree_path_composite_sum(reader: impl Read, writer: impl Write) {
138    prepare_io!(reader, writer);
139    sc!(n,
140        q,
141        value: [M; n],
142        (graph, edges): @TreeGraphScanner::<usize, (M, M)>::new(n));
143
144    let top_tree = graph.static_top_tree(0);
145    let mut dp = top_tree.dp::<Dp>(value, edges);
146
147    for _ in 0..q {
148        sc!(query: Query);
149        match query {
150            Query::SetVertex { v, x, r } => {
151                dp.set_vertex(v, x);
152                pp!(dp.fold_path(r).reverse.sum);
153            }
154            Query::SetEdge { e, a, b, r } => {
155                dp.set_edge(e, (a, b));
156                pp!(dp.fold_path(r).reverse.sum);
157            }
158        }
159    }
160}