From 9e4900548b9b139ba796482a49b228e197124000 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Wed, 22 Jul 2026 11:55:42 +0200 Subject: [PATCH 01/19] feat(solver): change reasoners api to allow plugging external, extra reasoners --- examples/scheduling/src/search.rs | 2 +- planning/planners/src/solver.rs | 6 +- planning/timelines/src/explain.rs | 2 +- solver/examples/tsp/main.rs | 4 +- solver/src/reasoners/mod.rs | 153 +++++++++++++++++++++++++++--- solver/src/solver/solver_impl.rs | 54 +++++++---- solver/src/solver/stats.rs | 13 +-- 7 files changed, 182 insertions(+), 52 deletions(-) diff --git a/examples/scheduling/src/search.rs b/examples/scheduling/src/search.rs index d5401716..8aa38917 100644 --- a/examples/scheduling/src/search.rs +++ b/examples/scheduling/src/search.rs @@ -82,7 +82,7 @@ pub fn get_solver(mut base_solver: Solver, strategy: &SearchStrategy, pb: &Encod "focused" => mode = Mode::Focused, x if x.starts_with("+lbd") => { let lvl = x.strip_prefix("+lbd").unwrap().parse().unwrap(); - base_solver.reasoners.sat.clauses.params.locked_lbd_level = lvl; + base_solver.reasoners.sat().clauses.params.locked_lbd_level = lvl; } "" => {} // ignore _ => panic!("Unsupported option: {opt}"), diff --git a/planning/planners/src/solver.rs b/planning/planners/src/solver.rs index e8bdd30a..21ea5ff9 100644 --- a/planning/planners/src/solver.rs +++ b/planning/planners/src/solver.rs @@ -381,7 +381,7 @@ pub fn init_solver(model: Model) -> Box { }; let mut solver = Box::new(aries_solver::solver::Solver::new(model)); - solver.reasoners.diff.config = stn_config; + solver.reasoners.diff().config = stn_config; solver } @@ -446,11 +446,11 @@ impl Strat { } Strat::ActivityBoolLight => { solver.set_brancher(ActivityBrancher::new_with_heuristic(ActivityBoolFirstHeuristic)); - solver.reasoners.diff.config.theory_propagation = TheoryPropagationLevel::Bounds; + solver.reasoners.diff().config.theory_propagation = TheoryPropagationLevel::Bounds; } Strat::Forward => { solver.set_brancher(ForwardSearcher::new(problem)); - solver.reasoners.diff.config.theory_propagation = TheoryPropagationLevel::Bounds; + solver.reasoners.diff().config.theory_propagation = TheoryPropagationLevel::Bounds; } Strat::Causal => { let strat = causal_brancher(problem, encoding); diff --git a/planning/timelines/src/explain.rs b/planning/timelines/src/explain.rs index 2c5504c8..1b678f18 100644 --- a/planning/timelines/src/explain.rs +++ b/planning/timelines/src/explain.rs @@ -46,7 +46,7 @@ impl ExplainableSolver { // enable stronger propagation than default in difference logic solver. // this is useful in planning models where bounds are not sufficient to reason on precedence between tasks - solver.reasoners.diff.config.theory_propagation = aries_solver::reasoners::stn::TheoryPropagationLevel::Full; + solver.reasoners.diff().config.theory_propagation = aries_solver::reasoners::stn::TheoryPropagationLevel::Full; Self { solver, diff --git a/solver/examples/tsp/main.rs b/solver/examples/tsp/main.rs index 836ae210..ee16de1f 100644 --- a/solver/examples/tsp/main.rs +++ b/solver/examples/tsp/main.rs @@ -202,9 +202,9 @@ fn solve_tsp(pb: &TspProblem, args: &Opt) -> Option { let mut solver = Solver::new(model); if args.no_lp { - solver.reasoners.lp.deactivate(); + solver.reasoners.lp().deactivate(); } else { - solver.reasoners.lp.activate(); + solver.reasoners.lp().activate(); } let solution_opt = match solver.minimize_with_callback( diff --git a/solver/src/reasoners/mod.rs b/solver/src/reasoners/mod.rs index 7769da4c..b77f1405 100644 --- a/solver/src/reasoners/mod.rs +++ b/solver/src/reasoners/mod.rs @@ -1,3 +1,5 @@ +use itertools::Itertools; + use crate::backtrack::Backtrack; use crate::core::Lit; use crate::core::state::{Cause, DomainsSnapshot, Explainer, InferenceCause}; @@ -24,6 +26,7 @@ pub enum ReasonerId { Cp, Tautologies, Lp, + Extra(u8), } impl ReasonerId { @@ -35,6 +38,7 @@ impl ReasonerId { impl Display for ReasonerId { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { use ReasonerId::*; + let mut _extra_str = String::new(); write!( f, "{}", @@ -44,12 +48,16 @@ impl Display for ReasonerId { Cp => "CP", Tautologies => "Optim", Lp => "LP", + Extra(i) => { + _extra_str = format!("Extra({i})"); + &_extra_str + } } ) } } -pub trait Theory: Backtrack + Send + 'static { +pub trait Theory: Backtrack + Send { fn identity(&self) -> ReasonerId; fn propagate(&mut self, model: &mut Domains) -> Result<(), Contradiction>; @@ -95,52 +103,168 @@ pub(crate) const REASONERS: [ReasonerId; 5] = [ ReasonerId::Lp, ]; -/// A set of inference modules for constraint propagation. -#[derive(Clone)] -pub struct Reasoners { +pub(crate) struct ReasonersTheories { pub sat: SatSolver, pub diff: StnTheory, pub(crate) cp: Cp, pub(crate) tautologies: Tautologies, pub lp: Lp, + pub extra: Vec>>, } -impl Reasoners { +impl Clone for ReasonersTheories { + fn clone(&self) -> Self { + Self { + sat: self.sat.clone(), + diff: self.diff.clone(), + cp: self.cp.clone(), + tautologies: self.tautologies.clone(), + lp: self.lp.clone(), + extra: self + .extra + .iter() + .map(|th| th.as_ref().map(|th| th.clone_box())) + .collect(), + } + } +} +impl ReasonersTheories { pub fn new() -> Self { - Reasoners { + ReasonersTheories { sat: SatSolver::new(ReasonerId::Sat), diff: StnTheory::new(Default::default()), cp: Cp::new(ReasonerId::Cp), tautologies: Tautologies::default(), lp: Lp::default(), + extra: vec![], } } + pub fn with_extra(extra: Vec>) -> Self { + assert!( + extra + .iter() + .map(|r| r.identity()) + .all(|rid| matches!(rid, ReasonerId::Extra(_))) + ); + assert!(extra.iter().map(|r| r.identity()).all_unique()); + + let extra = { + let mut res = vec![]; + for r in extra { + let ReasonerId::Extra(i) = r.identity() else { + unreachable!() + }; + let i = i as usize; + while res.len() <= i { + res.push(None); + } + res[i] = Some(r); + } + res + }; - pub fn reasoner(&self, id: ReasonerId) -> &dyn Theory { + ReasonersTheories { + sat: SatSolver::new(ReasonerId::Sat), + diff: StnTheory::new(Default::default()), + cp: Cp::new(ReasonerId::Cp), + tautologies: Tautologies::default(), + lp: Lp::default(), + extra, + } + } + pub fn get(&self, id: ReasonerId) -> &dyn Theory { match id { ReasonerId::Sat => &self.sat, ReasonerId::Diff => &self.diff, ReasonerId::Cp => &self.cp, ReasonerId::Tautologies => &self.tautologies, ReasonerId::Lp => &self.lp, + ReasonerId::Extra(id) => self.extra.get(id as usize).unwrap().as_ref().unwrap().as_ref(), } } - - pub fn reasoner_mut(&mut self, id: ReasonerId) -> &mut dyn Theory { + pub fn get_mut(&mut self, id: ReasonerId) -> &mut dyn Theory { match id { ReasonerId::Sat => &mut self.sat, ReasonerId::Diff => &mut self.diff, ReasonerId::Cp => &mut self.cp, ReasonerId::Tautologies => &mut self.tautologies, ReasonerId::Lp => &mut self.lp, + ReasonerId::Extra(id) => self.extra.get_mut(id as usize).unwrap().as_mut().unwrap().as_mut(), + } + } +} + +#[derive(Clone)] +pub(crate) struct ReasonersWriters { + writers: Vec, +} +impl ReasonersWriters { + pub fn new() -> Self { + Self { + writers: REASONERS.to_vec(), + } + } + pub fn with_extra(extra: &[Box]) -> Self { + assert!( + extra + .iter() + .map(|r| r.identity()) + .all(|rid| matches!(rid, ReasonerId::Extra(_))) + ); + assert!(extra.iter().map(|r| r.identity()).all_unique()); + + Self { + writers: REASONERS + .into_iter() + .chain(extra.iter().map(|r| r.identity())) + .collect(), } } + pub fn get(&self) -> &[ReasonerId] { + &self.writers + } +} - pub fn writers(&self) -> &'static [ReasonerId] { - &REASONERS +/// A set of inference modules for constraint propagation. +#[derive(Clone)] +pub struct Reasoners { + pub(crate) writers: ReasonersWriters, + pub(crate) theories: ReasonersTheories, +} +impl Reasoners { + pub fn new() -> Self { + Self { + writers: ReasonersWriters::new(), + theories: ReasonersTheories::new(), + } + } + pub fn with_extra(extra: Vec>) -> Self { + Self { + writers: ReasonersWriters::with_extra(&extra), + theories: ReasonersTheories::with_extra(extra), + } + } + + pub fn sat(&mut self) -> &mut SatSolver { + &mut self.theories.sat + } + pub fn diff(&mut self) -> &mut StnTheory { + &mut self.theories.diff + } + pub fn cp(&mut self) -> &mut Cp { + &mut self.theories.cp + } + pub fn tautologies(&mut self) -> &mut Tautologies { + &mut self.theories.tautologies + } + pub fn lp(&mut self) -> &mut Lp { + &mut self.theories.lp + } + pub fn extra(&mut self) -> impl Iterator> + '_ { + self.theories.extra.iter_mut().filter_map(|th| th.as_mut()) } - pub fn theories(&self) -> impl Iterator + '_ { - self.writers().iter().map(|w| (*w, self.reasoner(*w))) + pub fn iter(&self) -> impl Iterator + '_ { + self.writers.get().iter().map(|w| (*w, self.theories.get(*w))) } } @@ -152,7 +276,8 @@ impl Default for Reasoners { impl Explainer for Reasoners { fn explain(&mut self, cause: InferenceCause, literal: Lit, model: &DomainsSnapshot, explanation: &mut Explanation) { - self.reasoner_mut(cause.writer) + self.theories + .get_mut(cause.writer) .explain(literal, cause, model, explanation) } } diff --git a/solver/src/solver/solver_impl.rs b/solver/src/solver/solver_impl.rs index 34c1886c..b5ffec06 100644 --- a/solver/src/solver/solver_impl.rs +++ b/solver/src/solver/solver_impl.rs @@ -124,14 +124,24 @@ pub struct Solver { } impl Solver { pub fn new(model: Model) -> Solver { + Self::with_extra_reasoners(model, vec![]) + } + + pub fn with_extra_reasoners( + model: Model, + extra_reasoners: Vec>, + ) -> Solver { + let reasoners = Reasoners::with_extra(extra_reasoners); + let stats = Stats::with_reasoners(&reasoners); + Solver { model, next_unposted_constraint: 0, brancher: default_brancher(), - reasoners: Reasoners::new(), + reasoners, decision_level: DecLvl::ROOT, last_assumption_level: DecLvl::ROOT, - stats: Default::default(), + stats, sync: Synchro::new(), } } @@ -192,7 +202,7 @@ impl Solver { Constraint::Propagator(user_propagator) => { // black-box propagator, there is nothing we can do except posting it to the CP solver for prop in user_propagator.get_propagators() { - self.reasoners.cp.add_propagator(prop); + self.reasoners.cp().add_propagator(prop); } return Ok(()); } @@ -261,7 +271,7 @@ impl Solver { debug_assert!(factor > 0); // cst + factor *(b - a) <= 0 // (b - a) <= -cst / factor - self.reasoners.diff.add_half_reified_edge( + self.reasoners.diff().add_half_reified_edge( enabler, a, b, @@ -311,7 +321,9 @@ impl Solver { }; // add a dynamic edge to the STN, specifying that `tgt -src <= ub_var * ub_factor` // Each time a new upper bound is inferred on `ub_var` a new edge will temporarily added. - self.reasoners.diff.add_dynamic_edge(src, tgt, ub_var, ub_factor, doms) + self.reasoners + .diff() + .add_dynamic_edge(src, tgt, ub_var, ub_factor, doms) } None // we posted redendant constraint, but this is not an exact reformulation, we need to post the original one as well } else { @@ -321,13 +333,13 @@ impl Solver { self.post_constraint(&Constraint::HalfReified(reformulated, enabler)) } else { self.reasoners - .cp + .cp() .add_half_reif_linear_leq_constraint(lin, enabler, &self.model.state); // We check that enabler lit is always present as the lp reasonner can't handle optionnality if self.model.state.presence(enabler) == Lit::TRUE { self.reasoners - .lp + .lp() .add_linear_leq_constraint(lin, enabler, &self.model.state); } @@ -358,7 +370,7 @@ impl Solver { if propagatable.is_empty() { return self.model.state.set(!scope, Cause::Encoding).map(|_| ()); } - self.reasoners.sat.add_clause_scoped(propagatable.literals(), scope); + self.reasoners.sat().add_clause_scoped(propagatable.literals(), scope); Ok(()) } @@ -508,7 +520,7 @@ impl Solver { if let Some(dl) = self.backtrack_level_for_clause(&clause) { self.restore(dl); - self.reasoners.sat.add_clause(&clause); + self.reasoners.sat().add_clause(&clause); } else { return Ok(sat); } @@ -664,7 +676,7 @@ impl Solver { return Err(Exit::Interrupted); } InputSignal::LearnedClause(cl) => { - self.reasoners.sat.add_forgettable_clause(cl.literals()); + self.reasoners.sat().add_forgettable_clause(cl.literals()); requires_new_propagation = true; } InputSignal::SolutionFound(assignment) => { @@ -1010,10 +1022,10 @@ impl Solver { // clauses with a single literal are tautologies and can be given to the dedicated reasoner // note: a possible optimization would also be to not backjump to the root (always the case with a such clauses) // but instead to the first level where imposing it would not result in a conflict - self.reasoners.tautologies.add_tautology(expl.clause.literals()[0]) + self.reasoners.tautologies().add_tautology(expl.clause.literals()[0]) } else { // add clause to sat solver, making sure the asserted literal is set to true - self.reasoners.sat.add_learnt_clause(expl.clause.literals()); + self.reasoners.sat().add_learnt_clause(expl.clause.literals()); } true @@ -1094,16 +1106,16 @@ impl Solver { let num_events_at_start = self.model.state.num_events(); debug_assert_eq!( - self.reasoners.writers().iter().next(), + self.reasoners.writers.get().iter().next(), Some(&ReasonerId::Sat), "SAT propagator should propagate first to ensure none of its invariant are violated by others." ); // propagate all theories - for &i in self.reasoners.writers() { + for &i in self.reasoners.writers.get() { let trail_size = self.model.state.trail().len() as u64; let theory_propagation_start = StartCycleCount::now(); self.stats[i].propagation_loops += 1; - let th = self.reasoners.reasoner_mut(i); + let th = self.reasoners.theories.get_mut(i); match th.propagate(&mut self.model.state) { Ok(()) => (), @@ -1149,7 +1161,7 @@ impl Solver { pub fn print_stats(&self) { println!("{}", self.stats); - for (i, th) in self.reasoners.theories() { + for (i, th) in self.reasoners.iter() { println!("====== {i} ====="); th.print_stats(); } @@ -1170,8 +1182,8 @@ impl Backtrack for Solver { assert_eq!(self.model.save_state(), n); assert_eq!(self.brancher.save_state(), n); - for w in self.reasoners.writers() { - let th = self.reasoners.reasoner_mut(*w); + for w in self.reasoners.writers.get() { + let th = self.reasoners.theories.get_mut(*w); assert_eq!(th.save_state(), n); } n @@ -1182,7 +1194,7 @@ impl Backtrack for Solver { let n = self.decision_level.to_int(); assert_eq!(self.model.num_saved(), n); assert_eq!(self.brancher.num_saved(), n); - for (_, th) in self.reasoners.theories() { + for (_, th) in self.reasoners.iter() { assert_eq!(th.num_saved(), n); } true @@ -1202,8 +1214,8 @@ impl Backtrack for Solver { } self.model.restore(saved_id); self.brancher.restore(saved_id); - for w in self.reasoners.writers() { - let th = self.reasoners.reasoner_mut(*w); + for w in self.reasoners.writers.get() { + let th = self.reasoners.theories.get_mut(*w); th.restore(saved_id); } debug_assert_eq!(self.current_decision_level(), saved_id); diff --git a/solver/src/solver/stats.rs b/solver/src/solver/stats.rs index b96ebde2..8012faaa 100644 --- a/solver/src/solver/stats.rs +++ b/solver/src/solver/stats.rs @@ -1,7 +1,6 @@ use crate::backtrack::DecLvl; use crate::core::{IntCst, Lit}; -use crate::reasoners::REASONERS; -use crate::reasoners::ReasonerId; +use crate::reasoners::{ReasonerId, Reasoners}; use crate::utils::cpu_time::*; use aries_env_param::EnvParam; use std::collections::BTreeMap; @@ -39,9 +38,9 @@ pub struct ModuleStat { } impl Stats { - pub fn new() -> Stats { + pub fn with_reasoners(reasoners: &Reasoners) -> Stats { let mut per_mod = BTreeMap::new(); - for id in &REASONERS { + for id in reasoners.writers.get() { per_mod.insert(*id, ModuleStat::default()); } @@ -110,12 +109,6 @@ impl Stats { } } -impl Default for Stats { - fn default() -> Self { - Self::new() - } -} - impl Display for Stats { fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), Error> { // a right padded number with the given format string From 7c7fef6f6d9f260286393fe039144507c2d8a8d8 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 12:31:12 +0200 Subject: [PATCH 02/19] chore(solver): expose watches beyond `aries_solver` crate notably for use by external reasoners --- solver/src/core/literals.rs | 2 +- solver/src/core/literals/watches.rs | 6 +++--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/solver/src/core/literals.rs b/solver/src/core/literals.rs index fea19531..52465989 100644 --- a/solver/src/core/literals.rs +++ b/solver/src/core/literals.rs @@ -17,7 +17,7 @@ pub use conjunction::*; pub use disjunction::*; pub(crate) use implication_graph::*; pub use lit_set::*; -pub(crate) use watches::*; +pub use watches::*; use crate::prelude::*; use smallvec::SmallVec; diff --git a/solver/src/core/literals/watches.rs b/solver/src/core/literals/watches.rs index 478281d6..2de5eb31 100644 --- a/solver/src/core/literals/watches.rs +++ b/solver/src/core/literals/watches.rs @@ -4,7 +4,7 @@ use crate::core::*; /// A set of literals watches on bound changes. /// The event watches are all on the same bound (i.e. the lower or the upper bound) of a single variable. #[derive(Clone)] -pub(crate) struct WatchSet { +pub struct WatchSet { watches: Vec>, } impl WatchSet { @@ -82,7 +82,7 @@ impl Default for WatchSet { } #[derive(Copy, Clone)] -pub(crate) struct Watch { +pub struct Watch { pub(crate) watcher: Watcher, /// upper bound guard: IntCst, @@ -95,7 +95,7 @@ impl Watch { /// A datastructure for implementing watches, functionnally equivalent to a `Map>` #[derive(Clone)] -pub(crate) struct Watches { +pub struct Watches { watches: RefVec>, empty_watch_set: WatchSet, } From d8650048bb86fbadbd780bde36480a777d8a9bcb Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 12:32:21 +0200 Subject: [PATCH 03/19] feat(lp-highs): add an LP relaxation reasoner backed by HiGHS --- Cargo.toml | 1 + solver/lp_highs/Cargo.toml | 13 + solver/lp_highs/src/bindings.rs | 140 +++++++ solver/lp_highs/src/lib.rs | 633 ++++++++++++++++++++++++++++++++ solver/lp_highs/src/state.rs | 334 +++++++++++++++++ solver/lp_highs/src/types.rs | 132 +++++++ 6 files changed, 1253 insertions(+) create mode 100644 solver/lp_highs/Cargo.toml create mode 100644 solver/lp_highs/src/bindings.rs create mode 100644 solver/lp_highs/src/lib.rs create mode 100644 solver/lp_highs/src/state.rs create mode 100644 solver/lp_highs/src/types.rs diff --git a/Cargo.toml b/Cargo.toml index 27d640c1..7a475413 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -2,6 +2,7 @@ members = [ "solver", + "solver/lp_highs", "bench/data", "bench/bench", "utils/env_param", diff --git a/solver/lp_highs/Cargo.toml b/solver/lp_highs/Cargo.toml new file mode 100644 index 00000000..f486b58b --- /dev/null +++ b/solver/lp_highs/Cargo.toml @@ -0,0 +1,13 @@ +[package] +name = "aries-solver-lp-highs" +version = "0.1.0" +edition = "2024" + +[dependencies] +aries-solver = { path = ".." } +aries-env-param = { path = "../../utils/env_param" } +highs = { git = "https://github.com/nrealus/highs.git", branch = "dev" } +itertools = { workspace = true } +smallvec = { workspace = true } +idmap = { workspace = true } +tracing = { workspace = true } \ No newline at end of file diff --git a/solver/lp_highs/src/bindings.rs b/solver/lp_highs/src/bindings.rs new file mode 100644 index 00000000..af40c5ca --- /dev/null +++ b/solver/lp_highs/src/bindings.rs @@ -0,0 +1,140 @@ +use aries_solver::core::literals::Watches; +use aries_solver::core::views::Dom; + +use idmap::DirectIdMap; + +use crate::types::*; + +/// A binding from the main model to an LP column: whenever it applies, the column bound(s) it +/// states are pushed to the LP. +/// +/// Bindings only ever go in that direction. The LP is a relaxation of the main model, so a column +/// is constrained *by* the model and never the other way around. +/// +/// Both forms carry a `scope` and only apply while it is entailed, which is how the column of an +/// optional part of the model is left free rather than constrained while its presence is still +/// unknown. [`AriesLit::TRUE`] means "no scope". +#[derive(Clone, Copy, Debug)] +pub(super) enum Binding { + /// A *half* binding: the fixed column bound `lp_lit`, applied as soon as `scope` and `trigger` + /// are both entailed. + /// + /// This is the form to prefer. It is dispatched by a watch on each of its two literals, so it + /// costs nothing until one of them holds, and the cause it reports is exactly those two + /// literals: the weakest justification of the bound, rather than whatever the domain happened + /// to be when it fired. + Fixed { + scope: AriesLit, + trigger: AriesLit, + lp_lit: LpLit, + }, + /// A *full* binding: the column mirrors the variable's domain, both bounds following it. + /// + /// No finite set of watches can express this over a wide domain, so it is re-evaluated on every + /// event on `var` (or on the scope's variable). Each bound is justified by the variable's + /// matching bound at that moment, which is stronger than the trigger of an equivalent + /// [`Binding::Fixed`] would be, so prefer that one whenever the column only reacts to finitely + /// many thresholds — and note that this form constrains the column from *both* sides. + Tracking { scope: AriesLit, var: AriesVar, col: LpCol }, +} + +impl Binding { + /// The column bounds this binding currently states, each with the two main-model literals that + /// justify it. Empty when the binding does not apply. + fn eval(&self, dom: &impl Dom) -> impl Iterator { + let (lower, upper) = match *self { + Binding::Fixed { scope, trigger, lp_lit } => ( + (dom.entails(scope) && dom.entails(trigger)).then_some((lp_lit, scope, trigger)), + None, + ), + Binding::Tracking { scope, var, col } => { + if dom.entails(scope) { + let (lb, ub) = (dom.lb(var), dom.ub(var)); + ( + Some((LpLit::geq(col, int_cst_as_long(lb)), scope, AriesLit::geq(var, lb))), + Some((LpLit::leq(col, int_cst_as_long(ub)), scope, AriesLit::leq(var, ub))), + ) + } else { + (None, None) + } + } + }; + lower.into_iter().chain(upper) + } +} + +/// The bindings of an [`crate::LpRelax`] reasoner, indexed by what can make them apply. +#[derive(Default, Clone)] +pub(super) struct Bindings { + bindings: Vec, + /// Watches on the literals of the [`Binding::Fixed`] bindings, the watcher being the index of the binding. + /// A binding watching two literals is woken by either of them, and its evaluation then checks that the other one holds as well. + watches: Watches, + /// Indices of the [`Binding::Tracking`] bindings, by the variable they mirror. + on_var: DirectIdMap>, + /// Indices of the [`Binding::Tracking`] bindings, by the variable of their scope. + on_scope_var: DirectIdMap>, +} + +impl Bindings { + pub fn add(&mut self, binding: Binding) { + let index = self.bindings.len(); + self.bindings.push(binding); + + match binding { + Binding::Fixed { scope, trigger, .. } => { + assert!(scope != AriesLit::FALSE && trigger != AriesLit::FALSE); + // A tautological literal is entailed from the start and is never the subject of an + // event, so watching it would leave the binding unreachable. It is left unwatched, + // and `eval` checks entailment anyway. + for lit in [scope, trigger] { + if !lit.tautological() { + self.watches.add_watch(index, lit); + } + } + } + Binding::Tracking { scope, var, .. } => { + assert!(scope != AriesLit::FALSE); + assert!(var != AriesVar::ZERO); + Self::post(&mut self.on_var, var, index); + if !scope.tautological() { + Self::post(&mut self.on_scope_var, scope.variable(), index); + } + } + } + } + + fn post(map: &mut DirectIdMap>, var: AriesVar, index: usize) { + let var_idx = var.to_u32(); + if !map.contains_key(var_idx) { + map.insert(var_idx, Default::default()); + } + map[var_idx].push(index); + } + + /// The bindings that `lit` becoming entailed may have turned on: the ones watching it, and the + /// tracking ones on its variable, whether they mirror it or take it as scope. + /// + /// A binding may legitimately be yielded without anything having changed; applying a bound that + /// already holds is a no-op. + pub fn eval_on(&self, lit: AriesLit, dom: &impl Dom) -> impl Iterator { + let var_idx = lit.variable().to_u32(); + + let tracked = self.on_var.get(var_idx).into_iter().flatten().copied(); + let scoped = self.on_scope_var.get(var_idx).into_iter().flatten().copied(); + + self.watches + .watches_on(lit) + .chain(tracked) + .chain(scoped) + .flat_map(move |index| self.bindings[index].eval(dom)) + } + + /// Every binding that currently applies. + /// + /// Needed to seed the LP from the current domains: watches only fire on *future* events, so a + /// binding registered when its trigger is already entailed would otherwise never apply. + pub fn eval_all(&self, dom: &impl Dom) -> impl Iterator { + self.bindings.iter().flat_map(move |binding| binding.eval(dom)) + } +} diff --git a/solver/lp_highs/src/lib.rs b/solver/lp_highs/src/lib.rs new file mode 100644 index 00000000..00f38c6d --- /dev/null +++ b/solver/lp_highs/src/lib.rs @@ -0,0 +1,633 @@ +mod bindings; +mod state; +mod types; + +use aries_solver::{ + backtrack::{Backtrack, DecLvl, ObsTrailCursor}, + core::{ + literals::ConjunctionBuilder, + state::{Domains, DomainsSnapshot, Explanation, InferenceCause}, + }, + reasoners::{Contradiction, ReasonerId, Theory}, +}; + +use state::LpState; + +pub use types::*; + +use crate::bindings::{Binding, Bindings}; + +#[derive(Default, Clone)] +struct LpRelaxStats { + pub lpruns: u64, + pub lpruns_time: std::time::Duration, +} + +#[derive(Clone)] +pub struct LpOptions { + /// Used to activate/deactivate the propagation of the reasoner + pub propagation_active: bool, + /// If true, refined / minimized explanations will be computed. + pub is_explanation_refined: bool, +} +impl Default for LpOptions { + fn default() -> Self { + Self { + propagation_active: true, + is_explanation_refined: false, + } + } +} + +#[derive(Clone)] +pub struct LpRelax { + id: ReasonerId, + + model_events: ObsTrailCursor, + lp_state: LpState, + bindings: Bindings, + + stats: LpRelaxStats, + options: LpOptions, +} +unsafe impl Send for LpRelax {} +unsafe impl Sync for LpRelax {} + +impl Default for LpRelax { + fn default() -> Self { + Self { + id: ReasonerId::Extra(0), + model_events: Default::default(), + lp_state: Default::default(), + bindings: Default::default(), + stats: Default::default(), + options: Default::default(), + } + } +} +impl LpRelax { + pub fn with_options(options: LpOptions) -> Self { + Self { + options, + ..Default::default() + } + } + + pub fn activate_propagation(&mut self) { + self.options.propagation_active = true; + } + pub fn deactivate_propagation(&mut self) { + self.options.propagation_active = false; + } + + pub fn num_rows(&self) -> usize { + self.lp_state.num_rows() + } + pub fn num_columns(&self) -> usize { + self.lp_state.num_columns() + } + pub fn get_column_bounds(&self, col: LpCol) -> (LongCst, LongCst) { + self.lp_state.get_column_bounds(col) + } + + pub fn add_column_01(&mut self) -> LpCol { + assert!(self.lp_state.trail().trail.is_empty()); + self.add_column((Some(0), Some(1))) + } + pub fn add_column(&mut self, bounds: (Option, Option)) -> LpCol { + assert!(self.lp_state.trail().trail.is_empty()); + let bounds = (bounds.0.map(int_cst_as_long), bounds.1.map(int_cst_as_long)); + self.lp_state.add_column(bounds) + } + pub fn add_columns(&mut self, bounds: &[(Option, Option)]) -> Vec { + assert!(self.lp_state.trail().trail.is_empty()); + let bounds = bounds + .iter() + .map(|bounds| (bounds.0.map(int_cst_as_long), bounds.1.map(int_cst_as_long))); + self.lp_state.add_columns(bounds) + } + pub fn tighten_column(&mut self, col: LpCol, bounds: (Option, Option)) -> bool { + assert!(self.lp_state.trail().trail.is_empty()); + let bounds = (bounds.0.map(int_cst_as_long), bounds.1.map(int_cst_as_long)); + self.lp_state.tighten_column(col, bounds) + } + pub fn change_column(&mut self, col: LpCol, bounds: (Option, Option)) { + assert!(self.lp_state.trail().trail.is_empty()); + let bounds = (bounds.0.map(int_cst_as_long), bounds.1.map(int_cst_as_long)); + self.lp_state.change_column(col, bounds) + } + + pub fn add_row( + &mut self, + row_coefs: impl Iterator, + bounds: (Option, Option), + ) -> LpRow { + assert!(self.lp_state.trail().trail.is_empty()); + let row_coefs = row_coefs.map(|(col, c)| (col, int_cst_as_float(c))); + let bounds = (bounds.0.map(int_cst_as_float), bounds.1.map(int_cst_as_float)); + self.lp_state.add_row(row_coefs, bounds) + } + pub fn add_rows( + &mut self, + rows: impl Iterator, Option, impl Iterator)>, + ) -> Vec { + assert!(self.lp_state.trail().trail.is_empty()); + + self.lp_state.add_rows(rows.map(|(lb, ub, row)| { + ( + lb.map(int_cst_as_float), + ub.map(int_cst_as_float), + row.into_iter().map(|(col, c)| (col, int_cst_as_float(c))), + ) + })) + } + + pub fn add_objective_column( + &mut self, + main_var: AriesVar, + coefs: impl Iterator, + sense: LpObjectiveSense, + ) -> LpCol { + assert!(self.lp_state.trail().trail.is_empty()); + self.lp_state.add_objective_column(main_var, coefs, sense) + } + pub fn get_objective_column(&self) -> Option { + self.lp_state.get_objective_column() + } + pub fn get_objective_main_var(&self) -> Option { + self.lp_state.get_objective_main_var() + } + pub fn get_objective_sense(&self) -> Option { + self.lp_state.get_objective_sense() + } + + /// Applies `lp_lit` as soon as `scope` and `trigger` are both entailed. Either may be [`AriesLit::TRUE`]. + /// + /// Both are reported as the cause of the bound, so they must *entail* it. + pub fn half_bind_fixed(&mut self, scope: AriesLit, trigger: AriesLit, lp_lit: LpLit) { + assert!(self.lp_state.trail().trail.is_empty()); + self.bindings.add(Binding::Fixed { scope, trigger, lp_lit }); + } + + /// Makes `col` mirror the domain of `var` while `scope` is entailed. + /// + /// Unlike [`LpRelax::half_bind_fixed`], it constrains the column from both sides and is + /// re-evaluated on every event on `var` or on `scope`'s variable. + pub fn half_bind_tracking(&mut self, scope: AriesLit, var: AriesVar, col: LpCol) { + assert!(self.lp_state.trail().trail.is_empty()); + self.bindings.add(Binding::Tracking { scope, var, col }); + } + + fn process_model_events(&mut self, model: &mut Domains) -> Result<(), Contradiction> { + while let Some(main_event) = self.model_events.pop(model.trail()) { + // Ignore model events that originate from us (this reasoner), + // as they were already pushed to our (local) trail. + if let Some(x) = main_event.cause.as_external_inference() + && x.writer == self.identity() + { + continue; + } + + let derived = self + .bindings + .eval_on(main_event.new_literal(), model) + .collect::>(); + + for (lp_lit, scope, main_lit) in derived { + self.set_lp_lit( + lp_lit, + BoundCause::Some { + scope, + trigger: main_lit, + }, + )?; + } + } + Ok(()) + } + + /// Tightens a column bound. One crossing the opposite bound is a conflict, explained by the causes of both. + fn set_lp_lit(&mut self, lp_lit: LpLit, cause: BoundCause) -> Result<(), Contradiction> { + if let LpBoundUpdate::Emptied(cause) = self.lp_state.set_lp_lit(lp_lit, cause) { + // The bound it crosses is the one the column holds on the opposite side. + let opposite = match lp_lit.tpe { + LpLitType::GEQ => LpLitType::LEQ, + LpLitType::LEQ => LpLitType::GEQ, + }; + let mut expl = Explanation::new(); + cause.explain(&mut |l| expl.push(l)); + self.explain_column_bound(lp_lit.col, opposite, &mut |l| expl.push(l)); + Err(Contradiction::Explanation(expl)) + } else { + Ok(()) + } + } + + fn check_feasibility(&mut self) -> Result<(), Contradiction> { + match self.lp_state.solve_or_iis(&mut self.stats) { + Err(iis) => Err(self.build_contradiction(iis)), + _ => Ok(()), + } + } + + /// Pushes the main-model literals that justify the bound currently held by `col` on `bound`'s side. + fn explain_column_bound(&self, col: LpCol, bound: LpLitType, out: &mut impl FnMut(AriesLit)) { + self.lp_state.get_column_bound_cause(col, bound).explain(out); + } + + fn build_contradiction(&self, iis: LpIis) -> Contradiction { + let mut conjunction_builder = ConjunctionBuilder::new(); + let mut explained = 0usize; + + let mut explain_col = |col: LpCol, lower: bool, upper: bool| { + let mut push = |l: AriesLit| { + explained += 1; + conjunction_builder.push(l); + }; + if lower { + self.explain_column_bound(col, LpLitType::GEQ, &mut push); + } + if upper { + self.explain_column_bound(col, LpLitType::LEQ, &mut push); + } + }; + + for &(col, status) in iis.columns() { + match status { + highs::HighsIisBoundStatus::Lower => explain_col(col, true, false), + highs::HighsIisBoundStatus::Upper => explain_col(col, false, true), + highs::HighsIisBoundStatus::Boxed => explain_col(col, true, true), + highs::HighsIisBoundStatus::Free => (), + s => panic!("Unknown highs status {s:?}"), + } + } + + // An IIS that named nothing usable (e.g. HiGHS returned none) would leave an empty explanation, + // i.e. "infeasible whatever was decided". That's only true if no column has moved since the LP + // was built; otherwise fall back to every literal that moved one. + if explained == 0 && !self.lp_state.trail().trail.is_empty() { + for ev in &self.lp_state.trail().trail { + self.explain_column_bound(ev.col, ev.bound, &mut |l| conjunction_builder.push(l)); + } + } + + let mut expl = Explanation::new(); + expl.extend(conjunction_builder.build()); + Contradiction::Explanation(expl) + } +} + +impl Theory for LpRelax { + fn identity(&self) -> ReasonerId { + self.id + } + + fn propagate(&mut self, model: &mut Domains) -> Result<(), Contradiction> { + if self.lp_state.trail().trail.is_empty() { + let derived = self.bindings.eval_all(model).collect::>(); + for (lp_lit, scope, main_lit) in derived { + self.set_lp_lit( + lp_lit, + BoundCause::Some { + scope, + trigger: main_lit, + }, + )?; + } + } + + self.process_model_events(model)?; + + if !self.options.propagation_active { + return Ok(()); + } + + tracing::info!( + "|-[LPRELAX]- Solving LP at decision level {:?} (num events: {:?}) with HiGHS", + model.current_decision_level(), + model.num_events() + ); + self.check_feasibility() + } + + /// Should not be called: this reasoner never infers a literal, it only reports contradictions. + fn explain( + &mut self, + _literal: AriesLit, + _context: InferenceCause, + _model: &DomainsSnapshot, + _out_explanation: &mut Explanation, + ) { + unreachable!() + } + + fn print_stats(&self) { + println!("# lp runs: {}", self.stats.lpruns); + println!("# lp runs time: {:.6} s", self.stats.lpruns_time.as_secs_f64()); + } + + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl Backtrack for LpRelax { + fn save_state(&mut self) -> DecLvl { + self.lp_state.set_backtrack_point() + } + fn num_saved(&self) -> u32 { + self.lp_state.trail().num_saved() + } + fn restore_last(&mut self) { + self.lp_state.undo_to_last_backtrack_point(); + } +} + +#[cfg(test)] +pub mod test { + use aries_solver::backtrack::Backtrack; + use aries_solver::core::state::{Cause, Domains, Explanation}; + use aries_solver::core::views::Term; + use aries_solver::reasoners::{Contradiction, Theory}; + + use crate::LpRelax; + use crate::types::*; + + #[test] + fn test_trail_backtrack() { + let mut model = Domains::new(); + + let var2 = model.new_var(0, 10); + let var3 = model.new_var(0, 10); + + model.add_implication(var2.leq(5), var3.leq(5)); + + let mut theory = LpRelax::default(); + + let col2 = theory.add_column((Some(0), Some(10))); + let col3 = theory.add_column((Some(0), Some(10))); + + theory.half_bind_tracking(AriesLit::TRUE, var2.variable(), col2); + theory.half_bind_tracking(AriesLit::TRUE, var3.variable(), col3); + + let assert_col_bounds = |theory: &mut LpRelax, col: LpCol, col_bounds: (LongCst, LongCst)| { + assert_eq!(theory.get_column_bounds(col), col_bounds) + }; + + assert_col_bounds(&mut theory, col2, (0, 10)); + assert_col_bounds(&mut theory, col3, (0, 10)); + + model.save_state(); + theory.save_state(); + assert_eq!(model.set(var2.leq(8), Cause::Decision), Ok(true)); + assert!(theory.propagate(&mut model).is_ok()); + + assert_col_bounds(&mut theory, col2, (0, 8)); + assert_col_bounds(&mut theory, col3, (0, 10)); + + model.save_state(); + theory.save_state(); + assert_eq!(model.set(var3.leq(8), Cause::Decision), Ok(true)); + assert!(theory.propagate(&mut model).is_ok()); + + assert_col_bounds(&mut theory, col2, (0, 8)); + assert_col_bounds(&mut theory, col3, (0, 8)); + + model.restore_last(); + theory.restore_last(); + + assert_col_bounds(&mut theory, col2, (0, 8)); + assert_col_bounds(&mut theory, col3, (0, 10)); + + model.save_state(); + theory.save_state(); + + assert_eq!(model.set(var2.leq(5), Cause::Decision), Ok(true)); + assert!(theory.propagate(&mut model).is_ok()); + + assert_col_bounds(&mut theory, col2, (0, 5)); + assert_col_bounds(&mut theory, col3, (0, 5)); + + model.restore_last(); + theory.restore_last(); + + assert_col_bounds(&mut theory, col2, (0, 8)); + assert_col_bounds(&mut theory, col3, (0, 10)); + + model.restore_last(); + theory.restore_last(); + + assert_col_bounds(&mut theory, col2, (0, 10)); + assert_col_bounds(&mut theory, col3, (0, 10)); + } + + #[test] + fn test_infeas() { + let mut model = Domains::new(); + + let avar = model.new_var(0, 1); + let bvar = model.new_var(0, 1); + + let mut theory = LpRelax::default(); + + let acol = theory.add_column((Some(0), Some(1))); + let bcol = theory.add_column((Some(0), Some(1))); + + theory.half_bind_tracking(AriesLit::TRUE, avar.variable(), acol); + theory.half_bind_tracking(AriesLit::TRUE, bvar.variable(), bcol); + + theory.add_row([(acol, 1), (bcol, 1)].into_iter(), (Some(1), None)); + + let _ = model.set_ub(avar, 0, Cause::Decision).unwrap(); + let _ = model.set_ub(bvar, 0, Cause::Decision).unwrap(); + + let expl = match theory.propagate(&mut model) { + Err(Contradiction::Explanation(expl)) => expl, + _ => Explanation::new(), + }; + assert_eq!(expl.literals(), [avar.leq(0), bvar.leq(0)]); + } + + /// A fixed binding only applies within its scope, and stops applying again after a backtrack. + #[test] + fn test_binding_fixed_scope() { + let mut model = Domains::new(); + + let p = model.new_var(0, 1); + let q = model.new_var(0, 1); + let scope = model.new_var(0, 1); + + let mut theory = LpRelax::default(); + + let acol = theory.add_column((Some(0), Some(1))); + let bcol = theory.add_column((Some(0), Some(1))); + + theory.half_bind_fixed(AriesLit::TRUE, p.leq(0), LpLit::leq(acol, 0)); + theory.half_bind_fixed(scope.geq(1), q.leq(0), LpLit::leq(bcol, 0)); + + model.save_state(); + theory.save_state(); + + assert_eq!(model.set(p.leq(0), Cause::Decision), Ok(true)); + assert_eq!(model.set(q.leq(0), Cause::Decision), Ok(true)); + assert!(theory.propagate(&mut model).is_ok()); + + // `q <= 0` holds, but the binding that depends on it is still out of scope + assert_eq!(theory.get_column_bounds(acol), (0, 0)); + assert_eq!(theory.get_column_bounds(bcol), (0, 1)); + + model.save_state(); + theory.save_state(); + + assert_eq!(model.set(scope.geq(1), Cause::Decision), Ok(true)); + assert!(theory.propagate(&mut model).is_ok()); + + // entering the scope applies it, even though the event was on the scope rather than on `q` + assert_eq!(theory.get_column_bounds(bcol), (0, 0)); + + model.restore_last(); + theory.restore_last(); + + assert_eq!(theory.get_column_bounds(acol), (0, 0)); + assert_eq!(theory.get_column_bounds(bcol), (0, 1)); + } + + /// A conflict is explained by the triggers of the bindings involved, not by the bounds that the + /// domains happen to hold at that point. + #[test] + fn test_binding_fixed_explanation() { + let mut model = Domains::new(); + + let p = model.new_var(0, 10); + let scope = model.new_var(0, 1); + + let mut theory = LpRelax::default(); + + let acol = theory.add_column((Some(0), Some(1))); + let bcol = theory.add_column((Some(0), Some(1))); + + // both columns are pinned to 0 as soon as `p` is above 4, which the row forbids + theory.half_bind_fixed(AriesLit::TRUE, p.geq(5), LpLit::leq(acol, 0)); + theory.half_bind_fixed(scope.geq(1), p.geq(5), LpLit::leq(bcol, 0)); + theory.add_row([(acol, 1), (bcol, 1)].into_iter(), (Some(1), None)); + + assert_eq!(model.set(scope.geq(1), Cause::Decision), Ok(true)); + assert_eq!(model.set(p.geq(9), Cause::Decision), Ok(true)); + + let expl = match theory.propagate(&mut model) { + Err(Contradiction::Explanation(expl)) => expl, + _ => Explanation::new(), + }; + + // `p >= 5` is what the bindings asked for; `p >= 9` is what the domain holds + assert!(expl.literals().contains(&p.geq(5)), "{:?}", expl.literals()); + assert!(!expl.literals().contains(&p.geq(9)), "{:?}", expl.literals()); + assert!(expl.literals().contains(&scope.geq(1)), "{:?}", expl.literals()); + } + + /// Two bindings that empty a column between them conflict when the second is applied, + /// without the LP ever being solved, and are explained by the triggers of both. + #[test] + fn test_crossing_bounds_conflict() { + let mut model = Domains::new(); + + let p = model.new_var(0, 1); + let q = model.new_var(0, 1); + + let mut theory = LpRelax::default(); + let col = theory.add_column((Some(0), Some(1))); + theory.half_bind_fixed(AriesLit::TRUE, p.leq(0), LpLit::leq(col, 0)); + theory.half_bind_fixed(AriesLit::TRUE, q.leq(0), LpLit::geq(col, 1)); + + assert_eq!(model.set(p.leq(0), Cause::Decision), Ok(true)); + assert_eq!(model.set(q.leq(0), Cause::Decision), Ok(true)); + + let Err(Contradiction::Explanation(expl)) = theory.propagate(&mut model) else { + panic!("the column cannot be both <= 0 and >= 1") + }; + assert!(expl.literals().contains(&p.leq(0)), "{:?}", expl.literals()); + assert!(expl.literals().contains(&q.leq(0)), "{:?}", expl.literals()); + } + + /// Backtracking one level restores the bound the column had at that level, + /// not the one it was created with. + #[test] + fn test_bound_restored_to_intermediate_value() { + let mut model = Domains::new(); + + let p = model.new_var(0, 1); + let q = model.new_var(0, 1); + + let mut theory = LpRelax::default(); + let col = theory.add_column((Some(0), Some(10))); + theory.half_bind_fixed(AriesLit::TRUE, p.leq(0), LpLit::leq(col, 6)); + theory.half_bind_fixed(AriesLit::TRUE, q.leq(0), LpLit::leq(col, 3)); + + model.save_state(); + theory.save_state(); + assert_eq!(model.set(p.leq(0), Cause::Decision), Ok(true)); + assert!(theory.propagate(&mut model).is_ok()); + assert_eq!(theory.get_column_bounds(col), (0, 6)); + + model.save_state(); + theory.save_state(); + assert_eq!(model.set(q.leq(0), Cause::Decision), Ok(true)); + assert!(theory.propagate(&mut model).is_ok()); + assert_eq!(theory.get_column_bounds(col), (0, 3)); + + model.restore_last(); + theory.restore_last(); + assert_eq!(theory.get_column_bounds(col), (0, 6)); + + model.restore_last(); + theory.restore_last(); + assert_eq!(theory.get_column_bounds(col), (0, 10)); + } + + /// An empty row with a positive lower bound cannot be satisfied on its own, + /// but is only considered infeasible by HiGHS once the LP has a column. + /// Without any columns, it reports the problem as feasible, even though the row is inconsistent on its own. + #[test] + fn test_empty_row_with_positive_lower_bound() { + let mut model = Domains::new(); + + let mut theory = LpRelax::default(); + theory.add_row(std::iter::empty::<(LpCol, IntCst)>(), (Some(1), None)); + assert_eq!(theory.num_columns(), 0); + assert!(theory.propagate(&mut model).is_ok()); + + let mut theory = LpRelax::default(); + theory.add_column((Some(0), Some(1))); + theory.add_row(std::iter::empty::<(LpCol, IntCst)>(), (Some(1), None)); + + let Err(Contradiction::Explanation(expl)) = theory.propagate(&mut model) else { + panic!("a row demanding `0 >= 1` is infeasible") + }; + assert!(expl.literals().is_empty(), "{:?}", expl.literals()); + + let mut theory = LpRelax::default(); + let col = theory.add_column((Some(0), Some(1))); + theory.add_rows(std::iter::once((Some(1), Some(0), std::iter::once((col, 1))))); + + let Err(Contradiction::Explanation(expl)) = theory.propagate(&mut model) else { + panic!("a row bounded [1, 0] cannot be satisfied") + }; + assert!(expl.literals().is_empty(), "{:?}", expl.literals()); + } + + /// An LP that is infeasible from the bounds its columns were created with is infeasible whatever the search does, so its explanation is empty. + #[test] + fn test_root_infeasibility_is_explained_by_nothing() { + let mut model = Domains::new(); + + let mut theory = LpRelax::default(); + let acol = theory.add_column((Some(0), Some(0))); + let bcol = theory.add_column((Some(0), Some(0))); + theory.add_row([(acol, 1), (bcol, 1)].into_iter(), (Some(1), None)); + + let Err(Contradiction::Explanation(expl)) = theory.propagate(&mut model) else { + panic!("`a + b >= 1` cannot hold with both columns pinned to 0") + }; + assert!(expl.literals().is_empty(), "{:?}", expl.literals()); + } +} diff --git a/solver/lp_highs/src/state.rs b/solver/lp_highs/src/state.rs new file mode 100644 index 00000000..74999c04 --- /dev/null +++ b/solver/lp_highs/src/state.rs @@ -0,0 +1,334 @@ +use aries_solver::backtrack::{DecLvl, Trail}; + +use crate::types::*; + +#[derive(Debug, Clone)] +struct LpObjective { + pub col: LpCol, + pub main_var: AriesVar, + pub sense: LpObjectiveSense, +} + +/// Bounds of a column and the cause justifying each of them. +/// +/// The LP model holds the same bounds as floats. +/// This is also where the causes needed by the explanations are recorded. +#[derive(Debug, Clone)] +struct ColBounds { + lower: LongCst, + lower_cause: BoundCause, + upper: LongCst, + upper_cause: BoundCause, +} + +pub(super) struct LpState { + lp_trail: Trail, + col_bounds: Vec, + + lp_model: LpModel, + lp_obj: Option, +} + +impl Clone for LpState { + fn clone(&self) -> Self { + let mut lp_model = self.lp_model.clone(); + set_lp_model_options(&mut lp_model); + + Self { + lp_trail: self.lp_trail.clone(), + col_bounds: self.col_bounds.clone(), + lp_model, + lp_obj: self.lp_obj.clone(), + } + } +} +impl Default for LpState { + fn default() -> Self { + let mut lp_model = highs::ColProblem::default().optimise(LpObjectiveSense::Minimise); + set_lp_model_options(&mut lp_model); + + Self { + lp_trail: Default::default(), + col_bounds: Default::default(), + lp_model, + lp_obj: None, + } + } +} +fn set_lp_model_options(lp_model: &mut LpModel) { + //lp_model.set_option("time_limit", 2.0); // stop after 2 seconds + lp_model.set_option("parallel", "off"); // use 1 core + lp_model.set_option("threads", 1); // solve on 1 thread + lp_model.set_option("iis_strategy", 0); // https://github.com/ERGO-Code/HiGHS/blob/3be639f037e0001b617c59830d3965f246ab5beb/highs/interfaces/highs_c_api.h#L153 +} + +impl LpState { + pub fn trail(&self) -> &Trail { + &self.lp_trail + } + + pub fn num_rows(&self) -> usize { + self.lp_model.num_rows() + } + pub fn num_columns(&self) -> usize { + self.lp_model.num_cols() + } + + pub fn get_column_bounds(&self, col: LpCol) -> (LongCst, LongCst) { + let bounds = &self.col_bounds[col.index()]; + (bounds.lower, bounds.upper) + } + + /// Cause justifying one of the bounds of a column. + pub fn get_column_bound_cause(&self, col: LpCol, bound: LpLitType) -> BoundCause { + let bounds = &self.col_bounds[col.index()]; + match bound { + LpLitType::GEQ => bounds.lower_cause, + LpLitType::LEQ => bounds.upper_cause, + } + } + + /// Mirrors the bounds recorded for a column into the LP model. + fn update_model_with_column_bounds(&mut self, col: LpCol) { + let (lower, upper) = self.get_column_bounds(col); + self.lp_model + .change_column_bounds(col, long_cst_as_float(lower)..=long_cst_as_float(upper)); + } + + pub fn add_column(&mut self, bounds: (Option, Option)) -> LpCol { + self.add_columns(std::iter::once(bounds))[0] + } + + pub fn add_columns(&mut self, bounds: impl Iterator, Option)>) -> Vec { + let bounds = bounds + .map(|(lb, ub)| { + // An unbounded side is held as the extremum of the type, which converts to a + // magnitude HiGHS treats as infinite. + let (lower, upper) = (lb.unwrap_or(LongCst::MIN), ub.unwrap_or(LongCst::MAX)); + assert!(lower <= upper); + self.col_bounds.push(ColBounds { + lower, + lower_cause: BoundCause::None, + upper, + upper_cause: BoundCause::None, + }); + long_cst_as_float(lower)..=long_cst_as_float(upper) + }) + .collect::>(); + + let old_cols_num = self.lp_model.num_cols(); + self.lp_model.add_columns(bounds); + + debug_assert_eq!(self.col_bounds.len(), self.lp_model.num_cols()); + (old_cols_num..self.num_columns()).map(LpCol::from).collect() + } + + /// Restricts the bounds of a column unconditionally, i.e. with no cause to justify them. + /// Only meaningful before the search starts. Returns whether either bound was restricted. + pub fn tighten_column(&mut self, col: LpCol, bounds: (Option, Option)) -> bool { + debug_assert!(self.lp_trail.trail.is_empty()); + let (old_lower, old_upper) = self.get_column_bounds(col); + + let lower = bounds.0.filter(|lb| *lb > old_lower); + let upper = bounds.1.filter(|ub| *ub < old_upper); + + let entry = &mut self.col_bounds[col.index()]; + if let Some(lower) = lower { + entry.lower = lower; + } + if let Some(upper) = upper { + entry.upper = upper; + } + assert!(entry.lower <= entry.upper); + + self.update_model_with_column_bounds(col); + + lower.is_some() || upper.is_some() + } + + /// Overwrites the bounds of a column unconditionally. Only meaningful before the search starts. + pub fn change_column(&mut self, col: LpCol, bounds: (Option, Option)) { + debug_assert!(self.lp_trail.trail.is_empty()); + let entry = &mut self.col_bounds[col.index()]; + entry.lower = bounds.0.unwrap_or(LongCst::MIN); + entry.upper = bounds.1.unwrap_or(LongCst::MAX); + assert!(entry.lower <= entry.upper); + + self.update_model_with_column_bounds(col); + } + + pub fn add_row( + &mut self, + row_coefs: impl Iterator, + bounds: (Option, Option), + ) -> LpRow { + let lb = bounds.0.unwrap_or(FloatCst::MIN); + let ub = bounds.1.unwrap_or(FloatCst::MAX); + debug_assert!(lb <= ub, "row with inconsistent bounds [{lb}, {ub}]"); + + if lb > ub { + // HiGHS segfaults in its IIS routine on a row with `lb > ub` (see `add_rows`). + // An empty row with a positive lower bound states the same unconditional infeasibility and is handled safely. + return self.lp_model.add_row(1.0..FloatCst::MAX, std::iter::empty()); + } + + let num_columns = self.num_columns(); + self.lp_model.add_row( + lb..ub, + row_coefs.inspect(|(col, _)| debug_assert!(col.index() < num_columns)), + ) + } + + pub fn add_rows( + &mut self, + rows: impl Iterator< + Item = ( + Option, + Option, + impl Iterator, + ), + >, + ) -> Vec { + let num_cols = self.num_columns(); + + let rows = rows.map(move |(lb, ub, row)| { + let lb = lb.unwrap_or(FloatCst::MIN); + let ub = ub.unwrap_or(FloatCst::MAX); + + // HiGHS crashes on a row with `lb > ub`: it reports infeasibility from the bound check before + // making the matrix column-wise, and its IIS routine then reads out of bounds while building + // the 0-column IIS. An empty row with a positive lower bound is the same statement + // ("infeasible regardless of any column") in a shape HiGHS survives. + let consistent = lb <= ub; + let bounds = if consistent { lb..=ub } else { 1.0..=FloatCst::MAX }; + + // row ends up empty if not `consistent` (everything would be filtered out) + let row = row + .inspect(move |(col, _)| debug_assert!(col.index() < num_cols)) + .filter(move |_| consistent); + + (bounds, row) + }); + + let old_num_rows = self.num_rows(); + self.lp_model.add_rows(rows); + (old_num_rows..self.num_rows()).map(LpRow::from).collect() + } + + pub fn add_objective_column( + &mut self, + main_var: AriesVar, + coefs: impl Iterator, + sense: LpObjectiveSense, + ) -> LpCol { + assert!(self.lp_obj.is_none()); + + self.lp_obj = Some(LpObjective { + col: self.add_column((None, None)), + main_var, + sense, + }); + self.lp_model + .change_column_cost(self.get_objective_column().unwrap(), 1.); + + let factors = coefs + .into_iter() + .chain([(self.get_objective_column().unwrap(), -1.)]) + .collect::>(); + + self.add_row(factors.into_iter(), (Some(0.), Some(0.))); + self.lp_model.set_sense(sense); + + self.get_objective_column().unwrap() + } + pub fn get_objective_column(&self) -> Option { + self.lp_obj.as_ref().map(|lp_obj| lp_obj.col) + } + pub fn get_objective_main_var(&self) -> Option { + self.lp_obj.as_ref().map(|lp_obj| lp_obj.main_var) + } + pub fn get_objective_sense(&self) -> Option { + self.lp_obj.as_ref().map(|lp_obj| lp_obj.sense) + } + + /// Restricts one bound of a column, with the cause justifying it, and records what it takes to undo that on backtrack. + /// + /// A bound that is already held is left untouched, and one crossing the opposite bound is not applied at all (as that is a contradiction): + /// the caller is handed the cause back to explain this contradiction using it. + pub fn set_lp_lit(&mut self, lp_lit: LpLit, cause: BoundCause) -> LpBoundUpdate { + let col = lp_lit.col; + let (lower, upper) = self.get_column_bounds(col); + + let held = match lp_lit.tpe { + LpLitType::GEQ => LpLit::geq(col, lower), + LpLitType::LEQ => LpLit::leq(col, upper), + }; + if !lp_lit.strictly_entails(held) { + return LpBoundUpdate::Unchanged; + } + let (new_lower, new_upper) = match lp_lit.tpe { + LpLitType::GEQ => (lp_lit.val, upper), + LpLitType::LEQ => (lower, lp_lit.val), + }; + if new_lower > new_upper { + return LpBoundUpdate::Emptied(cause); + } + + let entry = &mut self.col_bounds[col.index()]; + let old_cause = match lp_lit.tpe { + LpLitType::GEQ => std::mem::replace(&mut entry.lower_cause, cause), + LpLitType::LEQ => std::mem::replace(&mut entry.upper_cause, cause), + }; + entry.lower = new_lower; + entry.upper = new_upper; + + self.lp_trail.push(LpEvent { + col, + bound: lp_lit.tpe, + old_val: held.val, + old_cause, + }); + self.update_model_with_column_bounds(col); + + LpBoundUpdate::Tightened + } + + pub fn set_backtrack_point(&mut self) -> DecLvl { + self.lp_trail.save_state() + } + + pub fn undo_to_last_backtrack_point(&mut self) { + self.lp_model.clear_solver(); + + let (col_bounds, lp_model) = (&mut self.col_bounds, &mut self.lp_model); + self.lp_trail.restore_last_with(|ev| { + let entry = &mut col_bounds[ev.col.index()]; + match ev.bound { + LpLitType::GEQ => { + entry.lower = ev.old_val; + entry.lower_cause = ev.old_cause; + } + LpLitType::LEQ => { + entry.upper = ev.old_val; + entry.upper_cause = ev.old_cause; + } + } + lp_model.change_column_bounds(ev.col, long_cst_as_float(entry.lower)..=long_cst_as_float(entry.upper)); + }); + } + + pub fn solve_or_iis(&mut self, stats: &mut crate::LpRelaxStats) -> Result { + let time = std::time::Instant::now(); + + let res = self.lp_model.solve_or_iis(); + + if res.is_err() { + self.lp_model.clear_solver(); + } + + stats.lpruns_time += time.elapsed(); + stats.lpruns += 1; + + res + } +} diff --git a/solver/lp_highs/src/types.rs b/solver/lp_highs/src/types.rs new file mode 100644 index 00000000..90f7b921 --- /dev/null +++ b/solver/lp_highs/src/types.rs @@ -0,0 +1,132 @@ +pub type AriesSignedVar = aries_solver::prelude::SignedVar; +pub type AriesVar = aries_solver::prelude::Var; +pub type AriesVarIdx = u32; +pub type AriesLit = aries_solver::prelude::Lit; +pub type AriesModelEvent = aries_solver::core::state::Event; + +pub use aries_solver::core::LongCst; +pub use aries_solver::prelude::{INT_CST_MAX, INT_CST_MIN, IntCst}; + +pub type LpCol = highs::Col; +pub type LpRow = highs::Row; +pub type LpSolution = highs::Solution; +pub type LpIis = highs::Iis; +pub type LpModel = highs::Model; +pub type LpObjectiveSense = highs::Sense; + +pub type FloatCst = f64; + +pub fn long_cst_as_float(value: LongCst) -> FloatCst { + // [`LongCst::MIN`] and [`LongCst::MAX`] stand for an unbounded column + match value { + LongCst::MIN => FloatCst::NEG_INFINITY, + LongCst::MAX => FloatCst::INFINITY, + value => value as FloatCst, + } +} + +pub fn int_cst_as_float(value: IntCst) -> FloatCst { + value as FloatCst +} + +/// Widening of a bound of the main model into the type column bounds are held in. +pub fn int_cst_as_long(value: IntCst) -> LongCst { + value as LongCst +} + +/// Analogous to a literal in the main aries solver, but on a column of the LP. +#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)] +pub struct LpLit { + pub col: LpCol, + pub tpe: LpLitType, + pub val: LongCst, +} + +#[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)] +pub enum LpLitType { + GEQ, + LEQ, +} + +impl LpLit { + pub fn new(col: LpCol, tpe: LpLitType, val: LongCst) -> Self { + Self { col, tpe, val } + } + pub fn leq(col: LpCol, val: LongCst) -> Self { + Self { + col, + tpe: LpLitType::LEQ, + val, + } + } + pub fn geq(col: LpCol, val: LongCst) -> Self { + Self { + col, + tpe: LpLitType::GEQ, + val, + } + } + pub fn entails(&self, other: Self) -> bool { + if self.tpe == other.tpe { + match self.tpe { + LpLitType::GEQ => self.val >= other.val, + LpLitType::LEQ => self.val <= other.val, + } + } else { + false + } + } + pub fn strictly_entails(&self, other: Self) -> bool { + self.entails(other) && self.val != other.val + } +} + +/// Store all the necessary information for backtracking after modifying a bound +/// +/// Only the old value and cause are necessary as they will overwrite the current ones +#[derive(Debug, Clone)] +pub(super) struct LpEvent { + /// Column affected by the bound change + pub col: LpCol, + /// Bound modified + pub bound: LpLitType, + /// Value the bound had before this change, restored when backtracking + pub old_val: LongCst, + /// Cause that justified `old_val`, restored along with it + pub old_cause: BoundCause, +} + +/// Justification of a bound of an LP column: the reason why that bound was set. +/// In other words, it corresponds to a sufficient condition for the bound to hold. +/// +/// - `Some { scope, trigger }`: the bound holds as long as **both** the scope and trigger literal are entailed in the main model. +/// `scope` is [`AriesLit::TRUE`] for the bounds that are not guarded by a scope, which is the common case. +/// - `None`: the bound holds unconditionally (initial bound of a column, or bound entailed at the root) +/// (it is functionally equivalent to `Some { scope: AriesLit::TRUE, trigger: AriesLit::TRUE }`) +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(super) enum BoundCause { + Some { scope: AriesLit, trigger: AriesLit }, + None, +} + +impl BoundCause { + /// Outputs the literals that must be entailed for a bound with this cause to hold. + /// An unconditional bound outputs nothing, and neither does a tautological literal. + pub(super) fn explain(&self, out: &mut impl FnMut(AriesLit)) { + if let BoundCause::Some { scope, trigger } = self { + if !scope.tautological() { + out(*scope); + } + if !trigger.tautological() { + out(*trigger); + } + } + } +} + +pub(super) enum LpBoundUpdate { + Unchanged, + Tightened, + /// Inconsistent update (lower bound higher than upper bound). Cause is handed back to explain the conflict. + Emptied(BoundCause), +} From 1d5b1a0b93db904fca6c651ca8fef64c8e469682 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 12:33:08 +0200 Subject: [PATCH 04/19] deps(lp-highs): depend on the highs bindings fork --- Cargo.lock | 35 +++++++++++++++++++++++++++++++++-- 1 file changed, 33 insertions(+), 2 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 1d38acac..87b72d86 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -206,6 +206,19 @@ dependencies = [ "tracing", ] +[[package]] +name = "aries-solver-lp-highs" +version = "0.1.0" +dependencies = [ + "aries-env-param", + "aries-solver", + "highs 2.4.0 (git+https://github.com/nrealus/highs.git?branch=dev)", + "idmap", + "itertools 0.14.0", + "smallvec", + "tracing", +] + [[package]] name = "aries-timelines" version = "0.1.0" @@ -1191,7 +1204,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c4255eb1d65eea322f73a326038f2bcec19dae13ca50fa557da19b03274b4d56" dependencies = [ "fnv", - "highs", + "highs 2.4.0 (registry+https://github.com/rust-lang/crates.io-index)", ] [[package]] @@ -1309,7 +1322,16 @@ version = "2.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7dfc791427e78fdc1a0a37983f79c20f86a0b7b7dfdb5d6a57a6f72d77cd9074" dependencies = [ - "highs-sys", + "highs-sys 1.15.0 (registry+https://github.com/rust-lang/crates.io-index)", + "log", +] + +[[package]] +name = "highs" +version = "2.4.0" +source = "git+https://github.com/nrealus/highs.git?branch=dev#680b09e7b8df6e0a22cdaa62dabf6ef106eed447" +dependencies = [ + "highs-sys 1.15.0 (git+https://github.com/nrealus/highs-sys.git?branch=dev)", "log", ] @@ -1323,6 +1345,15 @@ dependencies = [ "cmake", ] +[[package]] +name = "highs-sys" +version = "1.15.0" +source = "git+https://github.com/nrealus/highs-sys.git?branch=dev#c8cd96bdce362d9427ffe8e56c471700da94a3ef" +dependencies = [ + "bindgen", + "cmake", +] + [[package]] name = "http" version = "1.4.2" From de6f75e14bad24a159272b87e545084c78fb4946 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 15:24:33 +0200 Subject: [PATCH 05/19] chore(lp-highs): improve management of highs options --- solver/lp_highs/src/lib.rs | 39 ++++++++++++++++++++--- solver/lp_highs/src/state.rs | 60 ++++++++++++++++++++++-------------- 2 files changed, 72 insertions(+), 27 deletions(-) diff --git a/solver/lp_highs/src/lib.rs b/solver/lp_highs/src/lib.rs index 00f38c6d..b552c36d 100644 --- a/solver/lp_highs/src/lib.rs +++ b/solver/lp_highs/src/lib.rs @@ -15,7 +15,10 @@ use state::LpState; pub use types::*; -use crate::bindings::{Binding, Bindings}; +use crate::{ + bindings::{Binding, Bindings}, + state::HighsOptionValueWrapper, +}; #[derive(Default, Clone)] struct LpRelaxStats { @@ -38,8 +41,21 @@ impl Default for LpOptions { } } } +impl LpOptions { + fn get_highs_options(&self) -> impl Iterator>, HighsOptionValueWrapper)> { + [ + // ("time_limit", HighsOptionValueWrapper::Float(10.0)), + ("parallel", HighsOptionValueWrapper::Str("off")), // use 1 core + ("threads", HighsOptionValueWrapper::Int(1)), // solve on 1 thread + ( + "iis_strategy", + HighsOptionValueWrapper::Int(if self.is_explanation_refined { 4 } else { 0 }), + ), // https://github.com/ERGO-Code/HiGHS/blob/3be639f037e0001b617c59830d3965f246ab5beb/highs/interfaces/highs_c_api.h#L153 + ] + .into_iter() + } +} -#[derive(Clone)] pub struct LpRelax { id: ReasonerId, @@ -53,18 +69,33 @@ pub struct LpRelax { unsafe impl Send for LpRelax {} unsafe impl Sync for LpRelax {} +impl Clone for LpRelax { + fn clone(&self) -> Self { + let options = LpOptions::default(); + Self { + id: self.id, + model_events: self.model_events.clone(), + lp_state: self.lp_state.clone_with_options(options.get_highs_options()), + bindings: self.bindings.clone(), + stats: self.stats.clone(), + options, + } + } +} impl Default for LpRelax { fn default() -> Self { + let options = LpOptions::default(); Self { id: ReasonerId::Extra(0), model_events: Default::default(), - lp_state: Default::default(), + lp_state: LpState::default_with_options(options.get_highs_options()), bindings: Default::default(), stats: Default::default(), - options: Default::default(), + options, } } } + impl LpRelax { pub fn with_options(options: LpOptions) -> Self { Self { diff --git a/solver/lp_highs/src/state.rs b/solver/lp_highs/src/state.rs index 74999c04..a819c522 100644 --- a/solver/lp_highs/src/state.rs +++ b/solver/lp_highs/src/state.rs @@ -29,40 +29,54 @@ pub(super) struct LpState { lp_obj: Option, } -impl Clone for LpState { - fn clone(&self) -> Self { - let mut lp_model = self.lp_model.clone(); - set_lp_model_options(&mut lp_model); +#[allow(dead_code)] +pub(super) enum HighsOptionValueWrapper { + Str(&'static str), + Int(i32), + Float(f64), + Bool(bool), +} +impl LpState { + pub fn default_with_options(options: impl Iterator>, HighsOptionValueWrapper)>) -> Self { + let mut lp_model = highs::ColProblem::default().optimise(LpObjectiveSense::Minimise); + for (k, v) in options { + match v { + HighsOptionValueWrapper::Str(v) => lp_model.set_option(k, v), + HighsOptionValueWrapper::Int(v) => lp_model.set_option(k, v), + HighsOptionValueWrapper::Float(v) => lp_model.set_option(k, v), + HighsOptionValueWrapper::Bool(v) => lp_model.set_option(k, v), + } + } Self { - lp_trail: self.lp_trail.clone(), - col_bounds: self.col_bounds.clone(), + lp_trail: Default::default(), + col_bounds: Default::default(), lp_model, - lp_obj: self.lp_obj.clone(), + lp_obj: None, } } -} -impl Default for LpState { - fn default() -> Self { - let mut lp_model = highs::ColProblem::default().optimise(LpObjectiveSense::Minimise); - set_lp_model_options(&mut lp_model); + pub fn clone_with_options( + &self, + options: impl Iterator>, HighsOptionValueWrapper)>, + ) -> Self { + let mut lp_model = self.lp_model.clone(); + for (k, v) in options { + match v { + HighsOptionValueWrapper::Str(v) => lp_model.set_option(k, v), + HighsOptionValueWrapper::Int(v) => lp_model.set_option(k, v), + HighsOptionValueWrapper::Float(v) => lp_model.set_option(k, v), + HighsOptionValueWrapper::Bool(v) => lp_model.set_option(k, v), + } + } Self { - lp_trail: Default::default(), - col_bounds: Default::default(), + lp_trail: self.lp_trail.clone(), + col_bounds: self.col_bounds.clone(), lp_model, - lp_obj: None, + lp_obj: self.lp_obj.clone(), } } -} -fn set_lp_model_options(lp_model: &mut LpModel) { - //lp_model.set_option("time_limit", 2.0); // stop after 2 seconds - lp_model.set_option("parallel", "off"); // use 1 core - lp_model.set_option("threads", 1); // solve on 1 thread - lp_model.set_option("iis_strategy", 0); // https://github.com/ERGO-Code/HiGHS/blob/3be639f037e0001b617c59830d3965f246ab5beb/highs/interfaces/highs_c_api.h#L153 -} -impl LpState { pub fn trail(&self) -> &Trail { &self.lp_trail } From 4b9b73234e50ec2a4106c65777eae2825b0a14ba Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 15:26:53 +0200 Subject: [PATCH 06/19] chore(lp-highs): rename `LpRelax` to `Lp` to mirror the name of the existing `Lp` reasoner based on minilp --- solver/lp_highs/src/bindings.rs | 2 +- solver/lp_highs/src/lib.rs | 46 ++++++++++++++++----------------- solver/lp_highs/src/state.rs | 2 +- 3 files changed, 25 insertions(+), 25 deletions(-) diff --git a/solver/lp_highs/src/bindings.rs b/solver/lp_highs/src/bindings.rs index af40c5ca..5aed03c5 100644 --- a/solver/lp_highs/src/bindings.rs +++ b/solver/lp_highs/src/bindings.rs @@ -63,7 +63,7 @@ impl Binding { } } -/// The bindings of an [`crate::LpRelax`] reasoner, indexed by what can make them apply. +/// The bindings of an [`crate::Lp`] reasoner, indexed by what can make them apply. #[derive(Default, Clone)] pub(super) struct Bindings { bindings: Vec, diff --git a/solver/lp_highs/src/lib.rs b/solver/lp_highs/src/lib.rs index b552c36d..9b886bbb 100644 --- a/solver/lp_highs/src/lib.rs +++ b/solver/lp_highs/src/lib.rs @@ -21,7 +21,7 @@ use crate::{ }; #[derive(Default, Clone)] -struct LpRelaxStats { +struct LpStats { pub lpruns: u64, pub lpruns_time: std::time::Duration, } @@ -56,20 +56,20 @@ impl LpOptions { } } -pub struct LpRelax { +pub struct Lp { id: ReasonerId, model_events: ObsTrailCursor, lp_state: LpState, bindings: Bindings, - stats: LpRelaxStats, + stats: LpStats, options: LpOptions, } -unsafe impl Send for LpRelax {} -unsafe impl Sync for LpRelax {} +unsafe impl Send for Lp {} +unsafe impl Sync for Lp {} -impl Clone for LpRelax { +impl Clone for Lp { fn clone(&self) -> Self { let options = LpOptions::default(); Self { @@ -82,7 +82,7 @@ impl Clone for LpRelax { } } } -impl Default for LpRelax { +impl Default for Lp { fn default() -> Self { let options = LpOptions::default(); Self { @@ -96,7 +96,7 @@ impl Default for LpRelax { } } -impl LpRelax { +impl Lp { pub fn with_options(options: LpOptions) -> Self { Self { options, @@ -202,7 +202,7 @@ impl LpRelax { /// Makes `col` mirror the domain of `var` while `scope` is entailed. /// - /// Unlike [`LpRelax::half_bind_fixed`], it constrains the column from both sides and is + /// Unlike [`Lp::half_bind_fixed`], it constrains the column from both sides and is /// re-evaluated on every event on `var` or on `scope`'s variable. pub fn half_bind_tracking(&mut self, scope: AriesLit, var: AriesVar, col: LpCol) { assert!(self.lp_state.trail().trail.is_empty()); @@ -308,7 +308,7 @@ impl LpRelax { } } -impl Theory for LpRelax { +impl Theory for Lp { fn identity(&self) -> ReasonerId { self.id } @@ -362,7 +362,7 @@ impl Theory for LpRelax { } } -impl Backtrack for LpRelax { +impl Backtrack for Lp { fn save_state(&mut self) -> DecLvl { self.lp_state.set_backtrack_point() } @@ -381,7 +381,7 @@ pub mod test { use aries_solver::core::views::Term; use aries_solver::reasoners::{Contradiction, Theory}; - use crate::LpRelax; + use crate::Lp; use crate::types::*; #[test] @@ -393,7 +393,7 @@ pub mod test { model.add_implication(var2.leq(5), var3.leq(5)); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); let col2 = theory.add_column((Some(0), Some(10))); let col3 = theory.add_column((Some(0), Some(10))); @@ -401,7 +401,7 @@ pub mod test { theory.half_bind_tracking(AriesLit::TRUE, var2.variable(), col2); theory.half_bind_tracking(AriesLit::TRUE, var3.variable(), col3); - let assert_col_bounds = |theory: &mut LpRelax, col: LpCol, col_bounds: (LongCst, LongCst)| { + let assert_col_bounds = |theory: &mut Lp, col: LpCol, col_bounds: (LongCst, LongCst)| { assert_eq!(theory.get_column_bounds(col), col_bounds) }; @@ -459,7 +459,7 @@ pub mod test { let avar = model.new_var(0, 1); let bvar = model.new_var(0, 1); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); let acol = theory.add_column((Some(0), Some(1))); let bcol = theory.add_column((Some(0), Some(1))); @@ -488,7 +488,7 @@ pub mod test { let q = model.new_var(0, 1); let scope = model.new_var(0, 1); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); let acol = theory.add_column((Some(0), Some(1))); let bcol = theory.add_column((Some(0), Some(1))); @@ -532,7 +532,7 @@ pub mod test { let p = model.new_var(0, 10); let scope = model.new_var(0, 1); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); let acol = theory.add_column((Some(0), Some(1))); let bcol = theory.add_column((Some(0), Some(1))); @@ -565,7 +565,7 @@ pub mod test { let p = model.new_var(0, 1); let q = model.new_var(0, 1); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); let col = theory.add_column((Some(0), Some(1))); theory.half_bind_fixed(AriesLit::TRUE, p.leq(0), LpLit::leq(col, 0)); theory.half_bind_fixed(AriesLit::TRUE, q.leq(0), LpLit::geq(col, 1)); @@ -589,7 +589,7 @@ pub mod test { let p = model.new_var(0, 1); let q = model.new_var(0, 1); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); let col = theory.add_column((Some(0), Some(10))); theory.half_bind_fixed(AriesLit::TRUE, p.leq(0), LpLit::leq(col, 6)); theory.half_bind_fixed(AriesLit::TRUE, q.leq(0), LpLit::leq(col, 3)); @@ -622,12 +622,12 @@ pub mod test { fn test_empty_row_with_positive_lower_bound() { let mut model = Domains::new(); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); theory.add_row(std::iter::empty::<(LpCol, IntCst)>(), (Some(1), None)); assert_eq!(theory.num_columns(), 0); assert!(theory.propagate(&mut model).is_ok()); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); theory.add_column((Some(0), Some(1))); theory.add_row(std::iter::empty::<(LpCol, IntCst)>(), (Some(1), None)); @@ -636,7 +636,7 @@ pub mod test { }; assert!(expl.literals().is_empty(), "{:?}", expl.literals()); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); let col = theory.add_column((Some(0), Some(1))); theory.add_rows(std::iter::once((Some(1), Some(0), std::iter::once((col, 1))))); @@ -651,7 +651,7 @@ pub mod test { fn test_root_infeasibility_is_explained_by_nothing() { let mut model = Domains::new(); - let mut theory = LpRelax::default(); + let mut theory = Lp::default(); let acol = theory.add_column((Some(0), Some(0))); let bcol = theory.add_column((Some(0), Some(0))); theory.add_row([(acol, 1), (bcol, 1)].into_iter(), (Some(1), None)); diff --git a/solver/lp_highs/src/state.rs b/solver/lp_highs/src/state.rs index a819c522..e8444bcc 100644 --- a/solver/lp_highs/src/state.rs +++ b/solver/lp_highs/src/state.rs @@ -331,7 +331,7 @@ impl LpState { }); } - pub fn solve_or_iis(&mut self, stats: &mut crate::LpRelaxStats) -> Result { + pub fn solve_or_iis(&mut self, stats: &mut crate::LpStats) -> Result { let time = std::time::Instant::now(); let res = self.lp_model.solve_or_iis(); From 57e554fc5e5edd0e3f2f7a957ebdad074b27001d Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 15:35:28 +0200 Subject: [PATCH 07/19] feat(timelines): allow retrieving groundings of "empty source" in grounder interface --- planning/timelines/src/analysis/grounding/mod.rs | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/planning/timelines/src/analysis/grounding/mod.rs b/planning/timelines/src/analysis/grounding/mod.rs index 35fff5f5..64c33082 100644 --- a/planning/timelines/src/analysis/grounding/mod.rs +++ b/planning/timelines/src/analysis/grounding/mod.rs @@ -57,6 +57,11 @@ impl Groundings { .iter() .filter_map(|(k, v)| k.map(|task_id| (task_id, v.as_slice()))) } + + /// Groundings of the empty source. Since it is usually fully ground (no variables involved), there is usually exactly one, empty assignment. + pub fn empty_source_groundings(&self) -> &[ParametersAssignment] { + self.groundings.get(&None).map(|v| v.as_slice()).unwrap_or(&[]) + } } /// Ground all tasks appear in this problem. From 10733ba0a0faf4e82c17ecc8f3b88dac38a2a174 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 15:38:54 +0200 Subject: [PATCH 08/19] chore(timelines): refactor collecting ambiguous (nonsimple) conditions and effects --- planning/timelines/src/analysis/nonsimple.rs | 35 ++++++++++++-------- 1 file changed, 22 insertions(+), 13 deletions(-) diff --git a/planning/timelines/src/analysis/nonsimple.rs b/planning/timelines/src/analysis/nonsimple.rs index 2da73084..1c98e8df 100644 --- a/planning/timelines/src/analysis/nonsimple.rs +++ b/planning/timelines/src/analysis/nonsimple.rs @@ -49,24 +49,33 @@ fn collect_nonsimple_effects(ctx: &SchedEncoder) -> HashSet { fn collect_nonsimple_conditions(nonsimple_effects: &mut HashSet, ctx: &SchedEncoder) -> HashSet { let mut res = HashSet::new(); - for cl in ctx.causal_links.get_links() { - if res.contains(&cl.cond_id) { - nonsimple_effects.insert(cl.eff_id); - continue; + for (cond_id, cond) in ctx.causal_links.conditions.iter().enumerate() { + if !all_nonconstant_terms_are_included_in_source_terms( + cond.state_var.args.iter().chain(&[cond.value]).copied(), + cond.source, + ctx, + ) { + res.insert(cond_id); } - let cond = &ctx.causal_links.conditions[cl.cond_id]; + } - if nonsimple_effects.contains(&cl.eff_id) - || !all_nonconstant_terms_are_included_in_source_terms( - cond.state_var.args.iter().chain(&[cond.value]).copied(), - cond.source, - ctx, - ) - { - nonsimple_effects.insert(cl.eff_id); + // Propagate "nonsimple-ness": + // if a causal link uses an effect or condition marked as nonsimple, mark the other member (condition or effect) as nonsimple too. + for cl in ctx.causal_links.get_links() { + if nonsimple_effects.contains(&cl.eff_id) { res.insert(cl.cond_id); + } else if res.contains(&cl.cond_id) { + nonsimple_effects.insert(cl.eff_id); } } + + // A nonsimple effect / condition is such that for all its causal links, the other member is also nonsimple + debug_assert!( + ctx.causal_links + .get_links() + .all(|cl| nonsimple_effects.contains(&cl.eff_id) == res.contains(&cl.cond_id)) + ); + res } From fe841ae9ed5373bd8504adf684f71bab70018088 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 15:53:15 +0200 Subject: [PATCH 09/19] chore(timelines): derive more traits for `EffectOp` --- planning/timelines/src/effects.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/planning/timelines/src/effects.rs b/planning/timelines/src/effects.rs index 6354ed73..39b5756a 100644 --- a/planning/timelines/src/effects.rs +++ b/planning/timelines/src/effects.rs @@ -36,7 +36,7 @@ pub struct Effect { /// (mutex conditions). pub source: Option, } -#[derive(Clone, Eq, PartialEq)] +#[derive(Clone, Eq, PartialEq, PartialOrd, Ord, Hash)] pub enum EffectOp { /// Sets the state variable to an absolute value Assign(IntTerm), From e0b10fb316d5c7e988d97f0fb2f7eb3a09fba22c Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 15:59:30 +0200 Subject: [PATCH 10/19] feat(timelines): add `transitions` module a flag can be set to recover initial effects that may have been filtered out by the main encoding, due to being proven unable to support any condition. if set, these the transition store will include transitions corresponding to these ground initial effects, with their default value, which may not necessarily be the one that the original filtered-out effect may have had! (note, however, that this is not impactful for our use case) --- planning/timelines/src/analysis/mod.rs | 2 + .../transitions/closed_world_default.rs | 104 +++ .../timelines/src/analysis/transitions/mod.rs | 611 ++++++++++++++++++ .../analysis/transitions/tests/visitall.rs | 240 +++++++ 4 files changed, 957 insertions(+) create mode 100644 planning/timelines/src/analysis/transitions/closed_world_default.rs create mode 100644 planning/timelines/src/analysis/transitions/mod.rs create mode 100644 planning/timelines/src/analysis/transitions/tests/visitall.rs diff --git a/planning/timelines/src/analysis/mod.rs b/planning/timelines/src/analysis/mod.rs index 69c93b9b..0a2ddea2 100644 --- a/planning/timelines/src/analysis/mod.rs +++ b/planning/timelines/src/analysis/mod.rs @@ -1,5 +1,7 @@ pub mod grounding; mod nonsimple; +#[allow(dead_code)] +pub mod transitions; pub use nonsimple::collect_nonsimple_conditions_and_effects_to_relax; diff --git a/planning/timelines/src/analysis/transitions/closed_world_default.rs b/planning/timelines/src/analysis/transitions/closed_world_default.rs new file mode 100644 index 00000000..6f12d413 --- /dev/null +++ b/planning/timelines/src/analysis/transitions/closed_world_default.rs @@ -0,0 +1,104 @@ +use std::collections::{HashMap, HashSet}; + +use aries_solver::{core::IntCst, lang::Lit}; + +use super::EffectBasicInfo; +use crate::{EffectId, EffectOp, IntTerm, SchedEncoder, StateVar, encoder::CondId}; + +/// Closed-world default (ground) initial effects in place of those omitted in (or "missing" from) the main encoding +/// after being derived as unable to support any condition (see [`add_closed_world_negative_effects`]). +/// +/// Note that the values of the effects we're "recovering" here do not need to match those of the "original" omitted initial effects, +/// because if they were needed, they wouldn't have been pruned in the main encoding. +/// +/// Effect ids below `first_id` correspond to the "original" effects of the encoding. +#[derive(Clone, Default)] +pub(super) struct ClosedWorldDefaultEffects { + first_id: EffectId, + store: Vec, + pub ignored_fluents: HashSet, + original_initial_effects_ground_args: HashMap>>, +} +impl ClosedWorldDefaultEffects { + pub fn new( + ctx: &SchedEncoder, + effects_to_ignore: impl IntoIterator, + conditions_to_ignore: impl IntoIterator, + mut original_initial_effects_ground_args: HashMap>>, + ) -> Self { + for (_, entry) in original_initial_effects_ground_args.iter_mut() { + debug_assert!({ + use itertools::Itertools; + entry.iter().all_unique() + }); + entry.sort_unstable(); + } + + Self { + first_id: ctx.sched.effects.iter().count(), + store: vec![], + ignored_fluents: HashSet::::from_iter( + effects_to_ignore + .into_iter() + .map(|e_id| ctx.sched.effects.get(e_id).state_var.fluent.clone()) + .chain( + conditions_to_ignore + .into_iter() + .map(|c_id| ctx.causal_links.conditions.get(c_id).state_var.fluent.clone()), + ), + ), + original_initial_effects_ground_args, + } + } + pub fn is_empty(&self) -> bool { + self.store.is_empty() + } + pub fn contains(&self, eff_id: EffectId) -> bool { + eff_id >= self.first_id && eff_id < self.first_id + self.store.len() + } + pub fn get(&self, offset_eff_id: EffectId) -> &EffectBasicInfo { + &self.store[offset_eff_id - self.first_id] + } + pub fn add( + &mut self, + fluent: crate::Sym, + args: impl Into>, + value: IntCst, + ) -> Result { + let args = args.into(); + + // Ignore if there already is a (non-ignored) initial effect with these ground args. + if self + .original_initial_effects_ground_args + .get(&fluent) + .is_some_and(|known_grs| known_grs.contains(&args)) + { + return Err(()); + } + + let eff_view = EffectBasicInfo { + state_var: StateVar { + fluent, + args: args.into_iter().map(IntTerm::int_cst).collect(), + }, + operation: EffectOp::Assign(IntTerm::int_cst(value)), + prez: Lit::TRUE, + source: None, + }; + + debug_assert!( + eff_view + .state_var + .args + .iter() + .chain(match &eff_view.operation { + EffectOp::Assign(term) => [term], + EffectOp::Step(_term) => todo!(), + }) + .all(|term| term.is_cst()) + ); + + self.store.push(eff_view); + Ok(self.first_id + self.store.len() - 1) + } +} diff --git a/planning/timelines/src/analysis/transitions/mod.rs b/planning/timelines/src/analysis/transitions/mod.rs new file mode 100644 index 00000000..127b5b7a --- /dev/null +++ b/planning/timelines/src/analysis/transitions/mod.rs @@ -0,0 +1,611 @@ +mod closed_world_default; + +use closed_world_default::ClosedWorldDefaultEffects; + +use crate::analysis::{Source, collect_nonsimple_conditions_and_effects_to_relax}; +use crate::{EffectId, EffectOp, IntTerm, SchedEncoder, StateVar, TaskId, constraints::HasValueAt, encoder::CondId}; + +use aries_solver::{core::IntCst, lang::Lit}; +use idmap::DirectIdMap; + +pub type TransitionId = usize; + +#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq, PartialOrd, Ord)] +pub enum TransitionType { + Cond, + Eff, + CondEff, +} + +#[derive(Debug, Clone, Copy, Hash, PartialEq, Eq, PartialOrd, Ord)] +pub enum Transition { + Cond(CondId), + Eff(EffectId), + /// A condition and effect sharing the same source, presence literal, and state variable. + CondEff(CondId, EffectId), +} +impl Transition { + pub fn tpe(&self) -> TransitionType { + match self { + Transition::Cond(_) => TransitionType::Cond, + Transition::Eff(_) => TransitionType::Eff, + Transition::CondEff(_, _) => TransitionType::CondEff, + } + } +} + +/// Invariant: condition transitionss' `op` must be `EffectOp::Assign(val.unwrap())`. +#[derive(Debug, Clone, Hash, PartialEq, Eq, PartialOrd, Ord)] +struct TransitionTermsView<'a> { + // tpe: TransitionType, + args: &'a [IntTerm], + val: Option, + op: EffectOp, +} +impl<'a> TransitionTermsView<'a> { + // pub fn tpe(&self) -> TransitionType { + // self.tpe + // } + pub fn args(&'a self) -> &'a [IntTerm] { + self.args + } + pub fn val(&self) -> Option { + self.val + } + pub fn op(&self) -> &EffectOp { + &self.op + } + #[allow(dead_code)] + pub fn op_assign(&self) -> IntTerm { + match &self.op { + EffectOp::Assign(term) => *term, + _ => panic!("effect operation must be an assign"), + } + } +} + +/// An `None` value means the corresponding transition term doesn't actually appear in its source's terms +/// (e.g. because it's a constant, see [`collect_nonsimple_conditions_and_effects_to_relax`]). +#[derive(Clone)] +struct TransitionTermsIndicesInSource { + pub args: smallvec::SmallVec<[Option; 4]>, + pub val: Option>, + pub op: Option, +} + +struct EffectBasicInfoView<'a> { + source: Source, + prez: Lit, + state_var: &'a StateVar, + op: &'a EffectOp, +} +impl crate::Effect { + fn view<'a>(&'a self) -> EffectBasicInfoView<'a> { + EffectBasicInfoView { + source: self.source, + prez: self.prez, + state_var: &self.state_var, + op: &self.operation, + } + } +} +#[derive(Clone)] +struct EffectBasicInfo { + source: Source, + prez: Lit, + state_var: StateVar, + operation: EffectOp, +} +impl EffectBasicInfo { + fn view<'a>(&'a self) -> EffectBasicInfoView<'a> { + EffectBasicInfoView { + source: self.source, + prez: self.prez, + state_var: &self.state_var, + op: &self.operation, + } + } +} + +#[derive(Clone)] +pub(crate) struct Transitions { + /// Stores "unambiguous" transitions (i.e. not "nonsimple" ones, see [`collect_nonsimple_conditions_and_effects_to_relax`]). + store: Vec, + /// For each transition, stores the indices of its terms in its source's terms. + /// This is needed to evaluate a transition's grounding given a grounding of its source. + /// (Reminder: one of the requirements of an "unambiguous" transition is that its terms are either constant or appear in its source's terms / arguments). + transition_terms_indices_in_source: Vec, + + recovered_closed_world_defaults: ClosedWorldDefaultEffects, + + of_condition: DirectIdMap, + of_effect: DirectIdMap, + of_concrete_source: DirectIdMap>, + of_empty_source: Vec, +} + +impl Transitions { + #[allow(dead_code)] + pub fn get(&self, trans_id: TransitionId) -> Transition { + self.store[trans_id] + } + pub fn of_condition(&self, cond_id: CondId) -> Option { + self.of_condition.get(cond_id).copied() + } + pub fn of_effect(&self, eff_id: EffectId) -> Option { + self.of_effect.get(eff_id).copied() + } + pub fn of_source(&self, source: Source) -> &[TransitionId] { + if let Some(task_id) = source { + self.of_concrete_source + .get(task_id) + .map(|trs| trs.as_slice()) + .unwrap_or_default() + } else { + &self.of_empty_source + } + } + + #[allow(dead_code)] + pub fn iter(&self) -> impl Iterator { + self.store.iter().copied().enumerate() + } + #[allow(dead_code)] + pub fn iter_of_conditions(&self) -> impl Iterator { + self.of_condition.iter().map(|(cond_id, &trans_id)| (cond_id, trans_id)) + } + pub fn iter_of_effects(&self) -> impl Iterator { + self.of_effect.iter().map(|(eff_id, &trans_id)| (eff_id, trans_id)) + } + pub fn iter_of_sources(&self) -> impl Iterator)> { + std::iter::chain( + [(None, &self.of_empty_source)], + self.of_concrete_source + .iter() + .map(|(task_id, trans_ids)| (Some(task_id), trans_ids)), + ) + } + + pub fn get_condition<'a>(&self, trans_id: TransitionId, ctx: &'a SchedEncoder) -> Option<(CondId, &'a HasValueAt)> { + match self.store[trans_id] { + Transition::Cond(cond_id) | Transition::CondEff(cond_id, _) => { + Some((cond_id, ctx.causal_links.conditions.get(cond_id))) + } + Transition::Eff(_) => None, + } + } + fn get_effect_info<'a>( + &'a self, + trans_id: TransitionId, + ctx: &'a SchedEncoder, + ) -> Option<(EffectId, EffectBasicInfoView<'a>)> { + match self.store[trans_id] { + Transition::Eff(eff_id) | Transition::CondEff(_, eff_id) => { + let info_view: EffectBasicInfoView<'_> = if self.recovered_closed_world_defaults.contains(eff_id) { + self.recovered_closed_world_defaults.get(eff_id).view() + } else { + ctx.sched.effects.get(eff_id).view() + }; + Some((eff_id, info_view)) + } + Transition::Cond(_) => None, + } + } + pub fn get_prez(&self, trans_id: TransitionId, ctx: &SchedEncoder) -> Lit { + match self.store[trans_id].tpe() { + TransitionType::Cond => self.get_condition(trans_id, ctx).unwrap().1.prez, + TransitionType::Eff => self.get_effect_info(trans_id, ctx).unwrap().1.prez, + TransitionType::CondEff => { + let res = self.get_effect_info(trans_id, ctx).unwrap().1.prez; + debug_assert!(res == self.get_condition(trans_id, ctx).unwrap().1.prez); + res + } + } + } + pub fn get_source(&self, trans_id: TransitionId, ctx: &SchedEncoder) -> Source { + match self.store[trans_id].tpe() { + TransitionType::Cond => self.get_condition(trans_id, ctx).unwrap().1.source, + TransitionType::Eff => self.get_effect_info(trans_id, ctx).unwrap().1.source, + TransitionType::CondEff => { + let res = self.get_effect_info(trans_id, ctx).unwrap().1.source; + debug_assert!(res == self.get_condition(trans_id, ctx).unwrap().1.source); + res + } + } + } + pub fn get_state_var<'a>(&'a self, trans_id: TransitionId, ctx: &'a SchedEncoder) -> &'a StateVar { + match self.store[trans_id].tpe() { + TransitionType::Cond => &self.get_condition(trans_id, ctx).unwrap().1.state_var, + TransitionType::Eff => self.get_effect_info(trans_id, ctx).unwrap().1.state_var, + TransitionType::CondEff => { + let res = self.get_effect_info(trans_id, ctx).unwrap().1.state_var; + debug_assert!(*res == self.get_condition(trans_id, ctx).unwrap().1.state_var); + res + } + } + } + + pub(self) fn get_terms<'a>(&'a self, trans_id: TransitionId, ctx: &'a SchedEncoder) -> TransitionTermsView<'a> { + let args = self.get_state_var(trans_id, ctx).args.as_slice(); + match self.store[trans_id].tpe() { + TransitionType::Cond => { + let val = self.get_condition(trans_id, ctx).unwrap().1.value; + TransitionTermsView { + args, + val: Some(val), + op: EffectOp::Assign(val), + } + } + TransitionType::Eff => { + let op = self.get_effect_info(trans_id, ctx).unwrap().1.op.clone(); + TransitionTermsView { args, val: None, op } + } + TransitionType::CondEff => { + let val = self.get_condition(trans_id, ctx).unwrap().1.value; + let op = self.get_effect_info(trans_id, ctx).unwrap().1.op.clone(); + TransitionTermsView { + args, + val: Some(val), + op, + } + } + } + } + + #[allow(dead_code)] + pub fn is_recovered_closed_world_default(&self, eff_id: EffectId) -> bool { + self.recovered_closed_world_defaults.contains(eff_id) + } + pub fn are_recovered_closed_world_default_effects_empty(&self) -> bool { + self.recovered_closed_world_defaults.is_empty() + } + + /// Collects transitions from "unambiguous" conditions and effects (i.e. filtering out "nonsimple" ones) + pub fn new_unambiguous(ctx: &SchedEncoder, recover_closed_world_defaults: bool) -> Self { + // Collects nonsimple transitions to ignore / relax. + let (conditions_to_ignore, effects_to_ignore) = collect_nonsimple_conditions_and_effects_to_relax(ctx); + + // Group conditions and effects by sources + + let mut empty_source_conditions = vec![]; + let mut concrete_source_conditions = DirectIdMap::default(); + let mut empty_source_effects = vec![]; + let mut concrete_source_effects = DirectIdMap::default(); + + for (cond_id, c) in ctx.causal_links.conditions.iter().enumerate() { + if conditions_to_ignore.contains(&cond_id) { + continue; + } + if let Some(task_id) = c.source { + if !concrete_source_conditions.contains_key(task_id) { + concrete_source_conditions.insert(task_id, vec![]); + } + concrete_source_conditions.get_mut(task_id).unwrap().push((cond_id, c)); + } else { + empty_source_conditions.push((cond_id, c)); + } + } + for (eff_id, e) in ctx.sched.effects.iter().enumerate() { + if effects_to_ignore.contains(&eff_id) { + continue; + } + if let Some(task_id) = e.source { + if !concrete_source_effects.contains_key(task_id) { + concrete_source_effects.insert(task_id, vec![]); + } + concrete_source_effects.get_mut(task_id).unwrap().push((eff_id, e)); + } else { + empty_source_effects.push((eff_id, e)); + } + } + + // First, iterate over conditions (grouped by sources) and introduce corresponding Cond transitions. + // Then, iterate over effects (grouped by sources) and the conditions for those sources. + // When a compatible condition and effect are found, a corresponding CondEff transition is introduced, + // modifying the previously inserted Cond transition. + // If no compatible condition is found, a Eff transition is introduced. + // + // If a ground initial (empty source) Eff transition is introduced, remember that grounding. + // This is needed to avoid overriding it later when recovering "missing" closed world default initial effects. + + let mut store = vec![]; + + let mut of_condition = DirectIdMap::default(); + let mut of_effect = DirectIdMap::default(); + let mut of_empty_source = vec![]; + let mut of_concrete_source = DirectIdMap::default(); + + let mut initial_effects_ground_args = + std::collections::HashMap::>>::new(); + + let source_conds_iter = std::iter::chain( + [(None, &empty_source_conditions)], + concrete_source_conditions + .iter() + .map(|(task_id, conds)| (Some(task_id), conds)), + ); + let source_effs_iter = std::iter::chain( + [(None, &empty_source_effects)], + concrete_source_effects + .iter() + .map(|(task_id, effs)| (Some(task_id), effs)), + ); + + for (source, cs) in source_conds_iter { + if let Some(task_id) = source + && !of_concrete_source.contains_key(task_id) + { + of_concrete_source.insert(task_id, vec![]); + } + for &(cond_id, _) in cs { + let trans_id = store.len(); + of_condition.insert(cond_id, trans_id); + if let Some(task_id) = source { + of_concrete_source.get_mut(task_id).unwrap().push(trans_id); + } else { + of_empty_source.push(trans_id); + } + store.push(Transition::Cond(cond_id)); + } + } + for (source, es) in source_effs_iter { + if let Some(task_id) = source + && !of_concrete_source.contains_key(task_id) + { + of_concrete_source.insert(task_id, vec![]); + } + for &(eff_id, e) in es { + let mut compatible_conds_found = 0; + + // No CondEff pattern allowed for empty source. + if source.is_some() { + let cs = if let Some(task_id) = source { + concrete_source_conditions.get(task_id) + } else { + Some(&empty_source_conditions) + } + .into_iter() + .flatten(); + + for &(cond_id, c) in cs { + if e.state_var == c.state_var && e.prez == c.prez { + // Change the previously inserted Cond transition into a CondEff + let trans_id = *of_condition.get(cond_id).unwrap(); + of_effect.insert(eff_id, trans_id); + store[trans_id] = Transition::CondEff(cond_id, eff_id); + + compatible_conds_found += 1; + } + } + debug_assert!(compatible_conds_found <= 1); + } + + // Add a new Eff transition if the effect doesn't correspond to a CondEff + if compatible_conds_found == 0 { + let trans_id = store.len(); + of_effect.insert(eff_id, trans_id); + if let Some(task_id) = source { + of_concrete_source.get_mut(task_id).unwrap().push(trans_id); + } else { + of_empty_source.push(trans_id); + } + store.push(Transition::Eff(eff_id)); + } + + // Remember the args groundings of ground initial effects + if recover_closed_world_defaults + && source.is_none() + && e.state_var.args.iter().all(|term| term.is_cst()) + { + let ground_args = e.state_var.args.iter().map(|term| term.constant).collect(); + initial_effects_ground_args + .entry(e.state_var.fluent.to_string()) + .or_default() + .push(ground_args); + debug_assert!({ + use itertools::Itertools; + initial_effects_ground_args + .get(&e.state_var.fluent) + .unwrap() + .iter() + .all_unique() + }); + } + } + } + + // Loop over fluents and their parameter types' ground values. + // For each such grounding, introduce a default-valued initial effect (closed world default), + // if there wasn't already an effect with the same ground parameters encountered earlier + // (among the "explicit" known initial effects accessible from `ctx`). + + let recovered_closed_world_defaults = if !recover_closed_world_defaults { + ClosedWorldDefaultEffects::default() + } else { + let mut recovered_closed_world_defaults = ClosedWorldDefaultEffects::new( + ctx, + effects_to_ignore, + conditions_to_ignore, + initial_effects_ground_args, + ); + + for (sym, params, _) in ctx.sched.fluents.iter() { + if recovered_closed_world_defaults.ignored_fluents.contains(sym) { + continue; + } + + let args = crate::boxes::BBox::new(params.iter().map(|p| p.range).collect::>()); + let mut grs = args.as_ref().points(); + while let Some(gr) = streaming_iterator::StreamingIterator::next(&mut grs) { + let args_ground = Vec::from_iter(gr.iter().copied()); + + if let Ok(eff_id) = recovered_closed_world_defaults.add( + sym.to_string(), + args_ground, + ctx.sched.fluents.get_return(sym).unwrap().range.first, + ) { + let tr_id = store.len(); + of_effect.insert(eff_id, tr_id); + of_empty_source.push(tr_id); + store.push(Transition::Eff(eff_id)); + } else { + // Ignored (not added) as there already in an initial effect with these ground args. + }; + } + } + + recovered_closed_world_defaults + }; + + // For each transition, collect its terms' (args and values) indices in the list of its source's args. + // + // Note that currently, transitions whose terms contain auxiliary or reification variables + // that do not appearing in the the source's args are ignored anyway (filtered out as "nonsimple") + + let mut transition_terms_indices_in_source = Vec::with_capacity(store.len()); + + let get_source_terms = |source| { + if let Some(task_id) = source { + &ctx.sched.tasks[task_id].args + } else { + &ctx.sched.global_args + } + }; + let get_effect_info = |eff_id| { + if recovered_closed_world_defaults.contains(eff_id) { + recovered_closed_world_defaults.get(eff_id).view() + } else { + ctx.sched.effects.get(eff_id).view() + } + }; + let get_source = |transition| match transition { + Transition::Cond(cond_id) => ctx.causal_links.conditions.get(cond_id).source, + Transition::Eff(eff_id) => get_effect_info(eff_id).source, + Transition::CondEff(cond_id, eff_id) => { + let res = ctx.causal_links.conditions.get(cond_id).source; + debug_assert!(res == get_effect_info(eff_id).source); + debug_assert!(!recovered_closed_world_defaults.contains(eff_id)); + res + } + }; + let get_transition_terms = |transition| match transition { + Transition::Cond(cond_id) => TransitionTermsView { + args: &ctx.causal_links.conditions.get(cond_id).state_var.args, + val: Some(ctx.causal_links.conditions.get(cond_id).value), + op: EffectOp::Assign(ctx.causal_links.conditions.get(cond_id).value), + }, + Transition::Eff(eff_id) => TransitionTermsView { + args: &get_effect_info(eff_id).state_var.args, + val: None, + op: get_effect_info(eff_id).op.clone(), + }, + Transition::CondEff(cond_id, eff_id) => { + debug_assert!(!recovered_closed_world_defaults.contains(eff_id)); + TransitionTermsView { + args: &ctx.causal_links.conditions.get(cond_id).state_var.args, + val: Some(ctx.causal_links.conditions.get(cond_id).value), + op: get_effect_info(eff_id).op.clone(), + } + } + }; + + for transition in store.iter() { + let source_terms = get_source_terms(get_source(*transition)); + let transition_terms = get_transition_terms(*transition); + + let index_in_source = |term: IntTerm| -> Option { + if term.is_cst() { + return None; + } + let idx = source_terms.iter().position(|&t| t == term); + debug_assert!( + idx.is_some(), + "non-constant transition term absent from its source's args (such transitions are 'nonsimple' and must have been filtered out)" + ); + idx.map(|idx| idx as u16) + }; + + let transition_args_indices_in_source = + transition_terms.args().iter().copied().map(index_in_source).collect(); + let transition_val_index_in_source = transition_terms.val().map(index_in_source); + let transition_op_index_in_source = match transition_terms.op() { + EffectOp::Assign(term) => index_in_source(*term), + EffectOp::Step(_) => todo!(), + }; + + transition_terms_indices_in_source.push(TransitionTermsIndicesInSource { + args: transition_args_indices_in_source, + val: transition_val_index_in_source, + op: transition_op_index_in_source, + }); + } + + Self { + store, + transition_terms_indices_in_source, + of_condition, + of_effect, + of_empty_source, + of_concrete_source, + recovered_closed_world_defaults, + } + } +} + +#[cfg(test)] +mod tests { + pub(crate) mod visitall; + + use crate::analysis::collect_nonsimple_conditions_and_effects_to_relax; + use crate::analysis::transitions::tests::visitall::{VisitAllLine, build_and_encode_visitall_line}; + use crate::analysis::transitions::{TransitionType, Transitions}; + + #[test] + fn test_transitions_visitall_line() { + let encoder = build_and_encode_visitall_line( + &VisitAllLine { + num_locs: 5, + num_moves: 4, + }, + false, + ); + + let transitions = Transitions::new_unambiguous(&encoder, true); + + assert!({ + let (conditions_to_ignore, effects_to_ignore) = collect_nonsimple_conditions_and_effects_to_relax(&encoder); + conditions_to_ignore.is_empty() && effects_to_ignore.is_empty() + }); + + assert_eq!(transitions.iter().count(), 56); + assert_eq!( + transitions + .iter() + .filter(|(_, tr)| tr.tpe() == TransitionType::Cond) + .count(), + 9 + ); + assert_eq!( + transitions + .iter() + .filter(|(_, tr)| tr.tpe() == TransitionType::Eff) + .count(), + 43 + ); + assert_eq!( + transitions + .iter() + .filter(|(_, tr)| tr.tpe() == TransitionType::CondEff) + .count(), + 4 + ); + + assert_eq!( + transitions + .iter_of_effects() + .filter(|&(eff_id, _)| transitions.is_recovered_closed_world_default(eff_id)) + .count(), + 25 + ); + } +} diff --git a/planning/timelines/src/analysis/transitions/tests/visitall.rs b/planning/timelines/src/analysis/transitions/tests/visitall.rs new file mode 100644 index 00000000..8b0589dd --- /dev/null +++ b/planning/timelines/src/analysis/transitions/tests/visitall.rs @@ -0,0 +1,240 @@ +use aries_solver::core::state::Evaluable; +use aries_solver::lang::ModelView; +use aries_solver::prelude::*; + +use crate::boxes::Segment; +use crate::constraints::HasValueAt; +use crate::symbols::ObjectEncoding; +use crate::{ + Effect, EffectOp, FluentParam, FluentsEncoding, IntTerm, Sched, SchedEncoder, Solution, StateVar, Task, TaskId, + VarCst, +}; + +/// A `visitall` instance: `num_locs` locations *in a line*, every one of which must be visited, with `num_moves` `move` actions available. +#[derive(Debug)] +pub struct VisitAllLine { + pub num_locs: usize, + pub num_moves: usize, +} + +impl VisitAllLine { + fn locs_names(&self) -> Vec { + (0..self.num_locs).map(|i| format!("loc-x{i}")).collect() + } +} + +#[allow(dead_code)] +pub fn build_and_encode_visitall_line(pb: &VisitAllLine, print: bool) -> SchedEncoder { + let (sched, _) = build_visitall_line(pb); + + let mut encoder = sched.clone().encoder(); + for c in sched.constraints.iter() { + c.enforce(&mut encoder); + } + + if print { + let effs = encoder.sched.effects.iter().collect::>(); + let conds = encoder.causal_links.conditions.iter().collect::>(); + let causal_links = encoder.causal_links.get_links().collect::>(); + + println!("Effects:"); + for (eid, e) in effs.iter().enumerate() { + println!(" {eid}: {e:?}"); + } + println!("Conditions:"); + for (cid, c) in conds.iter().enumerate() { + println!(" {cid}: {c:?}"); + } + println!("Causal Links:"); + for cl in causal_links.iter() { + println!(" {cl:?}"); + } + } + + encoder +} + +pub fn build_visitall_line(pb: &VisitAllLine) -> (Sched, Vec) { + let names = pb.locs_names(); + + let objects = ObjectEncoding::build( + "object".into(), + |t| match t.as_str() { + "object" => vec!["loc".into()], + _ => vec![], + }, + { + let names = names.clone(); + move |t| match t.as_str() { + "loc" => names.clone(), + _ => vec![], + } + }, + ); + + let locs = objects.domain_of_type("loc").unwrap(); + let loc_range = Segment::new(locs.first, locs.last); + let bool_range = Segment::new(0, 1); + + let mut fluents = FluentsEncoding::empty(); + fluents.add( + "connected".into(), + &[FluentParam { range: loc_range }, FluentParam { range: loc_range }], + FluentParam { range: bool_range }, + ); + fluents.add( + "at-robot".into(), + &[FluentParam { range: loc_range }], + FluentParam { range: bool_range }, + ); + fluents.add( + "visited".into(), + &[FluentParam { range: loc_range }], + FluentParam { range: bool_range }, + ); + + let mut model = Sched::new(1, objects, fluents); + + let loc: Vec = names.iter().map(|n| model.objects.object_id(n).unwrap()).collect(); + + // (:init …): robot at loc-x0, which is already visited, and a bidirectional chain + init_bool(&mut model, "at-robot", &[loc[0]], true); + init_bool(&mut model, "visited", &[loc[0]], true); + for w in loc.windows(2) { + init_bool(&mut model, "connected", &[w[0], w[1]], true); + init_bool(&mut model, "connected", &[w[1], w[0]], true); + } + + // (:goal (and (visited loc-x0) … )) + for &l in &loc { + model.add_constraint(HasValueAt { + state_var: state_var("visited", vec![l.into()]), + value: IntTerm::TRUE, + timepoint: model.horizon, + prez: Lit::TRUE, + source: None, + }); + } + + let moves = (0..pb.num_moves).map(|_| add_move(&mut model, &loc)).collect(); + (model, moves) +} + +/// (:action move :parameters (?curpos ?nextpos - loc)) +fn add_move(model: &mut Sched, loc: &[IntCst]) -> Move { + let presence = model.new_bool_var(); + let start: VarCst = model.new_opt_timepoint(presence); + let end: VarCst = start + 1; + + let (first, last) = (*loc.first().unwrap(), *loc.last().unwrap()); + let curpos = model.new_optional_var(first, last, presence); + let nextpos = model.new_optional_var(first, last, presence); + + let task_id = model.add_task(Task { + name: "move".into(), + start, + end, + presence, + args: vec![curpos.into(), nextpos.into()], + }); + + // :precondition (and (at-robot ?curpos) (connected ?curpos ?nextpos)) + for sv in [ + state_var("at-robot", vec![curpos.into()]), + state_var("connected", vec![curpos.into(), nextpos.into()]), + ] { + model.add_constraint(HasValueAt { + state_var: sv, + value: IntTerm::TRUE, + timepoint: start, + prez: presence, + source: Some(task_id), + }); + } + + // :effect (and (at-robot ?nextpos) (not (at-robot ?curpos)) (visited ?nextpos)) + for (sv, value) in [ + (state_var("at-robot", vec![nextpos.into()]), IntTerm::TRUE), + (state_var("at-robot", vec![curpos.into()]), IntTerm::ZERO), + (state_var("visited", vec![nextpos.into()]), IntTerm::TRUE), + ] { + let mutex_end = model.new_opt_timepoint(presence); + model.add_effect(Effect { + transition_start: start, + transition_end: end, + mutex_end, + state_var: sv, + operation: EffectOp::Assign(value), + prez: presence, + source: Some(task_id), + }); + } + + Move { + presence, + start, + curpos, + nextpos, + _task_id: task_id, + } +} + +#[derive(Debug)] +pub struct Move { + presence: Lit, + start: VarCst, + curpos: Var, + nextpos: Var, + _task_id: TaskId, +} + +impl Evaluable for Move { + type Value = (IntCst, IntCst, IntCst); + + fn evaluate(&self, solution: &Solution) -> Option { + if !solution.entails(self.presence) { + return None; + } + Some(( + solution.eval(self.start).unwrap(), + solution.eval(self.curpos).unwrap(), + solution.eval(self.nextpos).unwrap(), + )) + } +} + +fn state_var(fluent: &str, args: Vec) -> StateVar { + StateVar { + fluent: fluent.into(), + args, + } +} + +fn init_bool(model: &mut Sched, fluent: &str, args: &[IntCst], value: bool) { + let mutex_end = model.new_timepoint(); + model.add_effect(Effect { + transition_start: model.origin, + transition_end: model.origin, + mutex_end, + state_var: state_var(fluent, args.iter().map(|&a| a.into()).collect()), + operation: EffectOp::Assign(if value { IntTerm::TRUE } else { IntTerm::ZERO }), + prez: Lit::TRUE, + source: None, + }); +} + +#[cfg(test)] +mod test { + use super::{VisitAllLine, build_and_encode_visitall_line}; + + #[test] + fn build_simple_visitall_line() { + build_and_encode_visitall_line( + &VisitAllLine { + num_locs: 5, + num_moves: 4, + }, + true, + ); + } +} From 3ed212461b0189d1bafd6b9f15367d546a44d2ad Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 16:01:36 +0200 Subject: [PATCH 11/19] feat(timelines): add convenience methods to evaluate a transition's terms on a grounding of its source --- .../src/analysis/transitions/ground.rs | 135 ++++++++++++++++++ .../timelines/src/analysis/transitions/mod.rs | 1 + 2 files changed, 136 insertions(+) create mode 100644 planning/timelines/src/analysis/transitions/ground.rs diff --git a/planning/timelines/src/analysis/transitions/ground.rs b/planning/timelines/src/analysis/transitions/ground.rs new file mode 100644 index 00000000..20410135 --- /dev/null +++ b/planning/timelines/src/analysis/transitions/ground.rs @@ -0,0 +1,135 @@ +use aries_solver::core::IntCst; + +use super::{TransitionId, Transitions}; + +use crate::analysis::grounding::ParametersAssignment; +use crate::{EffectOp, IntTerm, SchedEncoder}; + +#[derive(Debug, Clone, Hash, PartialEq, Eq, PartialOrd, Ord)] +pub struct TransitionTermsEvaluation { + // tpe: TransitionType, + args: smallvec::SmallVec<[IntCst; 4]>, + val: Option, + op: EffectOp, +} +impl TransitionTermsEvaluation { + pub fn new(args: impl Into>, val: Option, op: EffectOp) -> Self { + assert!(match op { + EffectOp::Assign(term) => term.is_cst(), + EffectOp::Step(_) => todo!(), + }); + Self { + /*tpe,*/ args: args.into(), + val, + op, + } + } + // pub fn tpe(&self) -> TransitionType { + // self.tpe + // } + pub fn args_evaluated(&self) -> &[IntCst] { + &self.args + } + pub fn val_evaluated(&self) -> Option { + self.val + } + // pub fn op(&self) -> &EffectOp { + // &self.op + // } + pub fn op_evaluated(&self) -> Option { + match &self.op { + EffectOp::Assign(term) => { + debug_assert!(term.is_cst()); + Some(term.constant) + } + EffectOp::Step(term) => { + debug_assert!(term.is_cst()); + Some(term.constant) + } // EffectOp::Erase => None, // IN THE FUTURE ? + } + } + pub fn op_evaluated_as_assign(&self) -> IntCst { + match &self.op { + EffectOp::Assign(term) => { + debug_assert!(term.is_cst()); + term.constant + } + _ => panic!("effect operation expected to be an assign"), + } + } +} + +impl Transitions { + pub fn evaluate_terms<'a>( + &'a self, + trans_id: TransitionId, + source_grounding: &ParametersAssignment, + ctx: &'a SchedEncoder, + ) -> TransitionTermsEvaluation { + let terms = self.get_terms(trans_id, ctx); + + let args = self.transition_terms_indices_in_source[trans_id] + .args + .iter() + .enumerate() + .map(|(j, i)| i.map_or(terms.args[j].constant, |i| source_grounding[i.into()])) + .collect::>(); + + let val = self.transition_terms_indices_in_source[trans_id] + .val + .map(|i| i.map_or(terms.val.unwrap().constant, |i| source_grounding[i.into()])); + + let op = self.transition_terms_indices_in_source[trans_id].op.map_or( + match terms.op { + EffectOp::Assign(term) => EffectOp::Assign(IntTerm::int_cst(term.constant)), + EffectOp::Step(_) => todo!(), + }, + |i| match terms.op { + EffectOp::Assign(_) => EffectOp::Assign(IntTerm::int_cst(source_grounding[i.into()])), + EffectOp::Step(_) => todo!(), + }, + ); + + TransitionTermsEvaluation::new(args, val, op) + } + pub fn iter_evaluated_non_constant_terms<'a>( + &'a self, + trans_id: TransitionId, + source_grounding: &ParametersAssignment, + ctx: &'a SchedEncoder, + ) -> impl Iterator { + let terms = self.get_terms(trans_id, ctx); + + let val = self.transition_terms_indices_in_source[trans_id].val.map(|i| { + ( + terms.val.unwrap(), + i.map_or(terms.val.unwrap().constant, |i| source_grounding[i.into()]), + ) + }); + + let op = self.transition_terms_indices_in_source[trans_id].op.map_or( + match terms.op { + EffectOp::Assign(term) => (term, term.constant), + EffectOp::Step(_) => todo!(), + }, + |i| match terms.op { + EffectOp::Assign(term) => (term, source_grounding[i.into()]), + EffectOp::Step(_) => todo!(), + }, + ); + + self.transition_terms_indices_in_source[trans_id] + .args + .iter() + .enumerate() + .map(|(j, i)| { + ( + terms.args[j], + i.map_or(terms.args[j].constant, |i| source_grounding[i.into()]), + ) + }) + .chain(val) + .chain([op]) + .filter(|(term, _)| !term.is_cst()) + } +} diff --git a/planning/timelines/src/analysis/transitions/mod.rs b/planning/timelines/src/analysis/transitions/mod.rs index 127b5b7a..601c65a2 100644 --- a/planning/timelines/src/analysis/transitions/mod.rs +++ b/planning/timelines/src/analysis/transitions/mod.rs @@ -1,4 +1,5 @@ mod closed_world_default; +pub mod ground; use closed_world_default::ClosedWorldDefaultEffects; From 3facab4e78408f4fc715ab9566cf482ec58fdcce Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 16:04:42 +0200 Subject: [PATCH 12/19] feat(timelines): allow collecting lifted (potential) supports between transitions a flag can be set to allow condition transitions to be considered as possible supporters --- .../timelines/src/analysis/transitions/mod.rs | 1 + .../src/analysis/transitions/supports.rs | 295 ++++++++++++++++++ 2 files changed, 296 insertions(+) create mode 100644 planning/timelines/src/analysis/transitions/supports.rs diff --git a/planning/timelines/src/analysis/transitions/mod.rs b/planning/timelines/src/analysis/transitions/mod.rs index 601c65a2..02ff16dd 100644 --- a/planning/timelines/src/analysis/transitions/mod.rs +++ b/planning/timelines/src/analysis/transitions/mod.rs @@ -1,5 +1,6 @@ mod closed_world_default; pub mod ground; +pub mod supports; use closed_world_default::ClosedWorldDefaultEffects; diff --git a/planning/timelines/src/analysis/transitions/supports.rs b/planning/timelines/src/analysis/transitions/supports.rs new file mode 100644 index 00000000..a748ef41 --- /dev/null +++ b/planning/timelines/src/analysis/transitions/supports.rs @@ -0,0 +1,295 @@ +use aries_solver::lang::Lit; + +use super::{TransitionId, TransitionType, Transitions}; +use crate::SchedEncoder; + +#[derive(Clone, Default)] +pub(crate) struct Supports { + pub with_condition_out_transitions: bool, + unsorted_out: Vec<((TransitionId, TransitionId), Option)>, +} +impl Supports { + pub fn from(transitions: &Transitions, ctx: &SchedEncoder, with_condition_out_transitions: bool) -> Self { + // Supporting stemming from the causal links in the main encoding. + + let supports_causal_links = ctx.causal_links.get_links().filter_map(|cl| { + // println!( + // "{:?} {:?}", + // (cl.eff_id, ctx.sched.effects.get(cl.eff_id)), + // (cl.cond_id, ctx.causal_links.conditions.get(cl.cond_id)) + // ); + let out_trans_id = transitions.of_effect(cl.eff_id)?; + let in_trans_id = transitions.of_condition(cl.cond_id)?; + + debug_assert_eq!( + transitions.get_state_var(out_trans_id, ctx).fluent, + transitions.get_state_var(in_trans_id, ctx).fluent, + ); + Some(((out_trans_id, in_trans_id), Some(cl.active))) + }); + + // Supports from effects (including "missing" ones) to other effects (non-initial) effects + // + // Supports from them to anything other than an effect are not needed (they otherwise wouldn't have been ignored in the main encoding). + // As such: This (obviously) assumes that the main encoding soundly omits initial effects that may not support the conditions. + + let iter_effects_info = || { + transitions + .iter_of_effects() + .map(|(_, trans_id)| transitions.get_effect_info(trans_id, ctx).unwrap()) + }; + let supports_from_effects_to_others = iter_effects_info().flat_map(move |(out_eff_id, _)| { + let out_trans_id = transitions.of_effect(out_eff_id).unwrap(); + + iter_effects_info().flat_map(move |(in_eff_id, _)| { + let in_trans_id = transitions.of_effect(in_eff_id).unwrap(); + + if out_trans_id == in_trans_id + || transitions.get(in_trans_id).tpe() == TransitionType::CondEff + || transitions.get_source(in_trans_id, ctx).is_none() // in-transition must not correspond to an initial effect + || transitions.get_state_var(out_trans_id, ctx).fluent + != transitions.get_state_var(in_trans_id, ctx).fluent + { + return None; + } + debug_assert!( + transitions.get(in_trans_id).tpe() == TransitionType::Eff + && transitions.get_source(in_trans_id, ctx).is_some() + ); + + let out_terms = transitions.get_terms(out_trans_id, ctx); + let in_terms = transitions.get_terms(in_trans_id, ctx); + + let incompatible_args = out_terms.args().iter().zip(in_terms.args()).any(|(out_term, in_term)| { + out_term.is_cst() && in_term.is_cst() && out_term.constant != in_term.constant + }); + + (!incompatible_args).then_some(((out_trans_id, in_trans_id), None::)) + }) + }); + + let iter_conditions = || { + transitions + .iter_of_conditions() + .map(|(_, trans_id)| transitions.get_condition(trans_id, ctx).unwrap()) + }; + + // If desired, see `with_condition_out_transitions` flag: + // Supports from (non-goal) conditions to other conditions and to (non-initial effects). + + let supports_from_conditions = with_condition_out_transitions + .then(|| { + iter_conditions().flat_map(move |(out_cond_id, _)| { + let out_trans_id = transitions.of_condition(out_cond_id).unwrap(); + + // out-transition must not correspond to a goal (i.e. initial) condition + // and must not be a cond-eff (these are already included in the original causal link supports) + (transitions.get_source(out_trans_id, ctx).is_some() + && transitions.get(out_trans_id).tpe() != TransitionType::CondEff) + .then(|| { + let to_effects = iter_effects_info().flat_map(move |(in_eff_id, _)| { + let in_trans_id = transitions.of_effect(in_eff_id).unwrap(); + + if out_trans_id == in_trans_id + || transitions.get(in_trans_id).tpe() == TransitionType::CondEff + || transitions.get_source(in_trans_id, ctx).is_none() // in-transition must not correspond to an initial effect + || transitions.get_state_var(out_trans_id, ctx).fluent + != transitions.get_state_var(in_trans_id, ctx).fluent + { + return None; + } + debug_assert!( + transitions.get(in_trans_id).tpe() == TransitionType::Eff + && transitions.get_source(in_trans_id, ctx).is_some() + ); + + let out_terms = transitions.get_terms(out_trans_id, ctx); + let in_terms = transitions.get_terms(in_trans_id, ctx); + + let incompatible_args = + out_terms.args().iter().zip(in_terms.args()).any(|(out_term, in_term)| { + out_term.is_cst() && in_term.is_cst() && out_term.constant != in_term.constant + }); + + (!incompatible_args).then_some(((out_trans_id, in_trans_id), None::)) + }); + let to_conditions = iter_conditions().flat_map(move |(in_cond_id, _)| { + let in_trans_id = transitions.of_condition(in_cond_id).unwrap(); + + if out_trans_id == in_trans_id + || transitions.get_state_var(out_trans_id, ctx).fluent + != transitions.get_state_var(in_trans_id, ctx).fluent + { + return None; + } + + let out_terms = transitions.get_terms(out_trans_id, ctx); + let in_terms = transitions.get_terms(in_trans_id, ctx); + + let incompatible_args = + out_terms.args().iter().zip(in_terms.args()).any(|(out_term, in_term)| { + out_term.is_cst() && in_term.is_cst() && out_term.constant != in_term.constant + }); + + let incompatible_vals = match (out_terms.val(), in_terms.val()) { + (Some(out_val), Some(in_val)) => { + out_val.is_cst() && in_val.is_cst() && out_val.constant != in_val.constant + } + _ => false, + }; + + (!incompatible_args && !incompatible_vals) + .then_some(((out_trans_id, in_trans_id), None::)) + }); + + to_effects.chain(to_conditions) + }) + .into_iter() + .flatten() + }) + }) + .into_iter() + .flatten(); + + let unsorted_out = (supports_causal_links + .chain(supports_from_effects_to_others) + .chain(supports_from_conditions)) + .filter(|&((out_trans_id, in_trans_id), _)| out_trans_id != in_trans_id) + .inspect(|&((out_trans_id, in_trans_id), _)| { + debug_assert!( + transitions.get(out_trans_id).tpe() != TransitionType::Cond + || transitions.get_source(out_trans_id, ctx).is_some() + ); + debug_assert!( + transitions.get(out_trans_id).tpe() != TransitionType::CondEff + || transitions.get_source(out_trans_id, ctx).is_some() + ); + debug_assert!( + transitions.get(in_trans_id).tpe() != TransitionType::Eff + || transitions.get_source(in_trans_id, ctx).is_some() + ); + }) + .collect(); + + Self { + unsorted_out, + with_condition_out_transitions, + } + } + + pub fn unsorted_out(&self) -> &[((TransitionId, TransitionId), Option)] { + &self.unsorted_out + } + + pub fn sort(&self) -> SupportsSorted { + let mut sorted_out = Vec::with_capacity(self.unsorted_out.len()); + let mut sorted_in = Vec::with_capacity(self.unsorted_out.len()); + + for &((out_trans_id, in_trans_id), active) in &self.unsorted_out { + sorted_out.push(((out_trans_id, in_trans_id), active)); + sorted_in.push((in_trans_id, out_trans_id)); + } + sorted_out.sort_unstable(); + sorted_in.sort_unstable(); + + SupportsSorted { sorted_out, sorted_in } + } +} + +#[derive(Clone, Default)] +pub(crate) struct SupportsSorted { + sorted_out: Vec<((TransitionId, TransitionId), Option)>, + sorted_in: Vec<(TransitionId, TransitionId)>, +} + +impl SupportsSorted { + pub fn sorted_out(&self) -> &[((TransitionId, TransitionId), Option)] { + &self.sorted_out + } + pub fn sorted_in(&self) -> &[(TransitionId, TransitionId)] { + &self.sorted_in + } +} + +#[cfg(test)] +mod tests { + use aries_solver::prelude::Lit; + use itertools::Itertools; + + use crate::analysis::transitions::tests::visitall::{VisitAllLine, build_and_encode_visitall_line}; + + use crate::analysis::transitions::{Transition, TransitionId, TransitionType, Transitions, supports::Supports}; + use crate::encoder::{CausalLink, SchedEncoder}; + + #[test] + fn test_supports() { + let encoder = build_and_encode_visitall_line( + &VisitAllLine { + num_locs: 5, + num_moves: 4, + }, + false, + ); + + let transitions = Transitions::new_unambiguous(&encoder, true); + + let supports_without_conditions_out_transitions = Supports::from(&transitions, &encoder, false); + for &((out_trans_id, in_trans_id), active) in supports_without_conditions_out_transitions.unsorted_out() { + test_supports_aux(&transitions, &encoder, out_trans_id, in_trans_id, active); + } + + let supports_with_conditions_out_transitions = Supports::from(&transitions, &encoder, true); + for &((out_trans_id, in_trans_id), active) in supports_with_conditions_out_transitions.unsorted_out() { + test_supports_aux(&transitions, &encoder, out_trans_id, in_trans_id, active); + } + } + + fn test_supports_aux( + transitions: &Transitions, + encoder: &SchedEncoder, + out_trans_id: TransitionId, + in_trans_id: TransitionId, + active: Option, + ) { + if let Some(active) = active { + assert!(transitions.get(in_trans_id).tpe() != TransitionType::Eff); + + assert!({ + let (eff_id, cond_id) = match (transitions.get(out_trans_id), transitions.get(in_trans_id)) { + ( + Transition::Eff(eff_id) | Transition::CondEff(_, eff_id), + Transition::Cond(cond_id) | Transition::CondEff(cond_id, _), + ) => (eff_id, cond_id), + _ => unreachable!(), + }; + encoder.causal_links.get_links().contains(&CausalLink { + eff_id, + cond_id, + active, + }) + }); + } else { + assert!(!matches!( + (transitions.get(out_trans_id), transitions.get(in_trans_id),), + ( + Transition::CondEff(_, _) | Transition::Eff(_), + Transition::Cond(_) | Transition::CondEff(_, _), + ) + )); + + assert!( + transitions + .get_effect_info(in_trans_id, encoder) + .is_none_or(|(in_eff_id, _)| !transitions.is_recovered_closed_world_default(in_eff_id)) + ); + + if transitions + .get_effect_info(out_trans_id, encoder) + .is_some_and(|(out_eff_id, _)| transitions.is_recovered_closed_world_default(out_eff_id)) + { + assert!(transitions.get(out_trans_id).tpe() == TransitionType::Eff); + assert!(transitions.get_source(out_trans_id, encoder).is_none()); + } + } + } +} From 7b2a3a644469a7c53bdccb90900abca5c4241e75 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 17:37:42 +0200 Subject: [PATCH 13/19] feat(timelines): add lprelax problem --- Cargo.lock | 1 + planning/timelines/Cargo.toml | 1 + planning/timelines/src/analysis/mod.rs | 1 - planning/timelines/src/constraints.rs | 1 + .../src/constraints/lprelax/encoder/encode.rs | 587 +++++++++++++ .../src/constraints/lprelax/encoder/ground.rs | 768 ++++++++++++++++++ .../lprelax/encoder/ground/utils.rs | 225 +++++ .../src/constraints/lprelax/encoder/mod.rs | 213 +++++ .../constraints/lprelax/encoder/problem.rs | 162 ++++ .../timelines/src/constraints/lprelax/mod.rs | 12 + 10 files changed, 1970 insertions(+), 1 deletion(-) create mode 100644 planning/timelines/src/constraints/lprelax/encoder/encode.rs create mode 100644 planning/timelines/src/constraints/lprelax/encoder/ground.rs create mode 100644 planning/timelines/src/constraints/lprelax/encoder/ground/utils.rs create mode 100644 planning/timelines/src/constraints/lprelax/encoder/mod.rs create mode 100644 planning/timelines/src/constraints/lprelax/encoder/problem.rs create mode 100644 planning/timelines/src/constraints/lprelax/mod.rs diff --git a/Cargo.lock b/Cargo.lock index 1d38acac..2043a4cf 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -213,6 +213,7 @@ dependencies = [ "aries-datalog", "aries-env-param", "aries-solver", + "hashbrown 0.16.1", "idmap", "itertools 0.14.0", "num-rational", diff --git a/planning/timelines/Cargo.toml b/planning/timelines/Cargo.toml index 6fdd5aff..73d70175 100644 --- a/planning/timelines/Cargo.toml +++ b/planning/timelines/Cargo.toml @@ -16,3 +16,4 @@ idmap = { workspace = true } smallvec = { workspace = true } num-rational = { workspace = true } streaming-iterator = { workspace = true } +hashbrown = { workspace = true } diff --git a/planning/timelines/src/analysis/mod.rs b/planning/timelines/src/analysis/mod.rs index 0a2ddea2..4bb8521a 100644 --- a/planning/timelines/src/analysis/mod.rs +++ b/planning/timelines/src/analysis/mod.rs @@ -1,6 +1,5 @@ pub mod grounding; mod nonsimple; -#[allow(dead_code)] pub mod transitions; pub use nonsimple::collect_nonsimple_conditions_and_effects_to_relax; diff --git a/planning/timelines/src/constraints.rs b/planning/timelines/src/constraints.rs index 756659cf..69d56a2b 100644 --- a/planning/timelines/src/constraints.rs +++ b/planning/timelines/src/constraints.rs @@ -1,4 +1,5 @@ pub mod grounding; +pub mod lprelax; pub mod symmetry; use aries_solver::lang::ModelView; diff --git a/planning/timelines/src/constraints/lprelax/encoder/encode.rs b/planning/timelines/src/constraints/lprelax/encoder/encode.rs new file mode 100644 index 00000000..507d5ed1 --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/encoder/encode.rs @@ -0,0 +1,587 @@ +use aries_solver::core::views::Dom; +use itertools::Itertools; + +use crate::SchedEncoder; +use crate::analysis::transitions::TransitionType; +use crate::constraints::lprelax::LpRelaxEncoder; +use crate::constraints::lprelax::encoder::problem::{ColTag, LpRelaxProblem, RowExpr, RowExprType}; + +pub fn encode_problem_lifted(encoder: &LpRelaxEncoder, ctx: &SchedEncoder, problem: &mut LpRelaxProblem) { + // [Lifted] A source is present if(f) its transitions are + // TODO: optimize iterations ? (flattened ?) + { + problem.push_row(RowExpr::new( + RowExprType::Eq, + vec![(1, ColTag::PresenceSource(None, None))], + 1, + 1, + )); + + for (source, trans_ids) in encoder.iter_sources() { + for &trans_id in trans_ids { + // It is possible that sometimes the presence literal of the transition differs from that of the source + // (in particular, as a result of complying with pddl set semantics, if the transition presence literal was replaced with a more specific one than the source's). + // However, the latter must always imply the former. + + debug_assert!(ctx.store.state.implies( + encoder.transitions.get_prez(trans_id, ctx), + encoder.get_source_prez(source, ctx) + )); + let presences_equivalent = ctx.store.state.implies( + encoder.get_source_prez(source, ctx), + encoder.transitions.get_prez(trans_id, ctx), + ); + + // When the transition's and source's presence literals are truly equivalent, we can enforce equality. Otherwise, we can only enforce one side. + let expr = if presences_equivalent { + RowExpr::new_eq_single_lhs(vec![ + (1, ColTag::PresenceTransition(trans_id, None)), + (1, ColTag::PresenceSource(source, None)), + ]) + } else { + RowExpr::new_leq_single_lhs(vec![ + (1, ColTag::PresenceTransition(trans_id, None)), + (1, ColTag::PresenceSource(source, None)), + ]) + }; + + problem.push_row(expr); + } + } + } + + // [Lifted] Support between two transitions implies presence of both of them + // NOTE: There's no need to enforce theses constraints for all cases, as the inflow and outflow constraints are stronger (see below). + // They're only actually needed for out-conditions when the in-transition is a "pure-condition", as this case is not implied by outflow constraints. + { + for &((out_trans_id, in_trans_id), _) in encoder.supports.unsorted_out() { + if !encoder.supports.with_condition_out_transitions + && encoder.transitions.get(in_trans_id).tpe() == TransitionType::Cond + { + problem.push_row(RowExpr::new_leq_single_lhs(vec![ + (1, ColTag::Support(out_trans_id, in_trans_id, None)), + (1, ColTag::PresenceTransition(out_trans_id, None)), + ])); + } + } + } + + // [Lifted] Forbid two transitions from mutually supporting each other ("trivial cycles") + { + let mut seen = vec![]; + + for &((out_trans_id, in_trans_id), _) in encoder.supports_sorted.as_ref().unwrap().sorted_out() { + seen.push((out_trans_id, in_trans_id)); + + if seen.binary_search(&(in_trans_id, out_trans_id)).is_ok() { + problem.push_row(RowExpr::new_leq_1(vec![ + (1, ColTag::Support(out_trans_id, in_trans_id, None)), + (1, ColTag::Support(in_trans_id, out_trans_id, None)), + ])); + } + } + debug_assert!(seen.is_sorted()); + } + + // [Lifted] "Inflow" constraints: [TODO] + { + let chunkby = encoder + .supports_sorted + .as_ref() + .unwrap() + .sorted_in() + .iter() + .chunk_by(|(in_trans_id, _)| in_trans_id); + + for (&in_trans_id, out_trans_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![(1, ColTag::PresenceTransition(in_trans_id, None))]; + res.append( + &mut out_trans_ids + .into_iter() + .map(|&(_, out_trans_id)| (1, ColTag::Support(out_trans_id, in_trans_id, None))) + .collect(), + ); + res + }; + debug_assert!(terms.len() >= 2); + + let expr = if encoder.transitions.are_recovered_closed_world_default_effects_empty() + && encoder.transitions.get(in_trans_id).tpe() == TransitionType::Eff + { + RowExpr::new_geq_single_lhs(terms) + } else { + RowExpr::new_eq_single_lhs(terms) + }; + + problem.push_row(expr); + } + } + + // [Lifted] "Outflow" constraints: [TODO] + { + let chunkby = encoder + .supports_sorted + .as_ref() + .unwrap() + .sorted_out() + .iter() + .chunk_by(|&((out_trans_id, _), _)| out_trans_id); + + for (&out_trans_id, in_trans_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![(1, ColTag::PresenceTransition(out_trans_id, None))]; + res.append( + &mut in_trans_ids + .into_iter() + .filter(|&&((_, in_trans_id), _)| { + encoder.supports.with_condition_out_transitions + || encoder.transitions.get(in_trans_id).tpe() != TransitionType::Cond + }) + .map(|&((_, in_trans_id), _)| (1, ColTag::Support(out_trans_id, in_trans_id, None))) + .collect(), + ); + res + }; + + if terms.len() >= 2 { + let expr = RowExpr::new_geq_single_lhs(terms); + problem.push_row(expr); + } + } + } +} + +pub fn encode_problem_ground(encoder: &LpRelaxEncoder, ctx: &SchedEncoder, problem: &mut LpRelaxProblem) { + // [Lifted-Ground] Source ground decomposition (a source is present iff one its groundings is) + { + for (source, _) in encoder.iter_sources() { + let terms = { + let mut res = vec![(1, ColTag::PresenceSource(source, None))]; + res.append( + &mut encoder + .sources_ground + .get(source) + .iter() + .map(|&source_grounding_id| (1, ColTag::PresenceSource(source, Some(source_grounding_id)))) + .collect(), + ); + res + }; + if terms.len() >= 2 { + let expr = RowExpr::new_eq_single_lhs(terms); + problem.push_row(expr); + } else { + // ? WARNING ? related to incomplete / partial groundings. [TODO] + continue; + } + } + } + + // [Lifted-Ground] Transition ground decomposition + { + let chunkby = encoder + .transitions_ground + .iter_all_unsourced() + .chunk_by(|&(trans_id, _)| trans_id); + + for (trans_id, trans_groundings_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![(1, ColTag::PresenceTransition(trans_id, None))]; + res.append( + &mut trans_groundings_ids + .into_iter() + .map(|(_, trans_grounding_id)| { + (1, ColTag::PresenceTransition(trans_id, Some(trans_grounding_id))) + }) + .collect(), + ); + res + }; + debug_assert!(terms.len() >= 2); + + let expr = RowExpr::new_eq_single_lhs(terms); + problem.push_row(expr); + } + } + + // [Lifted-Ground] Ground transition is only active if(f) a compatible grounding of its source is active + { + for (source, iter) in encoder.transitions_ground.iter_all_sourced() { + let chunkby = iter.chunk_by(|(trans_id, trans_grounding_id, _)| (trans_id, trans_grounding_id)); + + for ((&trans_id, &trans_grounding_id), source_groundings_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![(1, ColTag::PresenceTransition(trans_id, Some(trans_grounding_id)))]; + res.append( + &mut source_groundings_ids + .into_iter() + .map(|&(_, _, source_grounding_id)| { + (1, ColTag::PresenceSource(source, Some(source_grounding_id))) + }) + .collect(), + ); + res + }; + debug_assert!(terms.len() >= 2); + + // Same as for the lifted case: + // It is possible that sometimes the presence literal of the transition differs from that of the source + // (in particular, as a result of complying with pddl set semantics, if the transition presence literal was replaced with a more specific one than the source's). + // However, the latter must always imply the former. + + debug_assert!(ctx.store.state.implies( + encoder.transitions.get_prez(trans_id, ctx), + encoder.get_source_prez(source, ctx) + )); + let presences_equivalent = ctx.store.state.implies( + encoder.get_source_prez(source, ctx), + encoder.transitions.get_prez(trans_id, ctx), + ); + + // When the transition's and source's presence literals are truly equivalent, we can enforce equality. Otherwise, we can only enforce one side. + let expr = if presences_equivalent { + RowExpr::new_eq_single_lhs(terms) + } else { + RowExpr::new_leq_single_lhs(terms) + }; + + problem.push_row(expr); + } + } + } + + // [Ground] Support between two (ground) transitions implies presence of both of them + // NOTE: There's no need to enforce theses constraints for all cases, as the (ground) inflow and outflow constraints are stronger (see below). + // They're only actually needed for out-conditions when the in-transition is a "pure-condition", as this case is not implied by (ground) outflow constraints. + { + for &(out_trans_id, in_trans_id, trans_groundings_ids) in encoder.supports_ground.iter_all() { + let Some((out_trans_grounding_id, in_trans_grounding_id)) = trans_groundings_ids else { + continue; + }; + + if !encoder.supports.with_condition_out_transitions + && encoder.transitions.get(in_trans_id).tpe() == TransitionType::Cond + { + problem.push_row(RowExpr::new_leq_single_lhs(vec![ + ( + 1, + ColTag::Support( + out_trans_id, + in_trans_id, + Some((out_trans_grounding_id, in_trans_grounding_id)), + ), + ), + ( + 1, + ColTag::PresenceTransition(out_trans_id, Some(out_trans_grounding_id)), + ), + ])); + } + } + } + + // [Lifted-Ground] Supports ground decomposition + { + let chunkby = encoder + .supports_ground + .iter_all() + .chunk_by(|(out_trans_id, in_trans_id, _)| (out_trans_id, in_trans_id)); + + for ((&out_trans_id, &in_trans_id), trans_groundings_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![(1, ColTag::Support(out_trans_id, in_trans_id, None))]; + res.append( + &mut trans_groundings_ids + .into_iter() + .filter_map(|&(_, _, trans_groundings_ids)| { + trans_groundings_ids.map(|(out_trans_grounding_id, in_trans_grounding_id)| { + ( + 1, + ColTag::Support( + out_trans_id, + in_trans_id, + Some((out_trans_grounding_id, in_trans_grounding_id)), + ), + ) + }) + }) + .collect(), + ); + res + }; + + let expr = RowExpr::new_eq_single_lhs(terms); + problem.push_row(expr); + } + } + + // [Ground] "Inflow" constraints: [TODO] + { + let chunkby = encoder + .supports_ground + .iter_in_all() + .chunk_by(|(in_trans_id, in_trans_grounding_id, _)| (in_trans_id, in_trans_grounding_id)); + + for ((&in_trans_id, &in_trans_grounding_id), out_trans_groundings_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![(1, ColTag::PresenceTransition(in_trans_id, Some(in_trans_grounding_id)))]; + res.append( + &mut out_trans_groundings_ids + .into_iter() + .filter_map(|&(_, _, x)| { + x.map(|(out_trans_id, out_trans_grounding_id)| { + ( + 1, + ColTag::Support( + out_trans_id, + in_trans_id, + Some((out_trans_grounding_id, in_trans_grounding_id)), + ), + ) + }) + }) + .collect(), + ); + res + }; + + let expr = if encoder.transitions.are_recovered_closed_world_default_effects_empty() + && encoder.transitions.get(in_trans_id).tpe() == TransitionType::Eff + { + // In the case where we do not recover and use "missing" initial effects, + // the inflow constraints for (all) effects are slightly weaker. + RowExpr::new_geq_single_lhs(terms) + } else { + RowExpr::new_eq_single_lhs(terms) + }; + + problem.push_row(expr); + } + + if encoder.transitions.are_recovered_closed_world_default_effects_empty() { + // In the case where we do not recover and use "missing" initial effects, + // the inflow constraints for (all) effects are slightly weaker (see above). + // This is (partially? FIXME[proof?]) compensated by the following constraints, + // which state that the *sum* of inflows into (ground effects) with the *same state variable* (so, independent of their value) is upper bounded by 1. + + let chunkby = encoder + .supports_ground + .iter_in_all() + .filter(|(in_trans_id, _, _)| encoder.transitions.get(*in_trans_id).tpe() == TransitionType::Eff) + .map(|(in_trans_id, in_trans_grounding_id, out_trans_groundings)| { + (in_trans_grounding_id, in_trans_id, out_trans_groundings) + }) + .sorted_unstable_by_key(|&(in_trans_grounding_id, _, _)| *in_trans_grounding_id) + .chunk_by(|&(in_trans_grounding_id, _, _)| in_trans_grounding_id.state_var_grounding_id); + + for (_, x) in chunkby.into_iter() { + let terms = x + .into_iter() + .filter_map(|(&in_trans_grounding_id, &in_trans_id, out_trans_grounding)| { + out_trans_grounding.map(|(out_trans_id, out_trans_grounding_id)| { + ( + 1, + ColTag::Support( + out_trans_id, + in_trans_id, + Some((out_trans_grounding_id, in_trans_grounding_id)), + ), + ) + }) + }) + .collect::>(); + + if !terms.is_empty() { + let expr = RowExpr::new_leq_1(terms); + problem.push_row(expr); + } + } + } + } + + // [Ground] "Outflow" constraints: [TODO] + { + let chunkby = encoder + .supports_ground + .iter_out_all() + .chunk_by(|(out_trans_id, out_trans_grounding_id, _)| (out_trans_id, out_trans_grounding_id)); + + for ((&out_trans_id, &out_trans_grounding_id), in_trans_groundings_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![( + 1, + ColTag::PresenceTransition(out_trans_id, Some(out_trans_grounding_id)), + )]; + res.append( + &mut in_trans_groundings_ids + .into_iter() + .filter(|&&(_, _, (in_trans_id, _))| { + encoder.supports.with_condition_out_transitions + || encoder.transitions.get(in_trans_id).tpe() != TransitionType::Cond + }) + .map(|&(_, _, (in_trans_id, in_trans_grounding_id))| { + ( + 1, + ColTag::Support( + out_trans_id, + in_trans_id, + Some((out_trans_grounding_id, in_trans_grounding_id)), + ), + ) + }) + .collect(), + ); + res + }; + + if terms.len() >= 2 { + let expr = RowExpr::new_geq_single_lhs(terms); + problem.push_row(expr); + } + } + } + + // // [Ground] Forbid one (ground) transitions from mutually supporting each other + // // TODO: ? is this actually needed ? -> (this may not necessarily be useful / enough for cases of eff-eff supports, as any value could be usef (or eff-condeff)) + // if super::ARIES_LPRELAX_GROUND_2CYCLES.get() { + // let ground_supports_iter_sorted_fully = groundings.supports().iter_all().sorted(); + // let mut seen = vec![]; + // + // for &(out_trans_id, in_trans_id, trans_groundings_ids) in ground_supports_iter_sorted_fully { + // let Some((out_trans_grounding_id, in_trans_grounding_id)) = trans_groundings_ids else { + // continue; + // }; + // + // seen.push(( + // out_trans_id, + // in_trans_id, + // out_trans_grounding_id, + // in_trans_grounding_id, + // )); + // + // if seen + // .binary_search(&( + // in_trans_id, + // out_trans_id, + // in_trans_grounding_id, + // out_trans_grounding_id, + // )) + // .is_ok() + // { + // let expr = RowExpr::Leq1(vec![ + // ColTag::SupportGround( + // out_trans_id, + // in_trans_id, + // out_trans_grounding_id, + // in_trans_grounding_id, + // ), + // ColTag::SupportGround( + // in_trans_id, + // out_trans_id, + // in_trans_grounding_id, + // out_trans_grounding_id, + // ), + // ]); + // problem.push_row(expr); + // } + // } + // debug_assert!(seen.is_sorted()); + // } + + // [Ground] At most one of a term's groundings can be active + { + let chunkby = encoder + .terms_ground + .iter_sorted_all_only_assignments() + .chunk_by(|&(term, _)| term); + + for (term, values) in chunkby.into_iter() { + let terms = values + .into_iter() + .map(|(_, value)| (1, ColTag::TermGround(term, value))) + .collect::>(); + + debug_assert!(!terms.is_empty()); + + let expr = RowExpr::new_leq_1(terms); + problem.push_row(expr); + } + } + + // [Ground] A grounding of a term is active iff a ground transition using it is active + // [Ground] ---------------------------------------------- source -------------------- + { + let chunkby = encoder + .terms_ground + .iter_sorted_all_for_transitions() + .chunk_by(|&(term, value, trans_id, _)| (term, value, trans_id)); + + for ((&term, &value, &trans_id), trans_groundings_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![(1, ColTag::TermGround(term, value))]; + res.append( + &mut trans_groundings_ids + .into_iter() + .map(|&(_, _, _, trans_grounding_id)| { + (1, ColTag::PresenceTransition(trans_id, Some(trans_grounding_id))) + }) + .collect(), + ); + res + }; + debug_assert!(terms.len() >= 2); + + debug_assert!( + ctx.store + .state + .implies(encoder.transitions.get_prez(trans_id, ctx), ctx.store.presence(term)) + ); + let presences_equivalent = ctx + .store + .state + .implies(ctx.store.presence(term), encoder.transitions.get_prez(trans_id, ctx)); + + let expr = if presences_equivalent { + RowExpr::new_eq_single_lhs(terms) + } else { + RowExpr::new_geq_single_lhs(terms) + }; + problem.push_row(expr); + } + + let chunkby = encoder + .terms_ground + .iter_sorted_all_for_sources() + .chunk_by(|&(term, value, source, _)| (term, value, source)); + + for ((&term, &value, &source), sources_groundings_ids) in chunkby.into_iter() { + let terms = { + let mut res = vec![(1, ColTag::TermGround(term, value))]; + res.append( + &mut sources_groundings_ids + .into_iter() + .map(|&(_, _, _, source_grounding_id)| { + (1, ColTag::PresenceSource(source, Some(source_grounding_id))) + }) + .collect(), + ); + res + }; + debug_assert!(terms.len() >= 2); + + debug_assert!( + ctx.store + .state + .implies(encoder.get_source_prez(source, ctx), ctx.store.presence(term)) + && ctx + .store + .state + .implies(ctx.store.presence(term), encoder.get_source_prez(source, ctx)) + ); + + let expr = RowExpr::new_eq_single_lhs(terms); + problem.push_row(expr); + } + } +} diff --git a/planning/timelines/src/constraints/lprelax/encoder/ground.rs b/planning/timelines/src/constraints/lprelax/encoder/ground.rs new file mode 100644 index 00000000..21d6e0e2 --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/encoder/ground.rs @@ -0,0 +1,768 @@ +mod utils; + +use utils::{binary_search_range_by, merge_dedup_into}; + +use std::collections::HashMap; + +use aries_solver::core::IntCst; +use idmap::DirectIdMap; +use itertools::Itertools; + +use crate::IntTerm; +use crate::analysis::Source; +use crate::analysis::transitions::TransitionId; +use crate::analysis::transitions::supports::SupportsSorted; +use utils::merge_join_chunks_by_key; + +pub type SourceGrounding = crate::analysis::grounding::ParametersAssignment; +pub type SourceGroundingId = usize; + +#[derive(Clone, Default)] +pub(crate) struct SourcesGroundingsInfo { + /// First element: list of (unique) groundings of the "empty source" (i.e. groundings corresponding to the "initial/final" action) + /// Second element: list of (unique) groundings of a "concrete source" (action/task) + entries: ( + Vec, + DirectIdMap>, + ), + /// Index: id of a source grounding + groundings_rev: Vec<(Source, SourceGrounding)>, + + #[cfg(debug_assertions)] + groundings: HashMap<(Source, SourceGrounding), SourceGroundingId>, +} +impl SourcesGroundingsInfo { + /// Interns a grounding of a source. + /// WARNING: adding duplicate groundings (for the same source) will result in a panic (in debug mode). + /// + /// Does not do anything relating to ground supports. + pub fn post_ground_source(&mut self, source: Source, source_grounding: &SourceGrounding) -> SourceGroundingId { + #[cfg(debug_assertions)] + debug_assert!(!self.groundings.contains_key(&(source, source_grounding.clone()))); + + let source_grounding_id = self.groundings_rev.len(); + + #[cfg(debug_assertions)] + { + self.groundings + .insert((source, source_grounding.clone()), source_grounding_id); + } + self.groundings_rev.push((source, source_grounding.clone())); + + if let Some(task_id) = source { + if !self.entries.1.contains_key(task_id) { + self.entries.1.insert(task_id, vec![]); + } + self.entries.1[task_id].push(source_grounding_id); + } else { + self.entries.0.push(source_grounding_id); + } + source_grounding_id + } + + pub fn get(&self, source: Source) -> &[SourceGroundingId] { + if let Some(task_id) = source { + debug_assert!( + self.entries + .1 + .get(task_id) + .is_none_or(|entries| entries.iter().all_unique()) + ); + self.entries + .1 + .get(task_id) + .map(|entries| entries.as_slice()) + .unwrap_or_default() + } else { + debug_assert!(self.entries.0.iter().all_unique()); + &self.entries.0 + } + } +} + +type FluentId = usize; +type StateVarGrounding = smallvec::SmallVec<[IntCst; 4]>; +pub type StateVarGroundingId = usize; + +struct StateVarGroundingLookup<'a> { + fluent_id: FluentId, + state_var_grounding: &'a [IntCst], +} +impl<'a> hashbrown::Equivalent<(FluentId, StateVarGrounding)> for StateVarGroundingLookup<'a> { + fn equivalent(&self, key: &(FluentId, StateVarGrounding)) -> bool { + self.fluent_id == key.0 && self.state_var_grounding == key.1.as_slice() + } +} +impl<'a> std::hash::Hash for StateVarGroundingLookup<'a> { + fn hash(&self, state: &mut H) { + self.fluent_id.hash(state); + self.state_var_grounding.hash(state); + } +} + +#[derive(Clone, Default)] +pub(crate) struct StateVarsGroundingsInfo { + fluents_ids: HashMap, + /// Stores ids of state variable groundings (they are used in transition grounding ids as the first element of the triple) + entries: hashbrown::HashMap<(FluentId, StateVarGrounding), StateVarGroundingId>, + + #[cfg(debug_assertions)] + groundings_rev: Vec<(FluentId, StateVarGrounding)>, +} +impl StateVarsGroundingsInfo { + /// NOTE: will *not* panic if an already known grounding is given (unlike when interning source groundings). + /// To the contrary, this is used to retrieve the id of the given state variable grounding, if it was interned. + pub fn post_ground_state_var( + &mut self, + fluent: &crate::Sym, + state_var_grounding: &[IntCst], + ) -> StateVarGroundingId { + let fluent_id = if self.fluents_ids.contains_key(fluent) { + *self.fluents_ids.get(fluent).unwrap() + } else { + let fluent_id = self.fluents_ids.len(); + self.fluents_ids.insert(fluent.into(), fluent_id); + fluent_id + }; + + let lookup = StateVarGroundingLookup { + fluent_id, + state_var_grounding, + }; + + if let Some(&state_var_grounding_id) = self.entries.get(&lookup) { + state_var_grounding_id + } else { + let state_var_grounding_id = self.entries.len(); + self.entries + .insert((fluent_id, state_var_grounding.into()), state_var_grounding_id); + #[cfg(debug_assertions)] + { + debug_assert!(state_var_grounding_id == self.groundings_rev.len()); + self.groundings_rev.push((fluent_id, state_var_grounding.into())); + } + state_var_grounding_id + } + } +} + +#[derive(Clone, Default)] +pub(crate) struct TermsGroundingsInfo { + /// Flat storage of term assignments with the ground transitions in which they appear. + entries_merged_transitions: Vec<(IntTerm, IntCst, TransitionId, TransitionGroundingId)>, + /// Same as for transitions, but for ground sources. + entries_merged_sources: Vec<(IntTerm, IntCst, Source, SourceGroundingId)>, + + /// A buffer into which the entries are pushed before being deduped, sorted, and merged into the + /// corresponding 'merged' field (see above). + /// This could allow a more efficient incremental interning of new entries, + /// via interleaved sorting then merging, rather than global merging every time. + entries_pending_transitions: Vec<(IntTerm, IntCst, TransitionId, TransitionGroundingId)>, + /// Same as for transitions. + entries_pending_sources: Vec<(IntTerm, IntCst, Source, SourceGroundingId)>, +} +impl TermsGroundingsInfo { + // Interns a grounding of a source. + // WARNING: adding duplicate entries will result in a panic (in debug mode). + // + // Does not do anything relating to ground supports. + pub fn post_for_ground_transition( + &mut self, + term: IntTerm, + value: IntCst, + trans_id: TransitionId, + trans_grounding_id: TransitionGroundingId, + ) { + assert!(!term.is_cst()); + debug_assert!( + !self + .entries_merged_transitions + .contains(&(term, value, trans_id, trans_grounding_id)) + ); + + self.entries_pending_transitions + .push((term, value, trans_id, trans_grounding_id)); + } + pub fn post_for_ground_source( + &mut self, + term: IntTerm, + value: IntCst, + source: Source, + source_grounding_id: SourceGroundingId, + ) { + assert!(!term.is_cst()); + debug_assert!( + !self + .entries_merged_sources + .contains(&(term, value, source, source_grounding_id)) + ); + + self.entries_pending_sources + .push((term, value, source, source_grounding_id)); + } + + fn cmp_trs( + x: &(IntTerm, IntCst, TransitionId, TransitionGroundingId), + y: &(IntTerm, IntCst, TransitionId, TransitionGroundingId), + ) -> std::cmp::Ordering { + x.cmp(y) + } + fn cmp_src( + x: &(IntTerm, IntCst, Source, SourceGroundingId), + y: &(IntTerm, IntCst, Source, SourceGroundingId), + ) -> std::cmp::Ordering { + x.cmp(y) + } + + pub fn sort(&mut self) { + self.sort_and_merge_pending_transitions(); + self.sort_and_merge_pending_sources(); + } + + fn sort_and_merge_pending_transitions(&mut self) { + self.entries_pending_transitions.sort_unstable(); + self.entries_pending_transitions.dedup(); + merge_dedup_into( + &mut self.entries_pending_transitions, + &mut self.entries_merged_transitions, + Self::cmp_trs, + ); + debug_assert!(self.entries_pending_transitions.is_empty()); + debug_assert!(self.entries_merged_transitions.is_sorted()); + } + fn sort_and_merge_pending_sources(&mut self) { + self.entries_pending_sources.sort_unstable(); + self.entries_pending_sources.dedup(); + merge_dedup_into( + &mut self.entries_pending_sources, + &mut self.entries_merged_sources, + Self::cmp_src, + ); + debug_assert!(self.entries_pending_sources.is_empty()); + debug_assert!(self.entries_merged_sources.is_sorted()); + } + + fn debug_check_valid_for_transitions(&self) -> bool { + debug_assert!(self.entries_pending_transitions.is_empty()); + debug_assert!(is_sorted_and_no_dupes(self.entries_merged_transitions.iter())); + true + } + fn debug_check_valid_for_sources(&self) -> bool { + debug_assert!(self.entries_pending_sources.is_empty()); + debug_assert!(is_sorted_and_no_dupes(self.entries_merged_sources.iter())); + true + } + + /// Sorted, meaning the iterator can be chunked (without extra allocations). + pub fn iter_sorted_all_for_transitions( + &self, + ) -> impl Iterator { + debug_assert!(self.debug_check_valid_for_transitions()); + self.entries_merged_transitions.iter() + } + /// Sorted, meaning the iterator can be chunked (without extra allocations). + pub fn iter_sorted_all_for_sources(&self) -> impl Iterator { + debug_assert!(self.debug_check_valid_for_sources()); + self.entries_merged_sources.iter() + } + /// Sorted, meaning the iterator can be chunked (without extra allocations). + pub fn iter_sorted_all_only_assignments(&self) -> impl Iterator { + debug_assert!(self.debug_check_valid_for_sources()); + self.entries_merged_sources + .iter() + .map(|(term, value, _, _)| (*term, *value)) + .dedup() + } +} + +#[derive(Debug, Copy, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct TransitionGroundingId { + pub(super) state_var_grounding_id: usize, + pub(super) val_assignment: Option, + pub(super) op_assignment: Option, +} +impl TransitionGroundingId { + fn is_pure_eff(&self) -> bool { + self.op_assignment.is_some() && self.val_assignment.is_none() + } +} +#[derive(Clone, Default)] +pub(crate) struct TransitionsGroundingsInfo { + sources_of: DirectIdMap, + + /// Sorted + entries_sourced_empty_source: Vec<(TransitionId, TransitionGroundingId, SourceGroundingId)>, + /// Sorted + entries_sourced_concrete_sources: + DirectIdMap>, +} +impl TransitionsGroundingsInfo { + pub fn post_ground_transition( + &mut self, + trans_id: TransitionId, + trans_grounding_id: TransitionGroundingId, + source: Source, + source_grounding_id: SourceGroundingId, + ) { + if !self.sources_of.contains_key(trans_id) { + self.sources_of.insert(trans_id, source); + } else { + assert!(self.sources_of[trans_id] == source); + } + + if let Some(task_id) = source { + if !self.entries_sourced_concrete_sources.contains_key(task_id) { + self.entries_sourced_concrete_sources.insert(task_id, vec![]); + } + debug_assert!(!self.entries_sourced_concrete_sources[task_id].contains(&( + trans_id, + trans_grounding_id, + source_grounding_id + ))); + self.entries_sourced_concrete_sources[task_id].push((trans_id, trans_grounding_id, source_grounding_id)); + } else { + debug_assert!(!self.entries_sourced_empty_source.contains(&( + trans_id, + trans_grounding_id, + source_grounding_id + ))); + self.entries_sourced_empty_source + .push((trans_id, trans_grounding_id, source_grounding_id)); + } + } + + fn debug_check_valid(&self) -> bool { + debug_assert!(is_sorted_and_no_dupes(self.entries_sourced_empty_source.iter())); + debug_assert!(is_sorted_and_no_dupes(self.entries_sourced_concrete_sources.iter())); + true + } + + // fn sort_for(&mut self, source: Source) { + // if let Some(task_id) = source { + // if self.entries_sourced_concrete_sources.contains_key(task_id) { + // self.entries_sourced_concrete_sources[task_id].sort_unstable(); + // } + // } else { + // self.entries_sourced_empty_source + // .sort_unstable(); + // } + // } + pub fn sort_for_all(&mut self) { + self.entries_sourced_empty_source.sort_unstable(); + for (_, entries) in self.entries_sourced_concrete_sources.iter_mut() { + entries.sort_unstable(); + } + } + + fn source_of(&self, trans_id: TransitionId) -> Source { + self.sources_of[trans_id] + } + + pub fn iter_all_sourced( + &self, + ) -> impl Iterator< + Item = ( + Source, + impl Iterator, + ), + > { + debug_assert!(self.debug_check_valid()); + std::iter::chain( + [(None, self.entries_sourced_empty_source.iter())], + self.entries_sourced_concrete_sources + .iter() + .map(|(source, entries)| (Some(source), entries.iter())), + ) + } + + /// WARNING: although the the iterator can be chunked (without extra allocations) by transition, it may not necessarily be *sorted* !! + pub fn iter_all_unsourced(&self) -> impl Iterator { + debug_assert!(self.debug_check_valid()); + // result is not necessarily sorted, but chunked because a transition is "owned" only by a single source + std::iter::chain( + self.entries_sourced_empty_source + .iter() + .map(|(trans_id, trans_grounding_id, _)| (*trans_id, *trans_grounding_id)), + self.entries_sourced_concrete_sources.iter().flat_map(|(_, entries)| { + entries + .iter() + .map(|(trans_id, trans_grounding_id, _)| (*trans_id, *trans_grounding_id)) + }), + ) + .dedup() + } + + /// Returns the (sub)slice with the groundings of the given transition (if there are any) + /// Note that the groundings in the slice are ordered (notably meaning that it is chunkable) + /// + /// If the slice cannot be found, returns an empty one. + fn get_index_and_slice( + &self, + trans_id: TransitionId, + ) -> &[(TransitionId, TransitionGroundingId, SourceGroundingId)] { + debug_assert!(self.debug_check_valid()); + + let Some(&source) = self.sources_of.get(trans_id) else { + return [].as_slice(); + }; + let entries = if let Some(task_id) = source { + &self.entries_sourced_concrete_sources[task_id] + } else { + &self.entries_sourced_empty_source + }; + binary_search_range_by(entries, |(trans_id_, _, _)| (*trans_id_).cmp(&trans_id)) + .map_or_else(|| [].as_slice(), |(i, j)| &entries[i..j]) + } +} + +#[derive(Clone, Default)] +pub(crate) struct SupportsGroundingsInfo { + /// Sorted flat storage of ground supports: "outgoing view": + /// (out_trans_id, out_trans_grounding_id, (in_trans_id, in_trans_grounding_id)). + out: Vec<( + TransitionId, + TransitionGroundingId, + (TransitionId, TransitionGroundingId), + )>, + /// Sorted flat storage of ground supports: "incoming view": + /// (in_trans_id, in_trans_grounding_id, (out_trans_id, out_trans_grounding_id)). + #[allow(clippy::type_complexity)] + in_: Vec<( + TransitionId, + TransitionGroundingId, + Option<(TransitionId, TransitionGroundingId)>, + )>, + /// Sorted (over the first 2 element) flat storage of ground supports: "neutral view". + /// (out_trans_id, in_trans_id, (out_trans_grounding_id, in_trans_grounding_id)). + #[allow(clippy::type_complexity)] + entries: Vec<( + TransitionId, + TransitionId, + Option<(TransitionGroundingId, TransitionGroundingId)>, + )>, +} + +impl SupportsGroundingsInfo { + pub fn from(supports_sorted: &SupportsSorted, transitions_groundings: &TransitionsGroundingsInfo) -> Self { + debug_assert!(transitions_groundings.debug_check_valid()); + + let mut out = vec![]; + let mut in_ = vec![]; + let mut entries = vec![]; + + // (Pre-seeding) Acknowledge that an in-transition and a grounding of it exist. + // Later, after (in_trans_id, in_trans_grounding_id, Some(..)) entries are added, + // the pre-seeded entry (in_trans_id, in_trans_grounding_id, None) will be removed. + // If no (in_trans_id, in_trans_grounding_id, Some(..)) entries are added, + // then the (in_trans_id, in_trans_grounding_id, None) is kept and signifies + // that although the in-transition and the grounding exist, there exist no compatible out-transitions to support it. + in_.extend( + transitions_groundings + .iter_all_unsourced() + .filter_map(|(trans_id, trans_grounding_id)| { + let trans_is_non_initial_eff = + !trans_grounding_id.is_pure_eff() || transitions_groundings.source_of(trans_id).is_some(); + trans_is_non_initial_eff.then_some((trans_id, trans_grounding_id, None)) + }), + ); + + // Main loop + for &((out_trans_id, in_trans_id), _) in supports_sorted.sorted_out() { + let out_slice = transitions_groundings.get_index_and_slice(out_trans_id); + let in_slice = transitions_groundings.get_index_and_slice(in_trans_id); + + let entries_len_before_search = entries.len(); + + let mut push_supports = |out_trans_grounding_id: TransitionGroundingId, + in_trans_grounding_id: TransitionGroundingId| { + out.push(( + out_trans_id, + out_trans_grounding_id, + (in_trans_id, in_trans_grounding_id), + )); + in_.push(( + in_trans_id, + in_trans_grounding_id, + Some((out_trans_id, out_trans_grounding_id)), + )); + entries.push(( + out_trans_id, + in_trans_id, + Some((out_trans_grounding_id, in_trans_grounding_id)), + )); + }; + + // debug_assert!( + // out_slice + // .iter() + // .all(|&(trans_id, _, _)| trans_id == out_trans_id), + // ); + // debug_assert!( + // in_slice + // .iter() + // .all(|&(trans_id, _, _)| trans_id == in_trans_id) + // ); + + 'search: { + let in_transition_is_eff = in_slice + .first() + .is_some_and(|(_, in_trans_grounding_id, _)| in_trans_grounding_id.is_pure_eff()); + + let out_direct: Vec = out_slice + .iter() + .map(|&(_, trans_grounding_id, _)| trans_grounding_id) + .dedup() + .collect_vec(); + debug_assert!(is_sorted_and_no_dupes(out_direct.iter())); + + let in_direct: Vec = in_slice + .iter() + .map(|&(_, trans_grounding_id, _)| trans_grounding_id) + .dedup() + .collect_vec(); + debug_assert!(is_sorted_and_no_dupes(in_direct.iter())); + + // Case: in-transition ("consumer" of support) is an pure-effect (i.e. `val_assignment` is None for all its groundings). + // Any other (non-pure-cond) ground transition (with the same state var grounding) can support it, whatever its `op_assignment` + + if in_transition_is_eff { + for (out_chunk_same_sv, in_chunk_same_sv) in + merge_join_chunks_by_key(&out_direct, &in_direct, |trans_grounding_id| { + trans_grounding_id.state_var_grounding_id + }) + { + for (&out_trans_grounding_id, &in_trans_grounding_id) in + out_chunk_same_sv.iter().cartesian_product(in_chunk_same_sv.iter()) + { + push_supports(out_trans_grounding_id, in_trans_grounding_id); + } + } + break 'search; + } + + // Case: in-transition ("consumer" of support) is *not* a pure-effect (i.e. val_assignment is not None for all its groundings). + // it can be support by other (non-pure-cond) ground transitions (with the same state var grounding) and out_op_assignment = `in_val_assignment + + debug_assert!(!in_transition_is_eff); + + let in_direct: Vec<(StateVarGroundingId, Option, Option)> = in_direct + .into_iter() + .map(|trans_grounding_id| { + ( + trans_grounding_id.state_var_grounding_id, + trans_grounding_id.val_assignment, + trans_grounding_id.op_assignment, + ) + }) + .collect_vec(); + debug_assert!(is_sorted_and_no_dupes(in_direct.iter())); + + let out_inverted: Vec<(StateVarGroundingId, Option, Option)> = out_direct + .into_iter() + .map(|trans_grounding_id| { + ( + trans_grounding_id.state_var_grounding_id, + trans_grounding_id.op_assignment, + trans_grounding_id.val_assignment, + ) + }) + .sorted_unstable() + .collect_vec(); + debug_assert!(is_sorted_and_no_dupes(out_inverted.iter())); + + for (out_inv_chunk_same_sv, in_chunk_same_sv) in + merge_join_chunks_by_key(&out_inverted, &in_direct, |&(state_var_grounding_id, _, _)| { + state_var_grounding_id + }) + { + // Within a same-state-var chunk, `out_inverted` is sorted by `op_assignment` and `in_direct` by + // `val_assignment, so the same join applies to the values. + for (out_inv_chunk_matching, in_chunk_matching) in merge_join_chunks_by_key( + out_inv_chunk_same_sv, + in_chunk_same_sv, + |&(_, value, _)| { + value.expect( + "'out' side: `value` is 'out_op_assignment' (never None as an out transition is never a pure condition). + 'in' side: `value` is 'in_val_assignment' (never None as the pure-effect case was already handled above)." + ) + }, + ) { + for (&out_inverted_trans_grounding_id, &in_trans_grounding_id) in out_inv_chunk_matching + .iter() + .cartesian_product(in_chunk_matching.iter()) + { + let in_trans_grounding_id = TransitionGroundingId { + state_var_grounding_id: in_trans_grounding_id.0, + val_assignment: in_trans_grounding_id.1, + op_assignment: in_trans_grounding_id.2, + }; + // un-invert + let out_trans_grounding_id = TransitionGroundingId { + state_var_grounding_id: out_inverted_trans_grounding_id.0, + val_assignment: out_inverted_trans_grounding_id.2, + op_assignment: out_inverted_trans_grounding_id.1, + }; + + push_supports(out_trans_grounding_id, in_trans_grounding_id); + } + } + } + } + + if entries.len() == entries_len_before_search { + // Nothing was added: it means there no compatible groundings were found + entries.push((out_trans_id, in_trans_id, None)); + } + } + + debug_assert!(entries.is_sorted_by_key(|(out_trans_id, in_trans_id, _)| { (out_trans_id, in_trans_id) })); + debug_assert!( + // an entry with `None` exists iff there are no entries with Some (for the same (out_trans_id, in_trans_id) pair) + entries + .chunk_by(|a, b| (a.0, a.1) == (b.0, b.1)) + .all(|chunk| chunk.len() == 1 || chunk.iter().all(|e| e.2.is_some())) + ); + debug_assert!(entries.iter().all_unique()); + + debug_assert!(out.iter().all_unique()); + out.sort_unstable(); + + debug_assert!(in_.iter().all_unique()); + in_.sort_unstable(); + + // Remove pre-seeded (in_trans_id, in_trans_grounding_id, None) entries (see beginning) + // if (in_trans_id, in_trans_grounding_id, Some(..)) entries were found / added during the search. + in_.dedup_by(|next, prev| { + if prev.2.is_none() && (prev.0, prev.1) == (next.0, next.1) { + *prev = *next; // overwrite the sentinel with the support, then drop the duplicate + true + } else { + false + } + }); + debug_assert!( + // an entry with `None` exists iff there are no entries with Some (for the same (in_trans_id, in_trans_grounding_id) pair) + in_.chunk_by(|a, b| (a.0, a.1) == (b.0, b.1)) + .all(|chunk| chunk.len() == 1 || chunk.iter().all(|e| e.2.is_some())) + ); + + Self { entries, out, in_ } + } + + pub fn iter_all( + &self, + ) -> impl Iterator< + Item = &( + TransitionId, + TransitionId, + Option<(TransitionGroundingId, TransitionGroundingId)>, + ), + > { + self.entries.iter() + } + pub fn iter_out_all( + &self, + ) -> impl Iterator< + Item = &( + TransitionId, + TransitionGroundingId, + (TransitionId, TransitionGroundingId), + ), + > { + self.out.iter() + } + pub fn iter_in_all( + &self, + ) -> impl Iterator< + Item = &( + TransitionId, + TransitionGroundingId, + Option<(TransitionId, TransitionGroundingId)>, + ), + > { + self.in_.iter() + } +} + +fn is_sorted_and_no_dupes(iter: impl Iterator) -> bool { + iter.is_sorted_by(|a, b| a < b) +} + +#[cfg(test)] +mod tests { + use idmap::intid::IntegerId; + + use crate::TaskId; + + use super::*; + + #[test] + fn test_sources_groundings_addition() { + let mut sources_groundings = SourcesGroundingsInfo::default(); + + sources_groundings.post_ground_source(Some(TaskId::from_int(1)), &SourceGrounding::from(vec![1000, 50, 40])); + sources_groundings.post_ground_source(Some(TaskId::from_int(1)), &SourceGrounding::from(vec![1000, 50, 30])); + + sources_groundings.post_ground_source(Some(TaskId::from_int(0)), &SourceGrounding::from(vec![2, 10, 400, 300])); + + sources_groundings.post_ground_source(None, &SourceGrounding::from(vec![])); + + sources_groundings.post_ground_source(Some(TaskId::from_int(0)), &SourceGrounding::from(vec![2, 10, 401, 300])); + sources_groundings.post_ground_source(Some(TaskId::from_int(0)), &SourceGrounding::from(vec![2, 10, 400, 301])); + sources_groundings.post_ground_source(Some(TaskId::from_int(0)), &SourceGrounding::from(vec![2, 10, 401, 301])); + + assert!( + sources_groundings.entries == { + let mut res = DirectIdMap::new(); + res.insert(TaskId::from_int(0), vec![2, 4, 5, 6]); + res.insert(TaskId::from_int(1), vec![0, 1]); + (vec![3], res) + } + ); + } + + #[test] + fn test_sources_groundings_panic_on_dupe() { + fn catch_unwind_silent R + std::panic::UnwindSafe, R>(f: F) -> std::thread::Result { + let prev_hook = std::panic::take_hook(); + std::panic::set_hook(Box::new(|_| {})); + let result = std::panic::catch_unwind(f); + std::panic::set_hook(prev_hook); + result + } + let result = catch_unwind_silent(|| { + let mut sources_groundings = SourcesGroundingsInfo::default(); + + sources_groundings.post_ground_source(Some(TaskId::from_int(0)), &SourceGrounding::from(vec![1, 1, 1])); + sources_groundings.post_ground_source(Some(TaskId::from_int(0)), &SourceGrounding::from(vec![1, 1, 1])); + }); + assert!(result.is_err()); + } + + #[test] + fn test_state_vars_groundings() { + let mut state_vars_groundings = StateVarsGroundingsInfo::default(); + + assert_eq!( + state_vars_groundings.post_ground_state_var(&"a".to_string(), &[1, 2, 3, 4]), + 0 + ); + assert_eq!( + state_vars_groundings.post_ground_state_var(&"b".to_string(), &[1, 2, 3, 4]), + 1 + ); + assert_eq!( + state_vars_groundings.post_ground_state_var(&"b".to_string(), &[1, 2, 3, 5]), + 2 + ); + assert_eq!( + state_vars_groundings.post_ground_state_var(&"a".to_string(), &[1, 2, 3, 4]), + 0 + ); + } + + #[test] + fn test_terms_groundings() { + // TODO + } +} diff --git a/planning/timelines/src/constraints/lprelax/encoder/ground/utils.rs b/planning/timelines/src/constraints/lprelax/encoder/ground/utils.rs new file mode 100644 index 00000000..a04ec4ab --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/encoder/ground/utils.rs @@ -0,0 +1,225 @@ +/// Assuming `entries` is sorted w.r.t. the given comparison function +/// (i.e. it is a sequence of `Less`, then `Equal`, then `Greater` elements), +/// returns the `[start, end)` index range of the `Equal` elements, or `None` if there are none. +pub fn binary_search_range_by Ordering>(entries: &[T], mut cmp: F) -> Option<(usize, usize)> { + let start = entries.partition_point(|elt| cmp(elt) == Ordering::Less); + // `entries[start]`, if any, is `Equal` or `Greater`, so this is the length of the `Equal` run. + let len = entries[start..].partition_point(|elt| cmp(elt) == Ordering::Equal); + + (len != 0).then_some((start, start + len)) +} + +use std::cmp::Ordering; + +/// Co-iterates `a` and `b`, yielding the pairs of chunks (maximal runs) sharing a same key. +/// Keys present in only one of the two slices are skipped over. +/// +/// Both slices must be sorted by `key` ! +/// +/// Skips and runs are traversed by galloping (see [`gallop_run_end`]). +pub fn merge_join_chunks_by_key<'a, T, K: Ord>( + a: &'a [T], + b: &'a [T], + key: impl Fn(&T) -> K, +) -> impl Iterator { + let (mut i, mut j) = (0usize, 0usize); + + std::iter::from_fn(move || { + while i < a.len() && j < b.len() { + let ka = key(&a[i]); + let kb = key(&b[j]); + + match ka.cmp(&kb) { + // Jump straight to the first element that could match, not just past this run. + Ordering::Less => i = gallop_run_end(a, i, |x| key(x) < kb), + Ordering::Greater => j = gallop_run_end(b, j, |x| key(x) < ka), + Ordering::Equal => { + let i_end = gallop_run_end(a, i, |x| key(x) == ka); + let j_end = gallop_run_end(b, j, |x| key(x) == kb); + let chunks = (&a[i..i_end], &b[j..j_end]); + (i, j) = (i_end, j_end); + return Some(chunks); + } + } + } + None + }) +} + +/// Index of the first element after `start` for which `pred` is false, or `s.len()`. +/// Assumes `pred` is "true-then-false" over `s[start..]`, and true at `start`. +/// +/// Galloping: an exponential probe followed by a binary search bounded by the bracket it finds. +/// Costs a single comparison when the run has length 1, and `O(log d)` where `d` is the distance to the answer. +/// Unlike `partition_point`, which is `O(log n)` in the length of the whole remaining slice regardless of how near the answer is. +#[inline] +fn gallop_run_end(s: &[T], start: usize, pred: impl Fn(&T) -> bool) -> usize { + debug_assert!(start < s.len() && pred(&s[start])); + + let mut lo = start; + let mut step = 1; + while start + step < s.len() && pred(&s[start + step]) { + lo = start + step; + step *= 2; + } + // `pred` holds at `lo`, and fails at `hi` unless `hi == s.len()`. + let hi = (start + step).min(s.len()); + lo + 1 + s[lo + 1..hi].partition_point(&pred) +} + +/// Assuming `a` and `b` are sorted w.r.t the given comparison function, +/// merges `a` into `b`, replacing `b` and draining `a`. +/// *Consecutive* duplicate elements are ignored (i.e. exactly equal ones, not those 'equivalent' according to the comparison function). +/// But !! WARNING !! this means that if the comparison function is not a total order, +/// and the entries were sorted with an unstable sorted, +/// then some duplicate elements could end up not one after the other, and thus could be left in !!! +pub fn merge_dedup_into(a: &mut Vec, b: &mut Vec, cmp: impl Fn(&T, &T) -> Ordering) { + if a.is_empty() { + return; + } + + if b.is_empty() { + std::mem::swap(b, a); + return; + } + + // Save the capacity before draining `a`. + let capacity = a.len() + b.len(); + + let old_b = std::mem::take(b); + let mut a_iter = a.drain(..); + let mut b_iter = old_b.into_iter(); + + let mut result = Vec::with_capacity(capacity); + + while let (Some(a_ref), Some(b_ref)) = (a_iter.as_slice().first(), b_iter.as_slice().first()) { + match cmp(a_ref, b_ref) { + Ordering::Less => { + result.push(a_iter.next().unwrap()); + } + Ordering::Greater => { + result.push(b_iter.next().unwrap()); + } + Ordering::Equal => { + let a_elem = a_iter.next().unwrap(); + let b_elem = b_iter.next().unwrap(); + + if a_elem == b_elem { + result.push(a_elem); + } else { + result.push(a_elem); + result.push(b_elem); + } + } + } + } + + result.extend(a_iter); + result.extend(b_iter); + + *b = result; +} + +#[cfg(test)] +mod tests { + use super::{binary_search_range_by, merge_dedup_into, merge_join_chunks_by_key}; + + #[test] + fn test_binary_search_range_by() { + let entries = vec![(5, 3), (5, 6), (6, 3), (6, 5), (6, 8), (12, 10), (13, 8)]; + assert!(entries.is_sorted()); + + assert_eq!(binary_search_range_by(&entries, |elt| elt.0.cmp(&6)), Some((2, 5))); + assert_eq!( + binary_search_range_by(&entries, |elt| if elt.0 < 6 { + std::cmp::Ordering::Less + } else if elt.0 == 6 && elt.1 < 7 { + std::cmp::Ordering::Equal + } else { + std::cmp::Ordering::Greater + }), + Some((2, 4)) + ); + assert_eq!(binary_search_range_by(&entries, |elt| elt.0.cmp(&5)), Some((0, 2))); + assert_eq!(binary_search_range_by(&entries, |elt| elt.cmp(&(6, 5))), Some((3, 4))); + + assert_eq!(binary_search_range_by(&entries, |elt| elt.0.cmp(&13)), Some((6, 7))); + + assert_eq!(binary_search_range_by(&entries, |elt| elt.0.cmp(&1)), None); + assert_eq!(binary_search_range_by(&entries, |elt| elt.0.cmp(&7)), None); + assert_eq!(binary_search_range_by(&entries, |elt| elt.0.cmp(&20)), None); + + assert_eq!(binary_search_range_by(&[] as &[(i32, i32)], |elt| elt.0.cmp(&6)), None); + } + + #[test] + fn equal_sort_keys_are_both_retained() { + type Entry = (i32, i32, i32, i32); + + let cmp = |a: &Entry, b: &Entry| a.0.cmp(&b.0).then_with(|| a.1.cmp(&b.1)).then_with(|| a.2.cmp(&b.2)); + + let mut a = vec![(1, 10, 20, 100), (2, 10, 20, 200)]; + + let mut b = vec![(1, 10, 20, 101), (2, 10, 10, 100), (2, 10, 20, 200), (3, 10, 20, 300)]; + + merge_dedup_into(&mut a, &mut b, cmp); + + assert!(a.is_empty()); + assert!({ + let cand1 = vec![ + (1, 10, 20, 101), + (1, 10, 20, 100), + (2, 10, 10, 100), + (2, 10, 20, 200), + (3, 10, 20, 300), + ]; + let cand2 = vec![ + (1, 10, 20, 100), + (1, 10, 20, 101), + (2, 10, 10, 100), + (2, 10, 20, 200), + (3, 10, 20, 300), + ]; + b == cand1 || b == cand2 + }); + } + + #[test] + fn remaining_nonconsecutive_duplicates_with_partial_order() { + type Entry = (i32, i32, i32, i32); + let cmp = |a: &Entry, b: &Entry| a.0.cmp(&b.0).then_with(|| a.1.cmp(&b.1)).then_with(|| a.2.cmp(&b.2)); + + let mut a = vec![(1, 1, 1, 1), (1, 1, 1, 2), (1, 1, 1, 3), (1, 1, 1, 4)]; + let mut b = vec![(1, 1, 1, 0), (1, 1, 1, 1), (1, 1, 1, 2), (1, 1, 1, 3)]; + + merge_dedup_into(&mut a, &mut b, cmp); + + assert_eq!( + b, + vec![ + (1, 1, 1, 1), + (1, 1, 1, 0), + (1, 1, 1, 2), + (1, 1, 1, 1), // dup of index 0 + (1, 1, 1, 3), + (1, 1, 1, 2), // dup of index 2 + (1, 1, 1, 4), + (1, 1, 1, 3), // dup of index 4 + ] + ); + } + + #[test] + fn test_merge_join_chunks_by_key() { + let a = [(0, 'a'), (0, 'b'), (2, 'c'), (3, 'd'), (3, 'e')]; + let b = [(0, 'x'), (1, 'y'), (3, 'z')]; + + let got: Vec<_> = merge_join_chunks_by_key(&a, &b, |&(k, _)| k).collect(); + assert_eq!(got, vec![(&a[0..2], &b[0..1]), (&a[3..5], &b[2..3])]); + + // disjoint keys, and empty inputs + assert_eq!(merge_join_chunks_by_key(&a, &[(1, 'y')], |&(k, _)| k).count(), 0); + assert_eq!(merge_join_chunks_by_key(&a, &[], |&(k, _)| k).count(), 0); + assert_eq!(merge_join_chunks_by_key(&[], &b, |&(k, _)| k).count(), 0); + } +} diff --git a/planning/timelines/src/constraints/lprelax/encoder/mod.rs b/planning/timelines/src/constraints/lprelax/encoder/mod.rs new file mode 100644 index 00000000..cd70a7e7 --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/encoder/mod.rs @@ -0,0 +1,213 @@ +mod encode; +mod ground; +pub mod problem; + +use aries_solver::prelude::*; + +use encode::{encode_problem_ground, encode_problem_lifted}; +use ground::{SourceGrounding, TransitionGroundingId}; +pub use problem::LpRelaxProblem; + +use crate::analysis::Source; +use crate::analysis::grounding::ground_all_tasks; +use crate::analysis::transitions::supports::{Supports, SupportsSorted}; +use crate::analysis::transitions::{TransitionId, Transitions}; +use crate::{IntTerm, SchedEncoder, Task}; + +#[derive(Clone)] +pub(crate) struct LpRelaxEncoder { + pub(crate) transitions: Transitions, + pub(crate) supports: Supports, + supports_sorted: Option, + + sources_ground: ground::SourcesGroundingsInfo, + state_vars_ground: ground::StateVarsGroundingsInfo, + transitions_ground: ground::TransitionsGroundingsInfo, + terms_ground: ground::TermsGroundingsInfo, + supports_ground: ground::SupportsGroundingsInfo, +} + +impl LpRelaxEncoder { + pub fn with_transitions_from(ctx: &SchedEncoder) -> Self { + let transitions = Transitions::new_unambiguous(ctx, super::ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS.get()); + let supports = Supports::from( + &transitions, + ctx, + super::ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS.get(), + ); + Self { + transitions, + supports, + supports_sorted: None, + sources_ground: Default::default(), + state_vars_ground: Default::default(), + transitions_ground: Default::default(), + terms_ground: Default::default(), + supports_ground: Default::default(), + } + } + + fn post_ground_source(&mut self, source: Source, source_grounding: &SourceGrounding, ctx: &SchedEncoder) { + // Will panic (in debug mode) if the source and the grounding have already been interned. + let source_grounding_id = self.sources_ground.post_ground_source(source, source_grounding); + + // Intern each of the variable assignments in the source grounding, + // marking each of them as appearing in it. + { + let source_terms = self.get_source_terms(source, ctx); + + for (i, &term) in source_terms.iter().enumerate() { + let value = source_grounding[i]; + if !term.is_cst() { + self.terms_ground + .post_for_ground_source(term, value, source, source_grounding_id); + } + } + } + + // For each transition of the source, intern the corresponding grounding, + // marking each of them as being appearing in the source grounding. + // Also mark each of the variable assignments as appearing in the corresponding transition groundings. + { + for &trans_id in self.transitions.of_source(source) { + let trans_grounding = self.transitions.evaluate_terms(trans_id, source_grounding, ctx); + + let trans_grounding_id = TransitionGroundingId { + state_var_grounding_id: self.state_vars_ground.post_ground_state_var( + &self.transitions.get_state_var(trans_id, ctx).fluent, + trans_grounding.args_evaluated(), + ), + val_assignment: trans_grounding.val_evaluated(), + op_assignment: trans_grounding.op_evaluated(), + }; + + for (term, value) in self + .transitions + .iter_evaluated_non_constant_terms(trans_id, source_grounding, ctx) + { + debug_assert!(!term.is_cst()); + + self.terms_ground + .post_for_ground_transition(term, value, trans_id, trans_grounding_id); + } + + self.transitions_ground.post_ground_transition( + trans_id, + trans_grounding_id, + source, + source_grounding_id, + ); + } + } + } + + fn sort(&mut self) { + let time_all = std::time::Instant::now(); + let time = std::time::Instant::now(); + + self.supports_sorted = Some(self.supports.sort()); + + tracing::info!( + "|-[LPRELAX]----- Sorted lifted supports in {}s", + time.elapsed().as_secs_f64(), + ); + let time = std::time::Instant::now(); + + self.transitions_ground.sort_for_all(); + + tracing::info!( + "|-[LPRELAX]----- Sorted transitions groundings in {}s", + time.elapsed().as_secs_f64(), + ); + let time = std::time::Instant::now(); + + self.supports_ground = + ground::SupportsGroundingsInfo::from(self.supports_sorted.as_ref().unwrap(), &self.transitions_ground); + + tracing::info!( + "|-[LPRELAX]----- Built sorted support groundings in {}s", + time.elapsed().as_secs_f64(), + ); + let time = std::time::Instant::now(); + + self.terms_ground.sort(); + + tracing::info!( + "|-[LPRELAX]----- Sorted terms groundings in {}s", + time.elapsed().as_secs_f64(), + ); + + tracing::info!("|-[LPRELAX]--- Sorted all in {}s", time_all.elapsed().as_secs_f64(),); + } + + pub fn get_source<'a>(&self, source: Source, ctx: &'a SchedEncoder) -> Option<&'a Task> { + source.map(|task_id| &ctx.sched.tasks[task_id]) + } + pub fn get_source_prez(&self, source: Source, ctx: &SchedEncoder) -> Lit { + self.get_source(source, ctx).map_or(Lit::TRUE, |task| task.presence) + } + pub fn get_source_terms<'a>(&self, source: Source, ctx: &'a SchedEncoder) -> &'a [IntTerm] { + source + .map(|task_id| &ctx.sched.tasks[task_id].args) + .unwrap_or(&ctx.sched.global_args) + } + pub fn iter_sources(&self) -> impl Iterator)> { + self.transitions.iter_of_sources() + } + + pub fn encode(&mut self, ctx: &SchedEncoder) -> LpRelaxProblem { + let time = std::time::Instant::now(); + + let binding = ground_all_tasks(ctx); + let sources_groundings = [(None, binding.empty_source_groundings().to_vec())].into_iter().chain( + binding + .all_task_groundings() + .map(|(task, gs)| (Some(task), gs.to_vec())), + ); + + tracing::info!("|-[LPRELAX]--- Ran grounder in {}s", time.elapsed().as_secs_f64(),); + + let time = std::time::Instant::now(); + + let mut n = 0; + for (source, source_groundings) in sources_groundings { + for source_grounding in source_groundings { + self.post_ground_source(source, &source_grounding, ctx); + n += 1; + } + } + + tracing::info!( + "|-[LPRELAX]--- Interned {} groundings in {}s", + n, + time.elapsed().as_secs_f64(), + ); + + self.sort(); + + let mut problem = LpRelaxProblem::default(); + + let time = std::time::Instant::now(); + + encode_problem_lifted(self, ctx, &mut problem); + + tracing::info!( + "|-[LPRELAX]--- Collected lifted constraints in {}s", + time.elapsed().as_secs_f64(), + ); + let time = std::time::Instant::now(); + + encode_problem_ground(self, ctx, &mut problem); + + tracing::info!( + "|-[LPRELAX]--- Collected ground constraints in {}s", + time.elapsed().as_secs_f64(), + ); + + problem + } + + pub fn iter_sorted_all_only_assignments(&self) -> impl Iterator { + self.terms_ground.iter_sorted_all_only_assignments() + } +} diff --git a/planning/timelines/src/constraints/lprelax/encoder/problem.rs b/planning/timelines/src/constraints/lprelax/encoder/problem.rs new file mode 100644 index 00000000..b1c42c8b --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/encoder/problem.rs @@ -0,0 +1,162 @@ +use std::collections::HashMap; + +use aries_solver::core::IntCst; + +use crate::IntTerm; +use crate::analysis::transitions::TransitionId; +use crate::{analysis::Source}; + +use super::ground::{SourceGroundingId, TransitionGroundingId}; + +/// Represents a variable / column +#[derive(Clone, Copy, Debug, Eq, PartialEq, PartialOrd, Ord, Hash)] +pub enum ColTag { + PresenceSource(Source, Option), + PresenceTransition(TransitionId, Option), + Support( + TransitionId, + TransitionId, + Option<(TransitionGroundingId, TransitionGroundingId)>, + ), + TermGround(IntTerm, IntCst), +} +impl ColTag { + /// Whether this column tag is lifted (i.e. isn't specific to a grounding, i.e. corresponds to a variable in the main model). + pub fn is_lifted(&self) -> bool { + match self { + ColTag::PresenceSource(_, grounding) => grounding.is_none(), + ColTag::PresenceTransition(_, grounding) => grounding.is_none(), + ColTag::Support(_, _, groundings) => groundings.is_none(), + ColTag::TermGround(_, _) => true, + } + } +} + +/// Represents an expression of the form `lhs cmp rhs + cst`. +/// where `cmp` can be `=`, `=`, or `>=`, and coefficient of terms in `lhs` and `rhs` is 1. +#[derive(Debug, Clone)] +pub struct RowExpr { + pub tpe: RowExprType, + /// Index separating the lhs and rhs terms. As such, `terms[separator]` must be the first term of the rhs. + separator: usize, + terms: Vec<(IntCst, ColTag)>, + /// Constant part of the expression, in the rhs + cst: IntCst, +} +#[derive(Debug, Copy, Clone, Eq, PartialEq)] +pub enum RowExprType { + Eq, + Leq, + Geq, +} +impl RowExpr { + pub fn lhs(&self) -> &[(IntCst, ColTag)] { + self.terms.get(..self.separator).unwrap_or(&[]) + } + pub fn rhs(&self) -> &[(IntCst, ColTag)] { + self.terms.get(self.separator..).unwrap_or(&[]) + } + pub fn cst(&self) -> IntCst { + self.cst + } + #[allow(dead_code)] + pub fn new(tpe: RowExprType, terms: Vec<(IntCst, ColTag)>, separator: usize, cst: IntCst) -> Self { + Self { + tpe, + separator, + terms, + cst, + } + } + pub fn new_eq_single_lhs(terms: Vec<(IntCst, ColTag)>) -> Self { + Self { + tpe: RowExprType::Eq, + separator: 1, + terms, + cst: 0, + } + } + pub fn new_leq_single_lhs(terms: Vec<(IntCst, ColTag)>) -> Self { + Self { + tpe: RowExprType::Leq, + separator: 1, + terms, + cst: 0, + } + } + pub fn new_geq_single_lhs(terms: Vec<(IntCst, ColTag)>) -> Self { + Self { + tpe: RowExprType::Geq, + separator: 1, + terms, + cst: 0, + } + } + pub fn new_leq_1(terms: Vec<(IntCst, ColTag)>) -> Self { + Self { + tpe: RowExprType::Leq, + separator: terms.len(), + terms, + cst: 1, + } + } +} + +/// Represents a problem +#[derive(Clone, Default)] +pub struct LpRelaxProblem { + /// One entry per LP column (only once simplified). + cols: Option>, + /// Maps every column tag appearing in the rows to the index, in `cols`, of the column it resolves to. + /// With column merging enabled, several tags resolve to the same column; + /// keeping all of them here lets each of their bindings to the main model constrain that column. + col_index: Option>, + rows: Vec, +} +impl LpRelaxProblem { + pub fn push_row(&mut self, row: RowExpr) { + self.cols = None; + self.col_index = None; + self.rows.push(row); + } + pub fn rows(&self) -> &[RowExpr] { + &self.rows + } + pub fn cols(&self) -> Option<&[ColTag]> { + self.cols.as_deref() + } + + /// Numbers the columns of the rows as they are, without simplifying anything: + /// every tag is its own column. + /// + /// Afterwards, [`Self::cols`] and `col_index` are available just as after [`Self::simplify`]. + pub fn seal(&mut self) { + self.number_cols(std::iter::empty()); + } + + /// Numbers the columns appearing in the rows, in order of first appearance. + /// Each `(alias, tag)` pair then makes `alias` resolve to the column of `tag` + /// (the simplification uses this for tags merged into their class' representative). + fn number_cols(&mut self, aliases: impl IntoIterator) { + let mut cols = vec![]; + let mut col_index = HashMap::new(); + + for row in &self.rows { + for &(_, col_tag) in &row.terms { + col_index.entry(col_tag).or_insert_with(|| { + cols.push(col_tag); + cols.len() - 1 + }); + } + } + + for (alias, col_tag) in aliases { + if let Some(&col) = col_index.get(&col_tag) { + col_index.insert(alias, col); + } + } + + self.cols = Some(cols); + self.col_index = Some(col_index); + } +} diff --git a/planning/timelines/src/constraints/lprelax/mod.rs b/planning/timelines/src/constraints/lprelax/mod.rs new file mode 100644 index 00000000..4ba21e9e --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/mod.rs @@ -0,0 +1,12 @@ +#[allow(dead_code)] +mod encoder; + +use aries_env_param::EnvParam; + +pub(crate) use encoder::LpRelaxEncoder; + +pub static ARIES_LPRELAX_USE: EnvParam = EnvParam::new("ARIES_LPRELAX_USE", "false"); +pub static ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS: EnvParam = + EnvParam::new("ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS", "true"); +pub static ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS: EnvParam = + EnvParam::new("ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS", "true"); From ab216b51fc305e3de0f1b199718ec87d0968168e Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Thu, 1 Oct 2026 17:38:44 +0200 Subject: [PATCH 14/19] feat(timelines): add lprelax problem simplification a flag allows to choose whether equal columns are merged, using a union-find structure --- .../constraints/lprelax/encoder/problem.rs | 53 +- .../lprelax/encoder/problem/simplify.rs | 527 ++++++++++++++++++ .../timelines/src/constraints/lprelax/mod.rs | 2 + 3 files changed, 581 insertions(+), 1 deletion(-) create mode 100644 planning/timelines/src/constraints/lprelax/encoder/problem/simplify.rs diff --git a/planning/timelines/src/constraints/lprelax/encoder/problem.rs b/planning/timelines/src/constraints/lprelax/encoder/problem.rs index b1c42c8b..a768c356 100644 --- a/planning/timelines/src/constraints/lprelax/encoder/problem.rs +++ b/planning/timelines/src/constraints/lprelax/encoder/problem.rs @@ -1,10 +1,14 @@ +mod simplify; + use std::collections::HashMap; use aries_solver::core::IntCst; +use aries_solver::core::state::Domains; use crate::IntTerm; use crate::analysis::transitions::TransitionId; -use crate::{analysis::Source}; +use crate::encoder::SchedEncoder; +use crate::{analysis::Source, constraints::lprelax::LpRelaxEncoder}; use super::ground::{SourceGroundingId, TransitionGroundingId}; @@ -159,4 +163,51 @@ impl LpRelaxProblem { self.cols = Some(cols); self.col_index = Some(col_index); } + + /// Simplifies the problem, given what the domains already fix: + /// + /// - a column whose value is known is replaced by that value in every row; + /// - a row sitting at one of its bounds (i.e. equality row) pins all of its columns (e.g. `x + y = 0`, or a single-term `x = 1`) -- which in turn feeds the point above; + /// - rows that became empty are dropped, and a row that became contradictory turns the whole problem into a single trivially infeasible one; + /// - when `merge_equal_columns` is set, a row of the form `c*x - c*y = 0` merges the two columns into a single one, + /// (represented by a "lifted" column tag when possible). + /// + /// Afterwards, [`Self::cols`] holds one entry per remaining LP column. + /// + /// See [`Self::seal`] to get a usable problem without simplifying anything. + pub fn simplify( + &mut self, + encoder: &LpRelaxEncoder, + ctx: &SchedEncoder, + doms: &Domains, + merge_equal_columns: bool, + ) { + let prev_rows_len = self.rows().len(); + let time = std::time::Instant::now(); + + if simplify::try_simplify(self, encoder, ctx, doms, merge_equal_columns).is_err() { + self.make_infeasible(); + } + + tracing::info!( + "|-[LPRELAX]--- LPrelax problem simplification: {} rows removed in {}s (remaining: {} rows and {} columns)", + prev_rows_len - self.rows().len(), + time.elapsed().as_secs_f64(), + self.rows().len(), + self.cols().as_ref().unwrap().len(), + ); + } + + /// Replaces the problem by two rows over a single column, `x <= 0` and `x >= 1`: + /// infeasible whatever the column's bounds, while each row on its own has consistent bounds + /// (a row with inconsistent bounds makes HiGHS crash). + fn make_infeasible(&mut self) { + let col_tag = ColTag::PresenceSource(None, None); + + self.rows = vec![ + RowExpr::new(RowExprType::Leq, vec![(1, col_tag)], 1, 0), + RowExpr::new(RowExprType::Geq, vec![(1, col_tag)], 1, 1), + ]; + self.seal(); + } } diff --git a/planning/timelines/src/constraints/lprelax/encoder/problem/simplify.rs b/planning/timelines/src/constraints/lprelax/encoder/problem/simplify.rs new file mode 100644 index 00000000..416e06a2 --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/encoder/problem/simplify.rs @@ -0,0 +1,527 @@ +//! The simplification pass over an [`LpRelaxProblem`], and the equivalence classes it works on. + +use std::collections::HashMap; + +use aries_solver::core::IntCst; + +use crate::{Domains, SchedEncoder}; + +use super::super::LpRelaxEncoder; +use super::{ColTag, LpRelaxProblem, RowExpr, RowExprType}; + +/// Records the values that the domains already fix for the lifted presence and support columns, +/// and for the term grounding columns appearing in the rows. +fn seed_known_values( + pb: &LpRelaxProblem, + classes: &mut EquivClasses, + encoder: &LpRelaxEncoder, + ctx: &SchedEncoder, + doms: &Domains, +) -> Result<(), Infeasible> { + let lifted_presence_and_support_cols = encoder + .iter_sources() + .flat_map(|(source, trans_ids)| { + std::iter::chain( + [( + ColTag::PresenceSource(source, None), + encoder.get_source_prez(source, ctx), + )], + trans_ids.iter().map(|&trans_id| { + ( + ColTag::PresenceTransition(trans_id, None), + encoder.transitions.get_prez(trans_id, ctx), + ) + }), + ) + }) + .chain( + encoder + .supports_sorted + .as_ref() + .unwrap() + .sorted_out() + .iter() + .filter_map(|&((out_trans_id, in_trans_id), active)| { + active.map(|active| (ColTag::Support(out_trans_id, in_trans_id, None), active)) + }), + ); + + for (col_tag, lit) in lifted_presence_and_support_cols { + let value = if doms.entails(doms.presence(lit)) { + match doms.value(lit) { + Some(true) => 1, + Some(false) => 0, + None => continue, + } + } else if doms.entails(!doms.presence(lit)) { + 0 + } else { + continue; + }; + + if value == 1 && matches!(col_tag, ColTag::Support(..)) && encoder.supports.with_condition_out_transitions { + // In the case where conditions can be used as out-transitions, + // it is UNSOUND to derive all the columns corresponding to present and true causal link support literals as equal to 1. + // Indeed, that case forbids the LP relaxation from having an effect support 2 or more conditions, even though it is allowed in the main model. + // As such, this could force two support columns two 1, making their sum equal to 2, + // while this very sum would be constrained to be <= 1 by the lp relaxation (with conditions allowed to be out-transitions), + // which was contradictory. + continue; + } + + let class = classes.intern(col_tag); + classes.set_value(class, value)?; + } + + for row in &pb.rows { + for &(_, col_tag) in &row.terms { + let ColTag::TermGround(term, value) = col_tag else { + continue; + }; + + // Out of the term's bounds: the column is 0 whether or not the term is present. + // (It is 1 only if the term is known present and fixed to that value). + let known = if doms.entails(!doms.presence(term)) || value < doms.lb(term) || doms.ub(term) < value { + 0 + } else if doms.entails(doms.presence(term)) && doms.lb(term) == doms.ub(term) { + 1 + } else { + continue; + }; + + let class = classes.intern(col_tag); + classes.set_value(class, known)?; + } + } + + Ok(()) +} + +pub(super) fn try_simplify( + pb: &mut LpRelaxProblem, + encoder: &LpRelaxEncoder, + ctx: &SchedEncoder, + doms: &Domains, + merge_equal_columns: bool, +) -> Result<(), Infeasible> { + try_simplify_inner(pb, &mut EquivClasses::new(merge_equal_columns), encoder, ctx, doms) +} + +fn try_simplify_inner( + pb: &mut LpRelaxProblem, + classes: &mut EquivClasses, + encoder: &LpRelaxEncoder, + ctx: &SchedEncoder, + doms: &Domains, +) -> Result<(), Infeasible> { + seed_known_values(pb, classes, encoder, ctx, doms)?; + + // Learn values (and equalities) from the rows until nothing new comes out. + // Each round re-reads the original rows, whose canonical form only gets simpler as knowledge grows. + loop { + let mut learned = false; + for row in &pb.rows { + learned |= CanonicalRow::of(row, classes).learn(classes)?; + } + if !learned { + break; + } + } + + // Rewrite the rows over the representatives of the columns that are left. + let mut rows = Vec::with_capacity(pb.rows.len()); + + for row in &pb.rows { + let canon = CanonicalRow::of(row, classes); + if canon.coefs.is_empty() { + // Empty: `learn` has already checked that the constant satisfies the comparison. + continue; + } + + let terms = canon + .coefs + .iter() + .map(|&(coef, class)| (coef, classes.representative(class))) + .collect::>(); + + let separator = terms.len(); + rows.push(RowExpr::new(canon.tpe, terms, separator, canon.cst)); + } + pb.rows = rows; + + // Every "surviving" column tag resolves to its equivalence class' representative column. + // (Tags of classes with a known value aren't columns anymore: their representative appears in no row.) + let aliases = classes + .interned() + .into_iter() + .map(|(tag, class)| (tag, classes.representative(class))) + .collect::>(); + pb.number_cols(aliases); + + Ok(()) +} + +/// The problem was found trivially infeasible during simplification. +type Infeasible = (); + +/// Equivalence classes over the columns, together with the value / representative of a class (once it is known). +/// +/// (Only used / merged when `merging` is enabled -- otherwise each class stays a singleton). +struct EquivClasses { + merging: bool, + tags: Vec, + index: HashMap, + /// Union-find over the indices in `tags`. + parent: Vec, + size: Vec, + /// For each class (by root), the index of the tag it is represented by. + representative: Vec, + /// For each class (by root), its value once known. + values: HashMap, +} +impl EquivClasses { + pub(super) fn new(merging: bool) -> Self { + Self { + merging, + tags: vec![], + index: HashMap::new(), + parent: vec![], + size: vec![], + representative: vec![], + values: HashMap::new(), + } + } + + fn intern(&mut self, tag: ColTag) -> u32 { + if let Some(&i) = self.index.get(&tag) { + return i; + } + let i = self.tags.len() as u32; + self.tags.push(tag); + self.parent.push(i); + self.size.push(1); + self.representative.push(i); + self.index.insert(tag, i); + i + } + + /// All interned tags, with the index they were interned at. + fn interned(&self) -> Vec<(ColTag, u32)> { + self.index.iter().map(|(&tag, &i)| (tag, i)).collect() + } + + fn find(&mut self, mut i: u32) -> u32 { + while self.parent[i as usize] != i { + let grandparent = self.parent[self.parent[i as usize] as usize]; + self.parent[i as usize] = grandparent; + i = grandparent; + } + i + } + + fn value(&mut self, i: u32) -> Option { + let root = self.find(i); + self.values.get(&root).copied() + } + + /// The tag that the class of `i` is represented by (a lifted one whenever the class holds one). + fn representative(&mut self, i: u32) -> ColTag { + let root = self.find(i); + self.tags[self.representative[root as usize] as usize] + } + + /// Records that the class of `i` takes value `v`. Returns whether that was new. + fn set_value(&mut self, i: u32, v: IntCst) -> Result { + if !(0..=1).contains(&v) { + return Err(()); // columns live in `[0, 1]` + } + let root = self.find(i); + match self.values.get(&root) { + Some(&old) if old != v => Err(()), // leaves the known value in place, as `merge` does + Some(_) => Ok(false), + None => { + self.values.insert(root, v); + Ok(true) + } + } + } + + /// Merges the classes of `a` and `b`, some row having proven them equal. + /// Returns whether that was new. + /// + /// Does nothing when merging is disabled. + fn merge(&mut self, a: u32, b: u32) -> Result { + if !self.merging { + return Ok(false); + } + + let (mut kept, mut merged) = (self.find(a), self.find(b)); + if kept == merged { + return Ok(false); + } + if let (Some(va), Some(vb)) = (self.values.get(&kept), self.values.get(&merged)) + && va != vb + { + return Err(()); + } + + if self.size[kept as usize] < self.size[merged as usize] { + std::mem::swap(&mut kept, &mut merged); + } + self.parent[merged as usize] = kept; + self.size[kept as usize] += self.size[merged as usize]; + + if let Some(v) = self.values.remove(&merged) { + self.values.insert(kept, v); + } + + // Of the two representatives, prefer a lifted tag; break ties deterministically. + let (r_kept, r_merged) = (self.representative[kept as usize], self.representative[merged as usize]); + let key = |i: u32| { + let tag = self.tags[i as usize]; + (!tag.is_lifted(), tag) + }; + self.representative[kept as usize] = if key(r_kept) <= key(r_merged) { r_kept } else { r_merged }; + + Ok(true) + } +} + +/// A row rewritten with the representative columns / tags of the union-find +struct CanonicalRow { + tpe: RowExprType, + coefs: Vec<(IntCst, u32)>, + cst: IntCst, +} +impl CanonicalRow { + fn of(row: &RowExpr, classes: &mut EquivClasses) -> Self { + let mut coefs: Vec<(IntCst, u32)> = Vec::with_capacity(row.terms.len()); + let mut cst = row.cst(); + + for (i, &(coef, col_tag)) in row.terms.iter().enumerate() { + // `lhs cmp rhs + cst` reads as `sum(lhs) - sum(rhs) cmp cst`. + let coef = if i < row.separator { coef } else { -coef }; + let class = classes.intern(col_tag); + + match classes.value(class) { + Some(value) => cst -= coef * value, + None => { + let root = classes.find(class); + match coefs.iter_mut().find(|(_, c)| *c == root) { + Some((c, _)) => *c += coef, + None => coefs.push((coef, root)), + } + } + } + } + coefs.retain(|&(coef, _)| coef != 0); + + Self { + tpe: row.tpe, + coefs, + cst, + } + } + + /// Applies everything this row implies on its own. + /// Returns whether anything was learned. + fn learn(&self, classes: &mut EquivClasses) -> Result { + // The row reads `sum cmp cst`, where `sum` ranges over `[min, max]` given `[0, 1]` columns. + let min: IntCst = self.coefs.iter().map(|&(coef, _)| coef.min(0)).sum(); + let max: IntCst = self.coefs.iter().map(|&(coef, _)| coef.max(0)).sum(); + + let (lb, ub) = match self.tpe { + RowExprType::Eq => (Some(self.cst), Some(self.cst)), + RowExprType::Leq => (None, Some(self.cst)), + RowExprType::Geq => (Some(self.cst), None), + }; + + if lb.is_some_and(|lb| lb > max) || ub.is_some_and(|ub| ub < min) { + return Err(()); + } + + let mut learned = false; + + // At one of its bounds, every column of the row is pinned to the end of its own range. + if lb == Some(max) { + for &(coef, class) in &self.coefs { + learned |= classes.set_value(class, if coef > 0 { 1 } else { 0 })?; + } + } + if ub == Some(min) { + for &(coef, class) in &self.coefs { + learned |= classes.set_value(class, if coef > 0 { 0 } else { 1 })?; + } + } + + // `c*x - c*y = 0` proves the two columns equal. + if self.tpe == RowExprType::Eq + && self.cst == 0 + && let [(coef_a, a), (coef_b, b)] = self.coefs[..] + && coef_a == -coef_b + { + learned |= classes.merge(a, b)?; + } + + Ok(learned) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn prez_trans_lifted(trans: usize) -> ColTag { + ColTag::PresenceTransition(trans, None) + } + fn prez_source_ground(grounding: usize) -> ColTag { + ColTag::PresenceSource(None, Some(grounding)) + } + /// A row `sum(terms) cmp cst`, all terms on the lhs. + fn row(tpe: RowExprType, terms: Vec<(IntCst, ColTag)>, cst: IntCst) -> RowExpr { + let separator = terms.len(); + RowExpr::new(tpe, terms, separator, cst) + } + /// The coefficient that `canon` carries for the class of `tag` (0 if it carries none). + fn coef_of(canon: &CanonicalRow, classes: &mut EquivClasses, tag: ColTag) -> IntCst { + let class = classes.intern(tag); + let root = classes.find(class); + canon + .coefs + .iter() + .find(|&&(_, c)| c == root) + .map_or(0, |&(coef, _)| coef) + } + + /// An equivalence class holds at most one value, within `[0, 1]`, and a union carries it over. + #[test] + fn classes_track_values_and_reject_conflicts() { + let mut classes = EquivClasses::new(true); + let a = classes.intern(prez_trans_lifted(0)); + let b = classes.intern(prez_trans_lifted(1)); + + assert_eq!( + classes.intern(prez_trans_lifted(0)), + a, + "interning a tag again yields its class" + ); + assert_eq!(classes.value(a), None); + + assert!(classes.set_value(a, 1).unwrap()); // learned something new + assert!(!classes.set_value(a, 1).unwrap()); // already known + assert!(classes.set_value(a, 0).is_err()); // contradicts what is known + assert!(classes.set_value(b, 2).is_err()); // a column lives in [0, 1] + + assert!(classes.merge(a, b).unwrap()); + assert_eq!(classes.value(b), Some(1)); // the union carries the value over + + let mut classes = EquivClasses::new(true); + let a = classes.intern(prez_trans_lifted(0)); + let b = classes.intern(prez_trans_lifted(1)); + classes.set_value(a, 0).unwrap(); + classes.set_value(b, 1).unwrap(); + assert!(classes.merge(a, b).is_err()); // cannot be equal and differ + } + + /// A class is represented by a lifted tag whenever it holds one, + /// and is never merged at all when merging is disabled. + #[test] + fn merging_prefers_a_lifted_representative() { + let mut classes = EquivClasses::new(true); + let g = classes.intern(prez_source_ground(0)); + let l = classes.intern(prez_trans_lifted(7)); + + assert!(classes.merge(g, l).unwrap()); + assert_eq!(classes.find(g), classes.find(l)); + assert_eq!(classes.representative(g), prez_trans_lifted(7)); + assert_eq!(classes.representative(l), prez_trans_lifted(7)); + assert!(!classes.merge(g, l).unwrap()); // already merged + + let mut classes = EquivClasses::new(false); + let g = classes.intern(prez_source_ground(0)); + let l = classes.intern(prez_trans_lifted(7)); + + assert!(!classes.merge(g, l).unwrap()); + assert_ne!(classes.find(g), classes.find(l)); + } + + /// `lhs cmp rhs + cst` is canonicalised into `sum(lhs) - sum(rhs) cmp cst`, with the constant columns' values folded into the constant. + #[test] + fn canonicalisation_negates_the_rhs_and_folds_known_values() { + let mut classes = EquivClasses::new(true); + // `x = y + 1`, i.e. `x - y = 1` + let r = RowExpr::new( + RowExprType::Eq, + vec![(1, prez_trans_lifted(0)), (1, prez_trans_lifted(1))], + 1, + 1, + ); + + let canon = CanonicalRow::of(&r, &mut classes); + assert_eq!(canon.cst, 1); + assert_eq!(coef_of(&canon, &mut classes, prez_trans_lifted(0)), 1); + assert_eq!(coef_of(&canon, &mut classes, prez_trans_lifted(1)), -1); + + // knowing `y = 1`, the very same row now reads `x = 2` + let y = classes.intern(prez_trans_lifted(1)); + classes.set_value(y, 1).unwrap(); + + let canon = CanonicalRow::of(&r, &mut classes); + assert_eq!(canon.cst, 2); + assert_eq!(coef_of(&canon, &mut classes, prez_trans_lifted(0)), 1); + assert_eq!(coef_of(&canon, &mut classes, prez_trans_lifted(1)), 0); // gone into the constant + } + + /// An equality row sitting at one of its bounds pins each of its columns to the end of its own range, + /// while one of opposite coefficients merges them instead. + #[test] + fn an_equality_row_pins_or_merges_its_columns() { + // `x + y = 2` leaves no choice but `x = y = 1`, and `x + y = 0` none but `x = y = 0` + for (cst, pinned_to) in [(2, 1), (0, 0)] { + let mut classes = EquivClasses::new(true); + let r = row( + RowExprType::Eq, + vec![(1, prez_trans_lifted(0)), (1, prez_trans_lifted(1))], + cst, + ); + assert!(CanonicalRow::of(&r, &mut classes).learn(&mut classes).unwrap()); + + for tag in [prez_trans_lifted(0), prez_trans_lifted(1)] { + let class = classes.intern(tag); + assert_eq!(classes.value(class), Some(pinned_to)); + } + } + + let mut classes = EquivClasses::new(true); + // `x - y = 0` proves the two equal without fixing either + let r = row( + RowExprType::Eq, + vec![(1, prez_trans_lifted(0)), (-1, prez_trans_lifted(1))], + 0, + ); + assert!(CanonicalRow::of(&r, &mut classes).learn(&mut classes).unwrap()); + + let (x, y) = ( + classes.intern(prez_trans_lifted(0)), + classes.intern(prez_trans_lifted(1)), + ); + assert_eq!(classes.find(x), classes.find(y)); + assert_eq!(classes.value(x), None); + } + + /// A row out of the reach of its columns is infeasible, be it from its coefficients alone or only once the known values have been folded in. + #[test] + fn a_row_out_of_the_reach_of_its_columns_is_infeasible() { + let mut classes = EquivClasses::new(true); + // `x >= 2` cannot hold for a column in `[0, 1]` + let r = row(RowExprType::Geq, vec![(1, prez_trans_lifted(0))], 2); + assert!(CanonicalRow::of(&r, &mut classes).learn(&mut classes).is_err()); + + let mut classes = EquivClasses::new(true); + let x = classes.intern(prez_trans_lifted(0)); + classes.set_value(x, 0).unwrap(); + // `x >= 1` with `x = 0` reads as the empty row `0 >= 1` + let r = row(RowExprType::Geq, vec![(1, prez_trans_lifted(0))], 1); + assert!(CanonicalRow::of(&r, &mut classes).learn(&mut classes).is_err()); + } +} diff --git a/planning/timelines/src/constraints/lprelax/mod.rs b/planning/timelines/src/constraints/lprelax/mod.rs index 4ba21e9e..5e60efd1 100644 --- a/planning/timelines/src/constraints/lprelax/mod.rs +++ b/planning/timelines/src/constraints/lprelax/mod.rs @@ -10,3 +10,5 @@ pub static ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS: EnvParam = EnvParam::new("ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS", "true"); pub static ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS: EnvParam = EnvParam::new("ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS", "true"); +pub static ARIES_LPRELAX_MERGE_EQUAL_COLUMNS: EnvParam = + EnvParam::new("ARIES_LPRELAX_MERGE_EQUAL_COLUMNS", "true"); From dae607e6454b32b5162ea22fc977950c750282b0 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Fri, 2 Oct 2026 11:19:42 +0200 Subject: [PATCH 15/19] chore(lp_highs): expose reasoner stats publicly --- solver/lp_highs/src/lib.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/solver/lp_highs/src/lib.rs b/solver/lp_highs/src/lib.rs index 9b886bbb..0921de9f 100644 --- a/solver/lp_highs/src/lib.rs +++ b/solver/lp_highs/src/lib.rs @@ -21,7 +21,7 @@ use crate::{ }; #[derive(Default, Clone)] -struct LpStats { +pub struct LpStats { pub lpruns: u64, pub lpruns_time: std::time::Duration, } @@ -63,7 +63,7 @@ pub struct Lp { lp_state: LpState, bindings: Bindings, - stats: LpStats, + pub stats: LpStats, options: LpOptions, } unsafe impl Send for Lp {} From adce51a1cbdfb3ddd9966da8c426de63d5b6ec0e Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Fri, 2 Oct 2026 11:22:36 +0200 Subject: [PATCH 16/19] feat(timelines): add lprelax wrapper reasoner (HiGHS-based) --- Cargo.lock | 1 + planning/timelines/Cargo.toml | 1 + .../timelines/src/constraints/lprelax/mod.rs | 4 +- .../constraints/lprelax/wrappers/lp_highs.rs | 349 ++++++++++++++++++ .../src/constraints/lprelax/wrappers/mod.rs | 3 + planning/timelines/src/encoder.rs | 3 +- planning/timelines/src/explain.rs | 9 +- 7 files changed, 366 insertions(+), 4 deletions(-) create mode 100644 planning/timelines/src/constraints/lprelax/wrappers/lp_highs.rs create mode 100644 planning/timelines/src/constraints/lprelax/wrappers/mod.rs diff --git a/Cargo.lock b/Cargo.lock index a0c45c00..6040a850 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -226,6 +226,7 @@ dependencies = [ "aries-datalog", "aries-env-param", "aries-solver", + "aries-solver-lp-highs", "hashbrown 0.16.1", "idmap", "itertools 0.14.0", diff --git a/planning/timelines/Cargo.toml b/planning/timelines/Cargo.toml index 73d70175..8e50f204 100644 --- a/planning/timelines/Cargo.toml +++ b/planning/timelines/Cargo.toml @@ -10,6 +10,7 @@ edition = "2024" itertools = { workspace = true } tracing = { workspace = true } aries-solver = { path = "../../solver" } +aries-solver-lp-highs = { path = "../../solver/lp_highs" } aries-datalog = { path = "../../utils/datalog" } aries-env-param = { path = "../../utils/env_param" } idmap = { workspace = true } diff --git a/planning/timelines/src/constraints/lprelax/mod.rs b/planning/timelines/src/constraints/lprelax/mod.rs index 5e60efd1..01f1676b 100644 --- a/planning/timelines/src/constraints/lprelax/mod.rs +++ b/planning/timelines/src/constraints/lprelax/mod.rs @@ -1,5 +1,5 @@ -#[allow(dead_code)] -mod encoder; +pub(crate) mod encoder; +pub(crate) mod wrappers; use aries_env_param::EnvParam; diff --git a/planning/timelines/src/constraints/lprelax/wrappers/lp_highs.rs b/planning/timelines/src/constraints/lprelax/wrappers/lp_highs.rs new file mode 100644 index 00000000..875ca68b --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/wrappers/lp_highs.rs @@ -0,0 +1,349 @@ +use std::collections::HashMap; + +use itertools::Itertools; + +use aries_solver::backtrack::{Backtrack, DecLvl}; +use aries_solver::core::state::{Domains, DomainsSnapshot, Explanation, InferenceCause}; +use aries_solver::core::{IntCst, Lit, Var, views::Term}; +use aries_solver::reasoners::{Contradiction, ReasonerId, Theory}; + +use aries_solver_lp_highs::{Lp, LpCol, LpLit, LpOptions}; + +use crate::constraints::lprelax::encoder::problem::{ColTag, LpRelaxProblem, RowExprType}; +use crate::constraints::lprelax::{ARIES_LPRELAX_MERGE_EQUAL_COLUMNS, LpRelaxEncoder}; +use crate::{IntTerm, SchedEncoder}; + +/// Wrapper over the HiGHS-backed LP reasoner, specifically for the LP relaxation problem. +/// +/// For efficiency, the LP relaxation problem is posted to the reasoner after two steps: +/// - 1st, at the root level (but after the initial propagation), the [`LpRelaxEncoder`] is built, +/// together with a copy / cache of the domains at that level. +/// - 2nd, after all assumptions have been propagated, the [`LpRelaxEncoder`] builds the LP relaxation problem, +/// and simplifies it, notably using the cached propagated domains from the root level (*NOT* the domains after the propagation of assumptions). +/// It then binds the reasoner (the LP's columns) to the main model (events on its literals). +/// +/// This 2-stage approach allows to build the encoder more efficiently (as it will be done after the first propagation) +/// and to avoid building and solving the LP if it can be detected as unsatisfiable without it, after all assumptions are propagated. +/// +/// By default, we attempt solving the LP at most once. +#[derive(Clone)] +pub(crate) struct LpRelaxHighs { + lp: Lp, + /// WARNING NOTE: the `store` model within is not appropriate for use here, as it is stale ! + ctx: SchedEncoder, + + /// Number of assumption levels to wait for before building the relaxed problem. + num_assumptions: usize, + + /// The encoder of stage 1, with the domains it was built against. `None` until then. + encoder: Option<(LpRelaxEncoder, Domains)>, + /// Whether stage 2 has run. + posted: bool, + + num_events: u32, + propagation_calls: usize, +} + +impl LpRelaxHighs { + pub(crate) fn new(ctx: SchedEncoder, num_assumptions: usize) -> Self { + Self { + lp: Lp::with_options(LpOptions { + propagation_active: false, + ..Default::default() + }), + ctx, + num_assumptions, + encoder: None, + posted: false, + num_events: 0, + propagation_calls: 0, + } + } + + /// Stage 1: collect the transitions and their supports. Must be done at the root level. + fn build_encoder(&mut self, doms: &Domains) { + debug_assert!(self.encoder.is_none() && !self.posted); + let time = std::time::Instant::now(); + + let encoder = LpRelaxEncoder::with_transitions_from(&self.ctx); + + tracing::info!( + "|-[LPRELAX]- Built LP encoder after {} propagation calls (decision level {:?}, num events: {}) in {}s", + self.propagation_calls, + doms.current_decision_level(), + doms.num_events(), + time.elapsed().as_secs_f64(), + ); + + self.encoder = Some((encoder, doms.clone())); + } + + /// Stage 2: encode the problem, simplify it, and hand both it and its bindings to the reasoner. + fn post_relaxation(&mut self, doms: &Domains) { + debug_assert!(!self.posted); + let Some((encoder, base_doms)) = self.encoder.as_mut() else { + unreachable!("stage 2 only runs once stage 1 has") + }; + let time = std::time::Instant::now(); + let base_dec_lvl = base_doms.current_decision_level(); + debug_assert!(base_dec_lvl == DecLvl::ROOT); + + let mut problem = encoder.encode(&self.ctx); + problem.simplify(encoder, &self.ctx, base_doms, ARIES_LPRELAX_MERGE_EQUAL_COLUMNS.get()); + + let cols = post_columns_and_rows(&problem, &mut self.lp); + post_bindings(&cols, encoder, &self.ctx, base_doms, &mut self.lp); + + self.posted = true; + + tracing::info!( + "|-[LPRELAX]- Posted LP ({} columns, {} rows) after {} propagation calls (decision level {:?}, num events: {}) from the model at level {:?} in {}s", + self.lp.num_columns(), + self.lp.num_rows(), + self.propagation_calls, + doms.current_decision_level(), + doms.num_events(), + base_dec_lvl, + time.elapsed().as_secs_f64(), + ); + } +} + +impl Theory for LpRelaxHighs { + fn identity(&self) -> ReasonerId { + ReasonerId::Extra(42) + } + + fn propagate(&mut self, model: &mut Domains) -> Result<(), Contradiction> { + self.propagation_calls += 1; + + // Quiescence: the reasoners preceding this one (in propagation order) have inferred all they could. + let quiescent = { + let num_events = model.trail().num_events(); + let quiescent = num_events == self.num_events; + self.num_events = num_events; + quiescent + }; + + if quiescent { + if self.encoder.is_none() && model.current_decision_level() == DecLvl::ROOT { + self.build_encoder(model); + } + if self.encoder.is_some() + && !self.posted + && self.current_decision_level().to_int() as usize >= self.num_assumptions + { + self.post_relaxation(model); + } + } + + if quiescent && self.posted && self.lp.stats.lpruns == 0 { + self.lp.activate_propagation(); + } else { + self.lp.deactivate_propagation(); + } + self.lp.propagate(model) + } + + fn explain( + &mut self, + literal: Lit, + context: InferenceCause, + model: &DomainsSnapshot, + out_explanation: &mut Explanation, + ) { + self.lp.explain(literal, context, model, out_explanation); + } + + fn print_stats(&self) { + self.lp.print_stats(); + } + + fn clone_box(&self) -> Box { + Box::new(self.clone()) + } +} + +impl Backtrack for LpRelaxHighs { + fn save_state(&mut self) -> DecLvl { + self.lp.save_state() + } + fn num_saved(&self) -> u32 { + self.lp.num_saved() + } + fn restore_last(&mut self) { + self.lp.restore_last(); + } +} + +/// Adds one `[0, 1]` column per column of the problem, then one row per row, and returns the mapping from column tags to LP columns. +fn post_columns_and_rows(problem: &LpRelaxProblem, lp: &mut Lp) -> HashMap { + let problem_cols = problem + .cols() + .expect("the problem must have been simplified or sealed before being posted"); + + let cols: HashMap = problem_cols + .iter() + .copied() + .zip(lp.add_columns(&vec![(Some(0), Some(1)); problem_cols.len()])) + .collect(); + + debug_assert!( + problem.rows().iter().all(|row| { + row.lhs() + .iter() + .chain(row.rhs()) + .map(|&(_, tag)| cols[&tag]) + .all_unique() + }), + "a row uses the same LP column twice: its coefficients should have been summed", + ); + debug_assert!( + problem.rows().iter().all(|row| { + !row.lhs().is_empty() + || !row.rhs().is_empty() + || match row.tpe { + RowExprType::Eq => row.cst() != 0, + RowExprType::Leq => row.cst() < 0, + RowExprType::Geq => row.cst() > 0, + } + }), + "a row without any column should either have been dropped or be infeasible on its own", + ); + + lp.add_rows(problem.rows().iter().map(|row| { + let coefs = row + .lhs() + .iter() + .map(|&(coef, tag)| (cols[&tag], coef)) + .chain(row.rhs().iter().map(|&(coef, tag)| (cols[&tag], -coef))); + + let (lb, ub) = match row.tpe { + RowExprType::Eq => (Some(row.cst()), Some(row.cst())), + RowExprType::Leq => (None, Some(row.cst())), + RowExprType::Geq => (Some(row.cst()), None), + }; + + (lb, ub, coefs) + })); + + cols +} + +/// Ties the lifted columns (as well as term grounding columns) of the LP to the literals of the main model that decide them. +/// +/// Every binding is a half-binding: it constrains the column when the main model's literal becomes +/// entailed, and never the other way round. +fn post_bindings( + cols: &HashMap, + encoder: &LpRelaxEncoder, + ctx: &SchedEncoder, + doms: &Domains, + lp: &mut Lp, +) { + // Bind lifted presence columns of the LP with corresponding literals in the main CSP. + + let presence_lits_and_cols = { + let mut res = HashMap::>::new(); + + for (source, transitions) in encoder.iter_sources() { + for &trans_id in transitions { + let Some(&col) = cols.get(&ColTag::PresenceTransition(trans_id, None)) else { + continue; + }; + res.entry(encoder.transitions.get_prez(trans_id, ctx)) + .or_default() + .push(col); + } + let Some(&col) = cols.get(&ColTag::PresenceSource(source, None)) else { + continue; + }; + res.entry(encoder.get_source_prez(source, ctx)).or_default().push(col); + } + res + }; + + for (lit, lit_cols) in presence_lits_and_cols { + if lit.tautological() { + for &col in &lit_cols { + lp.tighten_column(col, (Some(1), None)); + } + } else if lit.absurd() { + for &col in &lit_cols { + lp.tighten_column(col, (None, Some(0))); + } + } else { + let p = lit.variable(); + debug_assert!(p != Var::ZERO && lit == p.geq(1)); + + for &col in &lit_cols { + lp.half_bind_fixed(doms.presence(p), lit, LpLit::geq(col, 1)); + lp.half_bind_fixed(Lit::TRUE, !lit, LpLit::leq(col, 0)); + } + } + } + + // Bind term grounding columns of the LP with corresponding literals in the main CSP. + + for (term, value) in encoder.iter_sorted_all_only_assignments() { + debug_assert!(!term.is_cst()); + let Some(&col) = cols.get(&ColTag::TermGround(term, value)) else { + continue; + }; + let var = term.variable(); + assert!(var != Var::ZERO); + + let Some(x) = var_value_of(term, value) else { + // There exists no `x` value that `var` could ever take such that `term = value` + lp.tighten_column(col, (None, Some(0))); + continue; + }; + + // The column is pinned to 0 as soon as we're sure the variable cannot take the value `x` + if let Some(above) = x.checked_add(1) { + lp.half_bind_fixed(Lit::TRUE, var.geq(above), LpLit::leq(col, 0)); + } + if let Some(below) = x.checked_sub(1) { + lp.half_bind_fixed(Lit::TRUE, var.leq(below), LpLit::leq(col, 0)); + } + } + + // Bind lifted support columns of the LP with corresponding literals in the main CSP. + + for &((out_trans_id, in_trans_id), active) in encoder.supports.unsorted_out() { + let Some(lit) = active else { continue }; + let Some(&col) = cols.get(&ColTag::Support(out_trans_id, in_trans_id, None)) else { + continue; + }; + debug_assert!(lit.variable() != Var::ZERO && lit == lit.variable().geq(1)); + + lp.half_bind_fixed(Lit::TRUE, !lit, LpLit::leq(col, 0)); + + // NOTE: the converse half-binding, raising the column's lower bound to 1 on an *active* link, is deliberately absent by default !! + // Indeed, it is *unsound* if condition transitions are allowed to support other transitions (i.e. act as out-transitions). + // See [`ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS`]. + // This is because we bound the "out-flow" of a support by 1, but the main model lets one effect support several conditions, + // so forcing two of those columns to 1 would make the LP relaxation report a contradiction that the main model does not have. + // + // // lp.half_bind_fixed(doms.presence(lit.variable()), lit, LpLit::geq(col, 1)); + } +} + +/// The value `x` of a term's variable that makes the term `a*x + b` equal to `value`, if any. +/// +/// A grounding records values of the *term*, while the literals of the main model are on its variable, hence `x = (value - b) / a`. +/// But most of the time `a` is 1 and `b` is 0. +/// +/// `None` means strictly that no integer `x` qualifies, and an arithmetic overflow panics instead of returning `None`. +/// Indeed, if an arithmetic overflow returned `None`, we could unsoundly pin a column to 0. +fn var_value_of(term: IntTerm, value: IntCst) -> Option { + const OVERFLOW_MSG: &str = "overflow while computing the variable value of a term grounding"; + + let factor = term.scaled_var.factor; + debug_assert!(factor != 0, "a non-constant term shouldn't have a 0 factor"); + + // The checked rem / div can only trip on `INT_CST_MIN` with `factor == -1`, + // as division by zero is excluded by the assert above. + let num = value.checked_sub(term.constant).expect(OVERFLOW_MSG); + (num.checked_rem(factor).expect(OVERFLOW_MSG) == 0).then(|| num.checked_div(factor).expect(OVERFLOW_MSG)) +} diff --git a/planning/timelines/src/constraints/lprelax/wrappers/mod.rs b/planning/timelines/src/constraints/lprelax/wrappers/mod.rs new file mode 100644 index 00000000..07728438 --- /dev/null +++ b/planning/timelines/src/constraints/lprelax/wrappers/mod.rs @@ -0,0 +1,3 @@ +mod lp_highs; + +pub(crate) use lp_highs::LpRelaxHighs; diff --git a/planning/timelines/src/encoder.rs b/planning/timelines/src/encoder.rs index 7126eee3..ee68c96f 100644 --- a/planning/timelines/src/encoder.rs +++ b/planning/timelines/src/encoder.rs @@ -6,6 +6,7 @@ use crate::*; /// Structure that provide all the context for encoding the scheduling problem /// into a CSP. +#[derive(Clone)] pub struct SchedEncoder { /// Scheduling problem that is being encoded pub sched: Arc, @@ -101,7 +102,7 @@ pub struct CausalLink { /// Accumulates the set of all [`CausalLink`]s in an encoding problem. /// /// These are accumulated when encoding [`HasValueAt`] constraints. -#[derive(Default)] +#[derive(Default, Clone)] pub struct CausalLinks { /// Debug util: this is used to make sure all causal links have been added before any read. /// If a new causal link is added *after* a read access, the corresponding method will panic. diff --git a/planning/timelines/src/explain.rs b/planning/timelines/src/explain.rs index 1b678f18..c6090c29 100644 --- a/planning/timelines/src/explain.rs +++ b/planning/timelines/src/explain.rs @@ -42,7 +42,14 @@ impl ExplainableSolver { c.enforce(&mut encoding); } } - let mut solver = Solver::new(encoding.store); + let mut solver = if crate::constraints::lprelax::ARIES_LPRELAX_USE.get() { + let model = encoding.store.clone(); + + let reasoner = crate::constraints::lprelax::wrappers::LpRelaxHighs::new(encoding, assumptions_map.len()); + Solver::with_extra_reasoners(model, vec![Box::new(reasoner)]) + } else { + Solver::new(encoding.store) + }; // enable stronger propagation than default in difference logic solver. // this is useful in planning models where bounds are not sufficient to reason on precedence between tasks From 278d4d65ff232e2a06ff75b99c7c3fe6f769da0e Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Fri, 2 Oct 2026 11:23:20 +0200 Subject: [PATCH 17/19] tests(timelines): add visitall unit test for lprelax --- .../timelines/src/analysis/transitions/mod.rs | 2 +- .../timelines/src/constraints/lprelax/mod.rs | 63 +++++++++++++++++++ 2 files changed, 64 insertions(+), 1 deletion(-) diff --git a/planning/timelines/src/analysis/transitions/mod.rs b/planning/timelines/src/analysis/transitions/mod.rs index 02ff16dd..21ab6017 100644 --- a/planning/timelines/src/analysis/transitions/mod.rs +++ b/planning/timelines/src/analysis/transitions/mod.rs @@ -555,7 +555,7 @@ impl Transitions { } #[cfg(test)] -mod tests { +pub(crate) mod tests { pub(crate) mod visitall; use crate::analysis::collect_nonsimple_conditions_and_effects_to_relax; diff --git a/planning/timelines/src/constraints/lprelax/mod.rs b/planning/timelines/src/constraints/lprelax/mod.rs index 01f1676b..9237ab21 100644 --- a/planning/timelines/src/constraints/lprelax/mod.rs +++ b/planning/timelines/src/constraints/lprelax/mod.rs @@ -12,3 +12,66 @@ pub static ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS: EnvParam = EnvParam::new("ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS", "true"); pub static ARIES_LPRELAX_MERGE_EQUAL_COLUMNS: EnvParam = EnvParam::new("ARIES_LPRELAX_MERGE_EQUAL_COLUMNS", "true"); + +#[cfg(test)] +mod tests { + + use super::wrappers::LpRelaxHighs; + use crate::analysis::transitions::tests::visitall::{VisitAllLine, build_and_encode_visitall_line}; + + #[test] + fn test_visitall_line() { + let sat_pb = &VisitAllLine { + num_locs: 4, + num_moves: 3, + }; + let encoder = build_and_encode_visitall_line(sat_pb, false); + let model = encoder.sched.clone().encode(); + + { + println!("Sat instance with lprelax (lprelax mustn't deem it unsat)"); + + let reasoner = LpRelaxHighs::new(encoder, 0); + let mut solver = aries_solver::solver::Solver::with_extra_reasoners(model, vec![Box::new(reasoner)]); + + assert!( + solver + .solve(aries_solver::solver::SearchLimit::None) + .is_ok_and(|sol| sol.is_some()) + ); + } + + let unsat_pb = &VisitAllLine { + num_locs: 5, + num_moves: 3, + }; + let encoder = build_and_encode_visitall_line(unsat_pb, false); + let model = encoder.sched.clone().encode(); + + { + println!("Unsat instance without lprelax (num decisions must be > 0)"); + + let mut solver = aries_solver::solver::Solver::with_extra_reasoners(model.clone(), vec![]); + + assert!( + solver + .solve(aries_solver::solver::SearchLimit::None) + .is_ok_and(|sol| sol.is_none()) + ); + assert!(solver.stats.num_decisions > 0); + + let reasoner = LpRelaxHighs::new(encoder, 0); + + println!("Unsat instance with lprelax (num decisions must be = 0, thanks to lprelax)"); + + let mut solver = aries_solver::solver::Solver::with_extra_reasoners(model, vec![Box::new(reasoner)]); + + assert!( + solver + .solve(aries_solver::solver::SearchLimit::None) + .is_ok_and(|sol| sol.is_none()) + ); + assert!(solver.stats.num_decisions == 0); + } + } +} From 552c7f2cf218ec57eca4d65c0b51785121184907 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Fri, 2 Oct 2026 11:24:39 +0200 Subject: [PATCH 18/19] ci(timelines): add ape solving validation with lprelax the purpose is to validate that the lprelax is not unsound (doesn't deem solvable problems as unsolvable) --- .github/workflows/aries.yml | 2 ++ justfile | 21 ++++++++++++++++++++- 2 files changed, 22 insertions(+), 1 deletion(-) diff --git a/.github/workflows/aries.yml b/.github/workflows/aries.yml index adff6a1d..ae8b941c 100644 --- a/.github/workflows/aries.yml +++ b/.github/workflows/aries.yml @@ -209,6 +209,8 @@ jobs: run: just ci-ape-val-opt - name: APE solver run: just ci-ape-solve 15 # increase the timeout to avoid flaky tests + - name: APE solver (LpRelaxHighs) + run: just ci-ape-solve-lprelax # timeouts should be large enough to avoid flaky tests tests: # Meta-job that only requires all test-jobs to pass diff --git a/justfile b/justfile index 8b843510..d08d1517 100644 --- a/justfile +++ b/justfile @@ -89,10 +89,29 @@ ci-pddl-parse-all lift="false" filter="d": ci-ape-val-opt: uv run ci/ape-val.py -# Checks that all problems marked are as solvable are indeed solved within their max-depth +# Checks that all problems marked as solvable are indeed solved within their max-depth ci-ape-solve timeout="5": uv run ci/ape-solve.py -t {{ timeout }} --from-toml ci/problems.toml +# Checks that all problems marked as solvable are indeed solved within their max-depth. With all lprelax configurations (unsoundness checks). +ci-ape-solve-lprelax: + ARIES_LPRELAX_USE=true ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS=false ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS=false ARIES_LPRELAX_MERGE_EQUAL_COLUMNS=false \ + uv run ci/ape-solve.py -t 30 --from-toml ci/problems.toml + ARIES_LPRELAX_USE=true ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS=false ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS=false ARIES_LPRELAX_MERGE_EQUAL_COLUMNS=true \ + uv run ci/ape-solve.py -t 90 --from-toml ci/problems.toml + ARIES_LPRELAX_USE=true ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS=false ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS=true ARIES_LPRELAX_MERGE_EQUAL_COLUMNS=false \ + uv run ci/ape-solve.py -t 90 --from-toml ci/problems.toml + ARIES_LPRELAX_USE=true ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS=false ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS=true ARIES_LPRELAX_MERGE_EQUAL_COLUMNS=true \ + uv run ci/ape-solve.py -t 90 --from-toml ci/problems.toml + ARIES_LPRELAX_USE=true ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS=true ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS=false ARIES_LPRELAX_MERGE_EQUAL_COLUMNS=false \ + uv run ci/ape-solve.py -t 90 --from-toml ci/problems.toml + ARIES_LPRELAX_USE=true ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS=true ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS=false ARIES_LPRELAX_MERGE_EQUAL_COLUMNS=true \ + uv run ci/ape-solve.py -t 90 --from-toml ci/problems.toml + ARIES_LPRELAX_USE=true ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS=true ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS=true ARIES_LPRELAX_MERGE_EQUAL_COLUMNS=false \ + uv run ci/ape-solve.py -t 90 --from-toml ci/problems.toml + ARIES_LPRELAX_USE=true ARIES_LPRELAX_RECOVER_CLOSED_WORLD_DEFAULTS=true ARIES_LPRELAX_WITH_CONDITION_OUT_TRANSITIONS=true ARIES_LPRELAX_MERGE_EQUAL_COLUMNS=true \ + uv run ci/ape-solve.py -t 90 --from-toml ci/problems.toml + bench-jsp name timeout="10": #!/usr/bin/env bash set -e # stop on first error From 0340724be3947161374d4cc06a68b4243d7b6880 Mon Sep 17 00:00:00 2001 From: Nika Beriachvili Date: Fri, 2 Oct 2026 11:25:46 +0200 Subject: [PATCH 19/19] chore(timelines): rename `is_lifted` to `is_lifted_or_term_grounding` --- planning/timelines/src/constraints/lprelax/encoder/problem.rs | 4 ++-- .../src/constraints/lprelax/encoder/problem/simplify.rs | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/planning/timelines/src/constraints/lprelax/encoder/problem.rs b/planning/timelines/src/constraints/lprelax/encoder/problem.rs index a768c356..ef3eb3c0 100644 --- a/planning/timelines/src/constraints/lprelax/encoder/problem.rs +++ b/planning/timelines/src/constraints/lprelax/encoder/problem.rs @@ -25,8 +25,8 @@ pub enum ColTag { TermGround(IntTerm, IntCst), } impl ColTag { - /// Whether this column tag is lifted (i.e. isn't specific to a grounding, i.e. corresponds to a variable in the main model). - pub fn is_lifted(&self) -> bool { + /// Whether this column tag is lifted -- i.e. isn't specific to a grounding -- or is a term grounding. + pub fn is_lifted_or_term_grounding(&self) -> bool { match self { ColTag::PresenceSource(_, grounding) => grounding.is_none(), ColTag::PresenceTransition(_, grounding) => grounding.is_none(), diff --git a/planning/timelines/src/constraints/lprelax/encoder/problem/simplify.rs b/planning/timelines/src/constraints/lprelax/encoder/problem/simplify.rs index 416e06a2..b46058b1 100644 --- a/planning/timelines/src/constraints/lprelax/encoder/problem/simplify.rs +++ b/planning/timelines/src/constraints/lprelax/encoder/problem/simplify.rs @@ -279,7 +279,7 @@ impl EquivClasses { let (r_kept, r_merged) = (self.representative[kept as usize], self.representative[merged as usize]); let key = |i: u32| { let tag = self.tags[i as usize]; - (!tag.is_lifted(), tag) + (!tag.is_lifted_or_term_grounding(), tag) }; self.representative[kept as usize] = if key(r_kept) <= key(r_merged) { r_kept } else { r_merged };