Skip to main content

celox_analysis/
dag_schedule.rs

1//! Bottom-up scheduling of a dependency DAG with explicit value edges and
2//! shared materialization tokens.
3//!
4//! Hard dependencies constrain order.  Value dependencies are a subset of
5//! those edges and additionally describe liveness: scheduling a user backward
6//! makes the producer value live, while scheduling the producer kills it.  A
7//! ready node with the smallest live-value delta is selected first.  This is
8//! the conventional register-pressure list-scheduling model without attaching
9//! target-specific widths to source IR values.  A materialization token is
10//! used by an arbitrary set of nodes rather than being defined by one DAG
11//! node.  It models cached source expressions whose producer is whichever use
12//! gets scheduled first.
13//!
14//! For `N` nodes, `E` hard edges, `V` value edges, and `I` node/token
15//! incidences, scheduling costs `O((N + E + V + I) log N)` time and
16//! `O(N + E + V + I)` space.  A value or token changes the contribution to
17//! ready-node priorities at most twice, so each incidence is revisited only a
18//! constant number of times.
19
20use std::cmp::Reverse;
21use std::collections::BTreeSet;
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
24pub enum DagScheduleError {
25    Shape,
26    InvalidNode,
27    InvalidToken,
28    DuplicateDependency,
29    DuplicateToken,
30    ValueIsNotDependency,
31    UsersAreNotReverseDependencies,
32    Cycle,
33    ArithmeticOverflow,
34}
35
36trait NodeRows {
37    fn len(&self) -> usize;
38    fn row(&self, node: usize) -> &[usize];
39}
40
41struct LocalNodeRows<'a>(&'a [Vec<usize>]);
42
43impl NodeRows for LocalNodeRows<'_> {
44    fn len(&self) -> usize {
45        self.0.len()
46    }
47
48    fn row(&self, node: usize) -> &[usize] {
49        &self.0[node]
50    }
51}
52
53/// Borrow rows indexed by a larger graph through a local-to-external node map.
54///
55/// The row values themselves are not remapped. This is suitable for stable
56/// global IDs such as materialization tokens.
57#[derive(Clone, Copy)]
58pub struct MappedNodeRows<'a> {
59    rows_by_external_node: &'a [Vec<usize>],
60    external_by_local: &'a [usize],
61}
62
63impl<'a> MappedNodeRows<'a> {
64    pub fn new(
65        rows_by_external_node: &'a [Vec<usize>],
66        external_by_local: &'a [usize],
67    ) -> Result<Self, DagScheduleError> {
68        if external_by_local
69            .iter()
70            .any(|external| *external >= rows_by_external_node.len())
71        {
72            return Err(DagScheduleError::InvalidNode);
73        }
74        Ok(Self {
75            rows_by_external_node,
76            external_by_local,
77        })
78    }
79}
80
81impl NodeRows for MappedNodeRows<'_> {
82    fn len(&self) -> usize {
83        self.external_by_local.len()
84    }
85
86    fn row(&self, node: usize) -> &[usize] {
87        &self.rows_by_external_node[self.external_by_local[node]]
88    }
89}
90
91trait GraphRows {
92    type Iter<'a>: Iterator<Item = usize>
93    where
94        Self: 'a;
95
96    fn len(&self) -> usize;
97    fn nodes(&self, row: usize) -> Self::Iter<'_>;
98    fn contains(&self, row: usize, node: usize) -> bool;
99}
100
101struct LocalGraphRows<'a>(&'a [Vec<usize>]);
102
103impl<'a> LocalGraphRows<'a> {
104    fn new(rows: &'a [Vec<usize>]) -> Result<Self, DagScheduleError> {
105        for row in rows {
106            validate_row(row, rows.len())?;
107        }
108        Ok(Self(rows))
109    }
110}
111
112impl GraphRows for LocalGraphRows<'_> {
113    type Iter<'a>
114        = std::iter::Copied<std::slice::Iter<'a, usize>>
115    where
116        Self: 'a;
117
118    fn len(&self) -> usize {
119        self.0.len()
120    }
121
122    fn nodes(&self, row: usize) -> Self::Iter<'_> {
123        self.0[row].iter().copied()
124    }
125
126    fn contains(&self, row: usize, node: usize) -> bool {
127        self.0[row].binary_search(&node).is_ok()
128    }
129}
130
131struct MappedGraphRowsIter<'a> {
132    external_nodes: std::slice::Iter<'a, usize>,
133    local_by_external: &'a [usize],
134}
135
136impl Iterator for MappedGraphRowsIter<'_> {
137    type Item = usize;
138
139    fn next(&mut self) -> Option<Self::Item> {
140        self.external_nodes.find_map(|external| {
141            let local = self.local_by_external[*external];
142            (local != usize::MAX).then_some(local)
143        })
144    }
145}
146
147/// Borrow adjacency rows from a larger graph and expose only nodes mapped into
148/// the current local DAG.
149#[derive(Clone, Copy)]
150pub struct MappedGraphRows<'a> {
151    rows_by_external_node: &'a [Vec<usize>],
152    external_by_local: &'a [usize],
153    local_by_external: &'a [usize],
154}
155
156impl<'a> MappedGraphRows<'a> {
157    pub fn new(
158        rows_by_external_node: &'a [Vec<usize>],
159        external_by_local: &'a [usize],
160        local_by_external: &'a [usize],
161    ) -> Result<Self, DagScheduleError> {
162        for (local, &external) in external_by_local.iter().enumerate() {
163            if external >= rows_by_external_node.len()
164                || external >= local_by_external.len()
165                || local_by_external[external] != local
166            {
167                return Err(DagScheduleError::InvalidNode);
168            }
169            validate_mapped_user_row(
170                &rows_by_external_node[external],
171                external_by_local,
172                local_by_external,
173            )?;
174        }
175        Ok(Self {
176            rows_by_external_node,
177            external_by_local,
178            local_by_external,
179        })
180    }
181}
182
183impl GraphRows for MappedGraphRows<'_> {
184    type Iter<'a>
185        = MappedGraphRowsIter<'a>
186    where
187        Self: 'a;
188
189    fn len(&self) -> usize {
190        self.external_by_local.len()
191    }
192
193    fn nodes(&self, row: usize) -> Self::Iter<'_> {
194        MappedGraphRowsIter {
195            external_nodes: self.rows_by_external_node[self.external_by_local[row]].iter(),
196            local_by_external: self.local_by_external,
197        }
198    }
199
200    fn contains(&self, row: usize, node: usize) -> bool {
201        self.rows_by_external_node[self.external_by_local[row]]
202            .binary_search(&self.external_by_local[node])
203            .is_ok()
204    }
205}
206
207/// Return a deterministic forward order for one DAG.
208///
209/// `dependencies[user]` contains every predecessor which must be scheduled
210/// before `user`. `value_dependencies[user]` contains the subset whose result
211/// remains live until this user. Rows must be sorted and duplicate-free.
212pub fn schedule_min_live_values(
213    dependencies: &[Vec<usize>],
214    value_dependencies: &[Vec<usize>],
215) -> Result<Vec<usize>, DagScheduleError> {
216    let tokens = vec![Vec::new(); dependencies.len()];
217    schedule_min_live_values_and_tokens(dependencies, value_dependencies, &tokens, &[])
218}
219
220/// Return a deterministic forward order while also minimizing shared
221/// materialization-token pressure.
222///
223/// `tokens_by_node[node]` contains the sorted, duplicate-free token IDs used
224/// by that node. `token_weights[token]` is the number of pressure units kept
225/// live between the first and last scheduled users of the token.
226pub fn schedule_min_live_values_and_tokens(
227    dependencies: &[Vec<usize>],
228    value_dependencies: &[Vec<usize>],
229    tokens_by_node: &[Vec<usize>],
230    token_weights: &[usize],
231) -> Result<Vec<usize>, DagScheduleError> {
232    let dependencies = LocalGraphRows::new(dependencies)?;
233    let value_dependencies = LocalGraphRows::new(value_dependencies)?;
234    let tokens_by_node = LocalNodeRows(tokens_by_node);
235    let token_users = validate_inputs(
236        &dependencies,
237        &value_dependencies,
238        &tokens_by_node,
239        token_weights,
240    )?;
241    let node_count = dependencies.len();
242    let mut successors = vec![Vec::<usize>::new(); node_count];
243    let mut value_users = vec![Vec::<usize>::new(); node_count];
244    for user in 0..node_count {
245        for definition in dependencies.nodes(user) {
246            successors[definition].push(user);
247        }
248        for definition in value_dependencies.nodes(user) {
249            value_users[definition].push(user);
250        }
251    }
252    let successors = LocalGraphRows::new(&successors)?;
253    let value_users = LocalGraphRows::new(&value_users)?;
254    schedule_with_users(
255        &dependencies,
256        &value_dependencies,
257        &tokens_by_node,
258        token_weights,
259        &token_users,
260        &successors,
261        &value_users,
262    )
263}
264
265/// Schedule a local DAG while borrowing dependency, user, and token rows from
266/// a larger graph.
267///
268/// This avoids reconstructing either direction of the hard- and value-edge
269/// adjacency. The mapped user rows must be the exact reverse of the mapped
270/// dependency rows.
271pub fn schedule_min_live_values_and_tokens_with_mapped_rows(
272    dependencies: MappedGraphRows<'_>,
273    value_dependencies: MappedGraphRows<'_>,
274    tokens_by_node: MappedNodeRows<'_>,
275    token_weights: &[usize],
276    successors: MappedGraphRows<'_>,
277    value_users: MappedGraphRows<'_>,
278) -> Result<Vec<usize>, DagScheduleError> {
279    let token_users = validate_inputs(
280        &dependencies,
281        &value_dependencies,
282        &tokens_by_node,
283        token_weights,
284    )?;
285    validate_reverse_users(&successors, &dependencies)?;
286    validate_reverse_users(&value_users, &value_dependencies)?;
287    schedule_with_users(
288        &dependencies,
289        &value_dependencies,
290        &tokens_by_node,
291        token_weights,
292        &token_users,
293        &successors,
294        &value_users,
295    )
296}
297
298fn validate_inputs<D: GraphRows, V: GraphRows, T: NodeRows>(
299    dependencies: &D,
300    value_dependencies: &V,
301    tokens_by_node: &T,
302    token_weights: &[usize],
303) -> Result<Vec<Vec<usize>>, DagScheduleError> {
304    let node_count = dependencies.len();
305    if value_dependencies.len() != node_count || tokens_by_node.len() != node_count {
306        return Err(DagScheduleError::Shape);
307    }
308    let mut token_users = vec![Vec::<usize>::new(); token_weights.len()];
309    for user in 0..node_count {
310        validate_token_row(tokens_by_node.row(user), token_weights.len())?;
311        if value_dependencies
312            .nodes(user)
313            .any(|value| !dependencies.contains(user, value))
314        {
315            return Err(DagScheduleError::ValueIsNotDependency);
316        }
317        for &token in tokens_by_node.row(user) {
318            token_users[token].push(user);
319        }
320    }
321    Ok(token_users)
322}
323
324fn validate_reverse_users<U: GraphRows, D: GraphRows>(
325    users: &U,
326    dependencies: &D,
327) -> Result<(), DagScheduleError> {
328    if users.len() != dependencies.len() {
329        return Err(DagScheduleError::Shape);
330    }
331    let expected_edges = (0..dependencies.len()).try_fold(0usize, |total, row| {
332        total
333            .checked_add(dependencies.nodes(row).count())
334            .ok_or(DagScheduleError::ArithmeticOverflow)
335    })?;
336    let mut actual_edges = 0usize;
337    for definition in 0..users.len() {
338        for user in users.nodes(definition) {
339            if !dependencies.contains(user, definition) {
340                return Err(DagScheduleError::UsersAreNotReverseDependencies);
341            }
342            actual_edges = actual_edges
343                .checked_add(1)
344                .ok_or(DagScheduleError::ArithmeticOverflow)?;
345        }
346    }
347    if actual_edges != expected_edges {
348        return Err(DagScheduleError::UsersAreNotReverseDependencies);
349    }
350    Ok(())
351}
352
353#[allow(clippy::too_many_arguments)]
354fn schedule_with_users<D: GraphRows, V: GraphRows, T: NodeRows, S: GraphRows, U: GraphRows>(
355    dependencies: &D,
356    value_dependencies: &V,
357    tokens_by_node: &T,
358    token_weights: &[usize],
359    token_users: &[Vec<usize>],
360    successors: &S,
361    value_users: &U,
362) -> Result<Vec<usize>, DagScheduleError> {
363    let node_count = dependencies.len();
364    if successors.len() != node_count || value_users.len() != node_count {
365        return Err(DagScheduleError::Shape);
366    }
367    let indegree = (0..node_count)
368        .map(|node| dependencies.nodes(node).count())
369        .collect();
370
371    let (topological, entry_depth) = topological_order(successors, indegree)?;
372    let mut exit_depth = vec![1usize; node_count];
373    for &node in topological.iter().rev() {
374        let successor_depth = exit_depth[node]
375            .checked_add(1)
376            .ok_or(DagScheduleError::ArithmeticOverflow)?;
377        for dependency in dependencies.nodes(node) {
378            exit_depth[dependency] = exit_depth[dependency].max(successor_depth);
379        }
380    }
381
382    let mut unscheduled_users = (0..node_count)
383        .map(|node| successors.nodes(node).count())
384        .collect::<Vec<_>>();
385    let mut live = vec![false; node_count];
386    let mut remaining_token_users = token_users.iter().map(Vec::len).collect::<Vec<_>>();
387    let mut live_tokens = vec![false; token_weights.len()];
388    let mut deltas = vec![0isize; node_count];
389    let mut present = vec![false; node_count];
390    let mut ready = BTreeSet::<(isize, Reverse<usize>, Reverse<usize>, Reverse<usize>)>::new();
391    for (node, users) in unscheduled_users.iter().enumerate() {
392        if *users == 0 {
393            insert_ready(
394                node,
395                dependencies,
396                value_dependencies,
397                tokens_by_node,
398                token_weights,
399                &live,
400                &live_tokens,
401                &remaining_token_users,
402                &entry_depth,
403                &exit_depth,
404                &mut deltas,
405                &mut present,
406                &mut ready,
407            )?;
408        }
409    }
410
411    let mut reverse = Vec::with_capacity(node_count);
412    while let Some(&(_, _, _, Reverse(selected))) = ready.first() {
413        remove_ready(
414            selected,
415            &entry_depth,
416            &exit_depth,
417            &deltas,
418            &mut present,
419            &mut ready,
420        );
421
422        if live[selected] {
423            live[selected] = false;
424            update_for_value(
425                selected,
426                1,
427                value_users,
428                &entry_depth,
429                &exit_depth,
430                &mut deltas,
431                &present,
432                &mut ready,
433            )?;
434        }
435        for definition in value_dependencies.nodes(selected) {
436            if !live[definition] {
437                live[definition] = true;
438                update_for_value(
439                    definition,
440                    -1,
441                    value_users,
442                    &entry_depth,
443                    &exit_depth,
444                    &mut deltas,
445                    &present,
446                    &mut ready,
447                )?;
448            }
449        }
450        for &token in tokens_by_node.row(selected) {
451            let before = remaining_token_users[token];
452            remaining_token_users[token] = before
453                .checked_sub(1)
454                .ok_or(DagScheduleError::ArithmeticOverflow)?;
455            let weight = isize::try_from(token_weights[token])
456                .map_err(|_| DagScheduleError::ArithmeticOverflow)?;
457            if !live_tokens[token] && before > 1 {
458                live_tokens[token] = true;
459                let adjustment = if before == 2 {
460                    weight
461                        .checked_mul(-2)
462                        .ok_or(DagScheduleError::ArithmeticOverflow)?
463                } else {
464                    -weight
465                };
466                update_for_token(
467                    token,
468                    adjustment,
469                    token_users,
470                    &entry_depth,
471                    &exit_depth,
472                    &mut deltas,
473                    &present,
474                    &mut ready,
475                )?;
476            } else if live_tokens[token] && before == 2 {
477                update_for_token(
478                    token,
479                    -weight,
480                    token_users,
481                    &entry_depth,
482                    &exit_depth,
483                    &mut deltas,
484                    &present,
485                    &mut ready,
486                )?;
487            } else if live_tokens[token] && before == 1 {
488                live_tokens[token] = false;
489            }
490        }
491        for dependency in dependencies.nodes(selected) {
492            unscheduled_users[dependency] = unscheduled_users[dependency]
493                .checked_sub(1)
494                .ok_or(DagScheduleError::ArithmeticOverflow)?;
495            if unscheduled_users[dependency] == 0 {
496                insert_ready(
497                    dependency,
498                    dependencies,
499                    value_dependencies,
500                    tokens_by_node,
501                    token_weights,
502                    &live,
503                    &live_tokens,
504                    &remaining_token_users,
505                    &entry_depth,
506                    &exit_depth,
507                    &mut deltas,
508                    &mut present,
509                    &mut ready,
510                )?;
511            }
512        }
513        reverse.push(selected);
514    }
515
516    if reverse.len() != node_count {
517        return Err(DagScheduleError::Cycle);
518    }
519    reverse.reverse();
520    Ok(reverse)
521}
522
523fn validate_row(row: &[usize], node_count: usize) -> Result<(), DagScheduleError> {
524    let mut previous = None;
525    for &node in row {
526        if node >= node_count {
527            return Err(DagScheduleError::InvalidNode);
528        }
529        if previous.is_some_and(|previous| previous >= node) {
530            return Err(DagScheduleError::DuplicateDependency);
531        }
532        previous = Some(node);
533    }
534    Ok(())
535}
536
537fn validate_mapped_user_row(
538    row: &[usize],
539    external_by_local: &[usize],
540    local_by_external: &[usize],
541) -> Result<(), DagScheduleError> {
542    let mut previous = None;
543    for &external in row {
544        if external >= local_by_external.len() {
545            return Err(DagScheduleError::InvalidNode);
546        }
547        if previous.is_some_and(|previous| previous >= external) {
548            return Err(DagScheduleError::DuplicateDependency);
549        }
550        let local = local_by_external[external];
551        if local != usize::MAX
552            && (local >= external_by_local.len() || external_by_local[local] != external)
553        {
554            return Err(DagScheduleError::InvalidNode);
555        }
556        previous = Some(external);
557    }
558    Ok(())
559}
560
561fn validate_token_row(row: &[usize], token_count: usize) -> Result<(), DagScheduleError> {
562    let mut previous = None;
563    for &token in row {
564        if token >= token_count {
565            return Err(DagScheduleError::InvalidToken);
566        }
567        if previous.is_some_and(|previous| previous >= token) {
568            return Err(DagScheduleError::DuplicateToken);
569        }
570        previous = Some(token);
571    }
572    Ok(())
573}
574
575fn topological_order<U: GraphRows>(
576    successors: &U,
577    mut indegree: Vec<usize>,
578) -> Result<(Vec<usize>, Vec<usize>), DagScheduleError> {
579    let mut ready = BTreeSet::new();
580    let mut entry_depth = vec![1usize; successors.len()];
581    for (node, degree) in indegree.iter().enumerate() {
582        if *degree == 0 {
583            ready.insert(node);
584        }
585    }
586    let mut order = Vec::with_capacity(successors.len());
587    while let Some(node) = ready.pop_first() {
588        order.push(node);
589        let next_depth = entry_depth[node]
590            .checked_add(1)
591            .ok_or(DagScheduleError::ArithmeticOverflow)?;
592        for successor in successors.nodes(node) {
593            entry_depth[successor] = entry_depth[successor].max(next_depth);
594            indegree[successor] = indegree[successor]
595                .checked_sub(1)
596                .ok_or(DagScheduleError::ArithmeticOverflow)?;
597            if indegree[successor] == 0 {
598                ready.insert(successor);
599            }
600        }
601    }
602    if order.len() != successors.len() {
603        return Err(DagScheduleError::Cycle);
604    }
605    Ok((order, entry_depth))
606}
607
608#[allow(clippy::too_many_arguments)]
609fn insert_ready<D: GraphRows, V: GraphRows, T: NodeRows>(
610    node: usize,
611    dependencies: &D,
612    value_dependencies: &V,
613    tokens_by_node: &T,
614    token_weights: &[usize],
615    live: &[bool],
616    live_tokens: &[bool],
617    remaining_token_users: &[usize],
618    entry_depth: &[usize],
619    exit_depth: &[usize],
620    deltas: &mut [isize],
621    present: &mut [bool],
622    ready: &mut BTreeSet<(isize, Reverse<usize>, Reverse<usize>, Reverse<usize>)>,
623) -> Result<(), DagScheduleError> {
624    if present[node] {
625        return Err(DagScheduleError::DuplicateDependency);
626    }
627    let missing = value_dependencies
628        .nodes(node)
629        .filter(|definition| !live[*definition])
630        .count();
631    let removed = usize::from(live[node]);
632    let value_delta = isize::try_from(missing)
633        .ok()
634        .and_then(|missing| {
635            isize::try_from(removed)
636                .ok()
637                .and_then(|removed| missing.checked_sub(removed))
638        })
639        .ok_or(DagScheduleError::ArithmeticOverflow)?;
640    let token_delta = tokens_by_node
641        .row(node)
642        .iter()
643        .try_fold(0isize, |delta, &token| {
644            let contribution = if !live_tokens[token] && remaining_token_users[token] > 1 {
645                isize::try_from(token_weights[token])
646                    .map_err(|_| DagScheduleError::ArithmeticOverflow)?
647            } else if live_tokens[token] && remaining_token_users[token] == 1 {
648                -isize::try_from(token_weights[token])
649                    .map_err(|_| DagScheduleError::ArithmeticOverflow)?
650            } else {
651                0
652            };
653            delta
654                .checked_add(contribution)
655                .ok_or(DagScheduleError::ArithmeticOverflow)
656        })?;
657    let delta = value_delta
658        .checked_add(token_delta)
659        .ok_or(DagScheduleError::ArithmeticOverflow)?;
660    debug_assert!(
661        value_dependencies
662            .nodes(node)
663            .all(|value| dependencies.contains(node, value))
664    );
665    deltas[node] = delta;
666    present[node] = true;
667    ready.insert((
668        delta,
669        Reverse(exit_depth[node]),
670        Reverse(entry_depth[node]),
671        Reverse(node),
672    ));
673    Ok(())
674}
675
676#[allow(clippy::too_many_arguments)]
677fn update_for_token(
678    token: usize,
679    adjustment: isize,
680    token_users: &[Vec<usize>],
681    entry_depth: &[usize],
682    exit_depth: &[usize],
683    deltas: &mut [isize],
684    present: &[bool],
685    ready: &mut BTreeSet<(isize, Reverse<usize>, Reverse<usize>, Reverse<usize>)>,
686) -> Result<(), DagScheduleError> {
687    for &candidate in &token_users[token] {
688        if !present[candidate] {
689            continue;
690        }
691        ready.remove(&(
692            deltas[candidate],
693            Reverse(exit_depth[candidate]),
694            Reverse(entry_depth[candidate]),
695            Reverse(candidate),
696        ));
697        deltas[candidate] = deltas[candidate]
698            .checked_add(adjustment)
699            .ok_or(DagScheduleError::ArithmeticOverflow)?;
700        ready.insert((
701            deltas[candidate],
702            Reverse(exit_depth[candidate]),
703            Reverse(entry_depth[candidate]),
704            Reverse(candidate),
705        ));
706    }
707    Ok(())
708}
709
710fn remove_ready(
711    node: usize,
712    entry_depth: &[usize],
713    exit_depth: &[usize],
714    deltas: &[isize],
715    present: &mut [bool],
716    ready: &mut BTreeSet<(isize, Reverse<usize>, Reverse<usize>, Reverse<usize>)>,
717) {
718    debug_assert!(present[node]);
719    ready.remove(&(
720        deltas[node],
721        Reverse(exit_depth[node]),
722        Reverse(entry_depth[node]),
723        Reverse(node),
724    ));
725    present[node] = false;
726}
727
728#[allow(clippy::too_many_arguments)]
729fn update_for_value<U: GraphRows>(
730    value: usize,
731    adjustment: isize,
732    value_users: &U,
733    entry_depth: &[usize],
734    exit_depth: &[usize],
735    deltas: &mut [isize],
736    present: &[bool],
737    ready: &mut BTreeSet<(isize, Reverse<usize>, Reverse<usize>, Reverse<usize>)>,
738) -> Result<(), DagScheduleError> {
739    for candidate in value_users.nodes(value).chain(std::iter::once(value)) {
740        if !present[candidate] {
741            continue;
742        }
743        ready.remove(&(
744            deltas[candidate],
745            Reverse(exit_depth[candidate]),
746            Reverse(entry_depth[candidate]),
747            Reverse(candidate),
748        ));
749        deltas[candidate] = deltas[candidate]
750            .checked_add(adjustment)
751            .ok_or(DagScheduleError::ArithmeticOverflow)?;
752        ready.insert((
753            deltas[candidate],
754            Reverse(exit_depth[candidate]),
755            Reverse(entry_depth[candidate]),
756            Reverse(candidate),
757        ));
758    }
759    Ok(())
760}
761
762#[cfg(test)]
763mod tests {
764    use super::*;
765
766    fn rows(successors: &[Vec<usize>]) -> Vec<Vec<usize>> {
767        let mut result = vec![Vec::new(); successors.len()];
768        for (definition, users) in successors.iter().enumerate() {
769            for &user in users {
770                result[user].push(definition);
771            }
772        }
773        for row in &mut result {
774            row.sort_unstable();
775            row.dedup();
776        }
777        result
778    }
779
780    fn maximum_live(order: &[usize], value_dependencies: &[Vec<usize>]) -> usize {
781        let mut users = vec![0usize; value_dependencies.len()];
782        for row in value_dependencies {
783            for &definition in row {
784                users[definition] += 1;
785            }
786        }
787        let mut live = 0usize;
788        let mut maximum = 0usize;
789        for &node in order {
790            if users[node] != 0 {
791                live += 1;
792                maximum = maximum.max(live);
793            }
794            for &definition in &value_dependencies[node] {
795                users[definition] -= 1;
796                if users[definition] == 0 {
797                    live -= 1;
798                }
799            }
800        }
801        maximum
802    }
803
804    fn maximum_live_tokens(order: &[usize], tokens_by_node: &[Vec<usize>]) -> usize {
805        let token_count = tokens_by_node
806            .iter()
807            .flatten()
808            .copied()
809            .max()
810            .map_or(0, |maximum| maximum + 1);
811        let mut remaining = vec![0usize; token_count];
812        for tokens in tokens_by_node {
813            for &token in tokens {
814                remaining[token] += 1;
815            }
816        }
817        let mut live = vec![false; token_count];
818        let mut live_count = 0usize;
819        let mut maximum = 0usize;
820        for &node in order {
821            for &token in &tokens_by_node[node] {
822                if !live[token] && remaining[token] > 1 {
823                    live[token] = true;
824                    live_count += 1;
825                    maximum = maximum.max(live_count);
826                }
827                remaining[token] -= 1;
828                if live[token] && remaining[token] == 0 {
829                    live[token] = false;
830                    live_count -= 1;
831                }
832            }
833        }
834        maximum
835    }
836
837    #[test]
838    fn independent_single_use_chains_stay_contiguous() {
839        let successors = vec![vec![2], vec![3], vec![4], vec![5], vec![], vec![]];
840        let dependencies = rows(&successors);
841        let scheduled = schedule_min_live_values(&dependencies, &dependencies).unwrap();
842
843        assert!(
844            maximum_live(&scheduled, &dependencies)
845                < maximum_live(&[0, 1, 2, 3, 4, 5], &dependencies)
846        );
847        for chain in [[0, 2, 4], [1, 3, 5]] {
848            let positions =
849                chain.map(|node| scheduled.iter().position(|item| *item == node).unwrap());
850            assert_eq!(positions[1], positions[0] + 1);
851            assert_eq!(positions[2], positions[1] + 1);
852        }
853    }
854
855    #[test]
856    fn order_only_edges_do_not_create_live_values() {
857        let dependencies = vec![vec![], vec![0], vec![1]];
858        let values = vec![vec![], vec![], vec![]];
859        assert_eq!(
860            schedule_min_live_values(&dependencies, &values).unwrap(),
861            vec![0, 1, 2]
862        );
863        assert_eq!(maximum_live(&[0, 1, 2], &values), 0);
864    }
865
866    #[test]
867    fn mapped_rows_match_owned_reverse_relations() {
868        let dependencies = vec![vec![], vec![0], vec![0, 1]];
869        let values = dependencies.clone();
870        let tokens = vec![vec![0], vec![0, 1], vec![1]];
871        let expected =
872            schedule_min_live_values_and_tokens(&dependencies, &values, &tokens, &[2, 1]).unwrap();
873
874        let external_by_local = vec![4, 1, 3];
875        let local_by_external = vec![usize::MAX, 1, usize::MAX, 2, 0];
876        let mut external_users = vec![Vec::new(); 5];
877        external_users[1] = vec![3];
878        external_users[4] = vec![1, 3];
879        let mut external_dependencies = vec![Vec::new(); 5];
880        external_dependencies[1] = vec![4];
881        external_dependencies[3] = vec![1, 4];
882        let mut external_tokens = vec![Vec::new(); 5];
883        external_tokens[4] = vec![0];
884        external_tokens[1] = vec![0, 1];
885        external_tokens[3] = vec![1];
886
887        let actual = schedule_min_live_values_and_tokens_with_mapped_rows(
888            MappedGraphRows::new(
889                &external_dependencies,
890                &external_by_local,
891                &local_by_external,
892            )
893            .unwrap(),
894            MappedGraphRows::new(
895                &external_dependencies,
896                &external_by_local,
897                &local_by_external,
898            )
899            .unwrap(),
900            MappedNodeRows::new(&external_tokens, &external_by_local).unwrap(),
901            &[2, 1],
902            MappedGraphRows::new(&external_users, &external_by_local, &local_by_external).unwrap(),
903            MappedGraphRows::new(&external_users, &external_by_local, &local_by_external).unwrap(),
904        )
905        .unwrap();
906
907        assert_eq!(actual, expected);
908    }
909
910    #[test]
911    fn mapped_rows_reject_an_incomplete_reverse_relation() {
912        let dependencies = vec![vec![], vec![0]];
913        let external_by_local = vec![0, 1];
914        let local_by_external = vec![0, 1];
915        let missing_users = vec![Vec::new(), Vec::new()];
916        let empty_tokens = vec![Vec::new(), Vec::new()];
917        let dependencies =
918            MappedGraphRows::new(&dependencies, &external_by_local, &local_by_external).unwrap();
919        let users =
920            MappedGraphRows::new(&missing_users, &external_by_local, &local_by_external).unwrap();
921
922        assert_eq!(
923            schedule_min_live_values_and_tokens_with_mapped_rows(
924                dependencies,
925                dependencies,
926                MappedNodeRows::new(&empty_tokens, &external_by_local).unwrap(),
927                &[],
928                users,
929                users,
930            ),
931            Err(DagScheduleError::UsersAreNotReverseDependencies)
932        );
933    }
934
935    #[test]
936    fn mapped_rows_reject_an_out_of_range_reverse_map_entry() {
937        let users = vec![vec![1], vec![]];
938
939        assert!(matches!(
940            MappedGraphRows::new(&users, &[0], &[0, 1]),
941            Err(DagScheduleError::InvalidNode)
942        ));
943    }
944
945    #[test]
946    fn mapped_rows_reject_an_alias_in_the_reverse_map() {
947        let users = vec![vec![1], vec![]];
948
949        assert!(matches!(
950            MappedGraphRows::new(&users, &[0], &[0, 0]),
951            Err(DagScheduleError::InvalidNode)
952        ));
953    }
954
955    #[test]
956    fn mapped_rows_preserve_self_edges_for_cycle_detection() {
957        let dependencies = vec![vec![0]];
958        let values = vec![vec![]];
959        let tokens = vec![vec![]];
960        let external_by_local = [0];
961        let local_by_external = [0];
962        let mapped_dependencies =
963            MappedGraphRows::new(&dependencies, &external_by_local, &local_by_external).unwrap();
964        let mapped_values =
965            MappedGraphRows::new(&values, &external_by_local, &local_by_external).unwrap();
966
967        assert_eq!(
968            schedule_min_live_values_and_tokens(&dependencies, &values, &tokens, &[]),
969            Err(DagScheduleError::Cycle)
970        );
971        assert_eq!(
972            schedule_min_live_values_and_tokens_with_mapped_rows(
973                mapped_dependencies,
974                mapped_values,
975                MappedNodeRows::new(&tokens, &external_by_local).unwrap(),
976                &[],
977                mapped_dependencies,
978                mapped_values,
979            ),
980            Err(DagScheduleError::Cycle)
981        );
982    }
983
984    #[test]
985    fn rejects_a_value_edge_without_a_hard_dependency() {
986        assert_eq!(
987            schedule_min_live_values(&[vec![], vec![]], &[vec![], vec![0]]),
988            Err(DagScheduleError::ValueIsNotDependency)
989        );
990    }
991
992    #[test]
993    fn wide_independent_ready_set_schedules_without_pairwise_state() {
994        const NODES: usize = 4096;
995        let graph = vec![Vec::new(); NODES];
996        let order = schedule_min_live_values(&graph, &graph).unwrap();
997        assert_eq!(order, (0..NODES).collect::<Vec<_>>());
998    }
999
1000    #[test]
1001    fn shared_materializations_stay_contiguous() {
1002        let dependencies = vec![vec![], vec![], vec![], vec![]];
1003        let values = dependencies.clone();
1004        let tokens = vec![vec![0], vec![1], vec![0], vec![1]];
1005        let order =
1006            schedule_min_live_values_and_tokens(&dependencies, &values, &tokens, &[1, 1]).unwrap();
1007
1008        assert!(maximum_live_tokens(&order, &tokens) < maximum_live_tokens(&[0, 1, 2, 3], &tokens));
1009        for users in [[0, 2], [1, 3]] {
1010            let positions = users.map(|node| order.iter().position(|item| *item == node).unwrap());
1011            assert_eq!(positions[1], positions[0] + 1);
1012        }
1013    }
1014
1015    #[test]
1016    fn materialization_pressure_never_breaks_hard_dependencies() {
1017        let dependencies = vec![vec![], vec![0], vec![], vec![1, 2]];
1018        let values = vec![vec![], vec![], vec![], vec![]];
1019        let tokens = vec![vec![0], vec![1], vec![0], vec![1]];
1020        let order =
1021            schedule_min_live_values_and_tokens(&dependencies, &values, &tokens, &[8, 1]).unwrap();
1022        let position = |node| order.iter().position(|item| *item == node).unwrap();
1023
1024        assert!(position(0) < position(1));
1025        assert!(position(1) < position(3));
1026        assert!(position(2) < position(3));
1027    }
1028
1029    #[test]
1030    fn sparse_materialization_incidence_scales_to_wide_ready_sets() {
1031        const NODES: usize = 4096;
1032        let graph = vec![Vec::new(); NODES];
1033        let tokens = (0..NODES).map(|node| vec![node / 2]).collect::<Vec<_>>();
1034        let weights = vec![1; NODES / 2];
1035        let order = schedule_min_live_values_and_tokens(&graph, &graph, &tokens, &weights).unwrap();
1036
1037        assert_eq!(order.len(), NODES);
1038        assert_eq!(maximum_live_tokens(&order, &tokens), 1);
1039    }
1040}