diff --git a/Cargo.lock b/Cargo.lock index ce9c22ff..74fa5e41 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1627,6 +1627,7 @@ dependencies = [ "thread-priority", "urlencoding", "wasm-bindgen", + "weak-table", ] [[package]] @@ -2862,6 +2863,12 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "weak-table" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "323f4da9523e9a669e1eaf9c6e763892769b1d38c623913647bfdc1532fe4549" + [[package]] name = "winapi" version = "0.3.9" diff --git a/Cargo.toml b/Cargo.toml index ce989175..c8965d93 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -59,6 +59,7 @@ embed_visualizer = [ ] # use nodejs to build frontend and embed in Python instead of outputing individual JSON files loose_sanity_check = [] # do not panic when check fails fast_ds = ["gxhash"] # use fast data structures to fast iterate +unsafe_pointer = [] # if turned on, use unsafe pointers [dependencies] pyo3 = { version = "0.23.4", features = [ @@ -107,6 +108,7 @@ bp = { path = "src/bp" } thread-priority = "1.2.0" lnexp = "0.2.1" gxhash = { version = "3.5.0", optional = true } +weak-table = "0.3.2" [dev-dependencies] test-case = "3.1.0" diff --git a/src/bin/aps2024_demo.rs b/src/bin/aps2024_demo.rs index 054ac783..38853f73 100644 --- a/src/bin/aps2024_demo.rs +++ b/src/bin/aps2024_demo.rs @@ -12,6 +12,7 @@ use mwpf::primal_module::*; use mwpf::primal_module_serial::*; use mwpf::util::*; use mwpf::visualize::*; +use mwpf::pointers::*; use num_traits::{FromPrimitive, Zero}; #[cfg(feature = "progress_bar")] use pbr::ProgressBar; @@ -23,8 +24,8 @@ fn debug_demo() { let mut code = CodeCapacityTailoredCode::new(3, 0., 0.01); let initializer = Arc::new(code.get_initializer()); let model_graph = Arc::new(ModelHyperGraph::new(initializer.clone())); - let mut dual_module = DualModulePQ::new_empty(&initializer); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); code.set_physical_errors(&[4]); let syndrome_pattern = Arc::new(code.get_syndrome()); let mut visualizer = Visualizer::new( @@ -55,12 +56,12 @@ fn debug_demo() { .snapshot_combined("begin".to_string(), vec![&interface_ptr, &dual_module]) .unwrap(); let decoding_graph = interface_ptr.read_recursive().decoding_graph.clone(); - let s0 = Arc::new(InvalidSubgraph::new_complete( + let s0 = Arc::new(InvalidSubgraph::new_complete_from_indices( fast_iter_set! {3}, fast_iter_set! {}, - &decoding_graph, + &mut dual_module, )); - let (_, s0_ptr) = interface_ptr.find_or_create_node(&s0, &mut dual_module); + let (_, s0_ptr) = interface_ptr.find_or_create_node(&s0, &mut dual_module, 0); dual_module.set_grow_rate(&s0_ptr, Rational::from_usize(1).unwrap()); for _ in 0..3 { dual_module.grow(Rational::new_raw(1.into(), 3.into())); @@ -69,12 +70,12 @@ fn debug_demo() { .unwrap(); } // create another node - let s1 = Arc::new(InvalidSubgraph::new_complete( + let s1 = Arc::new(InvalidSubgraph::new_complete_from_indices( fast_iter_set! {6}, fast_iter_set! {}, - &decoding_graph, + &mut dual_module, )); - let (_, s1_ptr) = interface_ptr.find_or_create_node(&s1, &mut dual_module); + let (_, s1_ptr) = interface_ptr.find_or_create_node(&s1, &mut dual_module, 0); dual_module.set_grow_rate(&s0_ptr, -Rational::from_usize(1).unwrap()); dual_module.set_grow_rate(&s1_ptr, Rational::from_usize(1).unwrap()); for _ in 0..3 { @@ -101,8 +102,8 @@ fn simple_demo() { let mut code = CodeCapacityTailoredCode::new(3, 0., 0.01); let initializer = Arc::new(code.get_initializer()); let model_graph = Arc::new(ModelHyperGraph::new(initializer.clone())); - let mut dual_module = DualModulePQ::new_empty(&initializer); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); code.set_physical_errors(&[4]); let syndrome_pattern = Arc::new(code.get_syndrome()); let mut visualizer = Visualizer::new( @@ -133,12 +134,12 @@ fn simple_demo() { .snapshot_combined("begin".to_string(), vec![&interface_ptr, &dual_module]) .unwrap(); let decoding_graph = interface_ptr.read_recursive().decoding_graph.clone(); - let s0 = Arc::new(InvalidSubgraph::new_complete( + let s0 = Arc::new(InvalidSubgraph::new_complete_from_indices( fast_iter_set! {3}, fast_iter_set! {}, - &decoding_graph, + &mut dual_module, )); - let (_, s0_ptr) = interface_ptr.find_or_create_node(&s0, &mut dual_module); + let (_, s0_ptr) = interface_ptr.find_or_create_node(&s0, &mut dual_module, 0); dual_module.set_grow_rate(&s0_ptr, Rational::from_usize(1).unwrap()); visualizer .snapshot_combined("create s0".to_string(), vec![&interface_ptr, &dual_module]) @@ -167,8 +168,8 @@ fn challenge_demo() { let mut code = CodeCapacityTailoredCode::new(5, 0., 0.01); let initializer = Arc::new(code.get_initializer()); let model_graph = Arc::new(ModelHyperGraph::new(initializer.clone())); - let mut dual_module = DualModulePQ::new_empty(&initializer); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); let syndrome_pattern = Arc::new(SyndromePattern::new_vertices(vec![10, 15, 16])); code.set_syndrome(&syndrome_pattern); let mut visualizer = Visualizer::new( @@ -236,11 +237,11 @@ fn challenge_demo() { while index >= s_ptr.len() { let (vertices, edges) = invalid_subgraphs[s_ptr.len()].clone(); let s = if vertices.is_empty() { - Arc::new(InvalidSubgraph::new(edges, &decoding_graph)) + Arc::new(InvalidSubgraph::new_from_indices(edges, dual_module)) } else { - Arc::new(InvalidSubgraph::new_complete(vertices, edges, &decoding_graph)) + Arc::new(InvalidSubgraph::new_complete_from_indices(vertices, edges, dual_module)) }; - let (_, ptr) = interface_ptr.find_or_create_node(&s, dual_module); + let (_, ptr) = interface_ptr.find_or_create_node(&s, dual_module, 0); dual_module.set_grow_rate(&ptr, Rational::zero()); s_ptr.push(ptr); } @@ -323,8 +324,8 @@ fn surface_code_example() { let mut code = CodeCapacityTailoredCode::new(9, p / 3., p / 3.); let initializer = Arc::new(code.get_initializer()); let model_graph = Arc::new(ModelHyperGraph::new(initializer.clone())); - let mut dual_module = DualModulePQ::new_empty(&initializer); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); let mut visualizer = Visualizer::new( Some(visualize_data_folder() + visualize_filename.as_str()), code.get_positions(), @@ -372,8 +373,8 @@ fn triangle_color_code_example() { let mut code = CodeCapacityColorCode::new(9, p); let initializer = Arc::new(code.get_initializer()); let model_graph = Arc::new(ModelHyperGraph::new(initializer.clone())); - let mut dual_module = DualModulePQ::new_empty(&initializer); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); let mut visualizer = Visualizer::new( Some(visualize_data_folder() + visualize_filename.as_str()), code.get_positions(), @@ -422,8 +423,8 @@ fn small_color_code_example() { let mut code = CodeCapacityColorCode::new(7, p); let initializer = Arc::new(code.get_initializer()); let model_graph = Arc::new(ModelHyperGraph::new(initializer.clone())); - let mut dual_module = DualModulePQ::new_empty(&initializer); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); let mut visualizer = Visualizer::new( Some(visualize_data_folder() + visualize_filename.as_str()), code.get_positions(), @@ -482,8 +483,8 @@ fn circuit_level_example() { ); let initializer = Arc::new(code.get_initializer()); let model_graph = Arc::new(ModelHyperGraph::new(initializer.clone())); - let mut dual_module = DualModulePQ::new_empty(&initializer); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); let mut visualizer = Visualizer::new( Some(visualize_data_folder() + visualize_filename.as_str()), code.get_positions(), diff --git a/src/bin/paper_figures.rs b/src/bin/paper_figures.rs index ea9f30d6..42c7857c 100644 --- a/src/bin/paper_figures.rs +++ b/src/bin/paper_figures.rs @@ -7,6 +7,7 @@ use mwpf::invalid_subgraph::*; use mwpf::model_hypergraph::*; use mwpf::util::*; use mwpf::visualize::*; +use mwpf::pointers::*; use num_traits::FromPrimitive; use std::sync::Arc; @@ -30,8 +31,8 @@ fn hyperedge_example() { // create dual module let initializer = Arc::new(code.get_initializer()); let model_graph = Arc::new(ModelHyperGraph::new(initializer.clone())); - let mut dual_module = DualModulePQ::new_empty(&initializer); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); // add syndrome let syndrome_pattern = Arc::new(SyndromePattern::new_vertices(vec![1, 2, 4, 6])); interface_ptr.write().decoding_graph.set_syndrome(syndrome_pattern.clone()); @@ -47,8 +48,8 @@ fn hyperedge_example() { (fast_iter_set! {1}, 0.5), ]; for (vertices, dual_variable) in dual_variables.into_iter() { - let s1 = Arc::new(InvalidSubgraph::new_complete(vertices, fast_iter_set! {}, &decoding_graph)); - let (_, s1_ptr) = interface_ptr.find_or_create_node(&s1, &mut dual_module); + let s1 = Arc::new(InvalidSubgraph::new_complete_from_indices(vertices, fast_iter_set! {}, &mut dual_module)); + let (_, s1_ptr) = interface_ptr.find_or_create_node(&s1, &mut dual_module, 0); dual_module.set_grow_rate(&s1_ptr, Rational::from_f64(dual_variable).unwrap()); } dual_module.grow(Rational::from_f64(1.).unwrap()); diff --git a/src/cli.rs b/src/cli.rs index 65bd37ad..1a57417b 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -18,6 +18,8 @@ use serde::Serialize; use serde_variant::to_variant_name; use std::env; use std::sync::Arc; +use crate::num_traits::Zero; +use crate::dual_module_pq::{EdgePtr, EdgeWeak, VertexPtr, VertexWeak, Edge, Vertex}; const TEST_EACH_ROUNDS: usize = 100; @@ -271,42 +273,100 @@ impl TypedValueParser for SerdeJsonParser { impl MatrixSpeedClass { pub fn run(&self, parameters: MatrixSpeedParameters, samples: Vec, bool)>>) { + let (vertices, edges) = Self::initialize_vertex_edges_for_matrix_testing( + (0..parameters.height).collect(), + (0..parameters.width).collect(), + ); + let vertices_weak: Vec<_> = vertices.into_iter().map(|v| v.downgrade()).collect(); + let edges_weak: Vec<_> = edges.into_iter().map(|e| e.downgrade()).collect(); + match *self { MatrixSpeedClass::EchelonTailTight => { let mut matrix = Echelon::>>::new(); - for edge_index in 0..parameters.width { - matrix.add_tight_variable(edge_index); + for edge_weak in edges_weak.iter() { + matrix.add_tight_variable(edge_weak.clone()); } - Self::run_on_matrix_interface(&matrix, samples) + Self::run_on_matrix_interface(&matrix, samples, &vertices_weak, &edges_weak); } MatrixSpeedClass::EchelonTight => { let mut matrix = Echelon::>::new(); - for edge_index in 0..parameters.width { - matrix.add_tight_variable(edge_index); + for edge_weak in edges_weak.iter() { + matrix.add_tight_variable(edge_weak.clone()); } - Self::run_on_matrix_interface(&matrix, samples) + Self::run_on_matrix_interface(&matrix, samples, &vertices_weak, &edges_weak) } MatrixSpeedClass::Echelon => { let mut matrix = Echelon::::new(); - for edge_index in 0..parameters.width { - matrix.add_variable(edge_index); + for edge_weak in edges_weak.iter() { + matrix.add_variable(edge_weak.clone()); } - Self::run_on_matrix_interface(&matrix, samples) + Self::run_on_matrix_interface(&matrix, samples, &vertices_weak, &edges_weak) } } } - pub fn run_on_matrix_interface(matrix: &M, samples: Vec, bool)>>) { + pub fn run_on_matrix_interface( + matrix: &M, + samples: Vec, bool)>>, + vertices: &Vec, + edges: &Vec, + ) { for parity_checks in samples.iter() { let mut matrix = matrix.clone(); for (vertex_index, (incident_edges, parity)) in parity_checks.iter().enumerate() { - matrix.add_constraint(vertex_index, incident_edges, *parity); + let incident_edges_weak: Vec = incident_edges.iter().map(|&i| edges[i].clone()).collect(); + matrix.add_constraint(vertices[vertex_index].clone(), &incident_edges_weak, *parity); } // for a MatrixView, visiting the columns and rows is sufficient to update its internal state matrix.columns(); matrix.rows(); } } + + fn initialize_vertex_edges_for_matrix_testing( + vertex_indices: Vec, + edge_indices: Vec, + ) -> (Vec, Vec) { + // create edges + let edges: Vec = edge_indices + .into_iter() + .map(|edge_index| { + EdgePtr::new_value( + Edge { + edge_index, + weight: Rational::zero(), + dual_nodes: vec![], + vertices: vec![], + last_updated_time: Rational::zero(), + growth_at_last_updated_time: Rational::zero(), + grow_rate: Rational::zero(), + // unit_index: Some(0), // dummy value + // connected_to_boundary_vertex: false, // dummy value + #[cfg(feature = "incr_lp")] + cluster_weights: hashbrown::HashMap::new(), + }, + (edge_index, edge_index), + ) + }) + .collect(); + + // create vertices + let vertices: Vec = vertex_indices + .into_iter() + .map(|vertex_index| { + VertexPtr::new_value( + Vertex { + vertex_index, + is_defect: false, + edges: vec![], + }, + (vertex_index, vertex_index), + ) + }) + .collect(); + + (vertices, edges) + } } impl Cli { @@ -930,7 +990,7 @@ impl ResultVerifier for VerifierActualError { } else { Rational::from( self.initializer - .get_subgraph_total_weight(&OutputSubgraph::new(error_pattern.clone(), Default::default())), + .get_subgraph_total_weight(&OutputSubgraph::new(error_pattern.clone(), Default::default(), vec![])), ) }; let (subgraph, weight_range) = solver.subgraph_range_visualizer(visualizer); diff --git a/src/cluster.rs b/src/cluster.rs index 6387e470..5b283962 100644 --- a/src/cluster.rs +++ b/src/cluster.rs @@ -1,6 +1,7 @@ use crate::dual_module::*; use crate::matrix::*; use crate::util::*; +use crate::dual_module_pq::{EdgePtr, VertexPtr}; use derivative::Derivative; @@ -8,11 +9,11 @@ use derivative::Derivative; #[derivative(Debug)] pub struct Cluster { /// vertices of the cluster - pub vertices: FastIterSet, + pub vertices: FastIterSet, /// tight edges of the cluster - pub edges: FastIterSet, + pub edges: FastIterSet, /// edges incident to the vertices but are not tight - pub hair: FastIterSet, + pub hair: FastIterSet, /// dual variables of the cluster pub nodes: FastIterSet, /// parity matrix of the cluster @@ -39,17 +40,17 @@ impl Cluster { } /// Add a vertex to the cluster - pub fn add_vertex(&mut self, vertex: VertexIndex) { + pub fn add_vertex(&mut self, vertex: VertexPtr) { self.vertices.insert(vertex); } /// Add an edge to the cluster - pub fn add_edge(&mut self, edge: EdgeIndex) { + pub fn add_edge(&mut self, edge: EdgePtr) { self.edges.insert(edge); } /// Add a hair to the cluster - pub fn add_hair(&mut self, hair: EdgeIndex) { + pub fn add_hair(&mut self, hair: EdgePtr) { self.hair.insert(hair); } @@ -87,6 +88,8 @@ pub mod tests { initializer.uniform_weights(Rational::one()); let mut solver = SolverSerialJointSingleHair::new(&Arc::new(initializer), json!({})); solver.solve_visualizer(syndrome, Some(&mut visualizer)); + let dual_module = solver.get_dual_module(); + let vertex_ptr = dual_module.get_vertex_ptr(2); if cfg!(feature = "embed_visualizer") { let html = visualizer.generate_html(json!({})); assert!(visualizer_path.ends_with(".json")); @@ -95,11 +98,11 @@ pub mod tests { println!("visualizer path: {}", &html_path); } // generate the cluster - let cluster = solver.get_cluster(2); + let cluster = solver.get_cluster(vertex_ptr); println!("cluster: {cluster:?}"); - assert_eq!(cluster.vertices, expected_vertices); - assert_eq!(cluster.edges, expected_edges); - assert_eq!(cluster.hair, expected_hair); + assert_eq!(cluster.vertices.iter().map(|v| v.read_recursive().vertex_index).collect::>(), expected_vertices); + assert_eq!(cluster.edges.iter().map(|e| e.read_recursive().edge_index).collect::>(), expected_edges); + assert_eq!(cluster.hair.iter().map(|e| e.read_recursive().edge_index).collect::>(), expected_hair); cluster } diff --git a/src/decoding_hypergraph.rs b/src/decoding_hypergraph.rs index aba39d77..5b0085e9 100644 --- a/src/decoding_hypergraph.rs +++ b/src/decoding_hypergraph.rs @@ -1,4 +1,3 @@ -use crate::matrix::*; use crate::model_hypergraph::*; use crate::util::*; use crate::visualize::*; @@ -53,52 +52,6 @@ impl DecodingHyperGraph { pub fn new_defects(model_graph: Arc, defect_vertices: Vec) -> Self { Self::new(model_graph, Arc::new(SyndromePattern::new_vertices(defect_vertices))) } - - pub fn find_valid_subgraph( - &self, - edges: &FastIterSet, - vertices: &FastIterSet, - ) -> Option { - let mut matrix = Echelon::::new(); - for &edge_index in edges.iter() { - matrix.add_variable(edge_index); - } - - for &vertex_index in vertices.iter() { - let incident_edges = self.get_vertex_neighbors(vertex_index); - let parity = self.is_vertex_defect(vertex_index); - matrix.add_constraint(vertex_index, incident_edges, parity); - } - matrix.get_solution() - } - - pub fn find_valid_subgraph_auto_vertices(&self, edges: &FastIterSet) -> Option { - self.find_valid_subgraph(edges, &self.get_edges_neighbors(edges)) - } - - pub fn is_valid_cluster(&self, edges: &FastIterSet, vertices: &FastIterSet) -> bool { - self.find_valid_subgraph(edges, vertices).is_some() - } - - pub fn is_valid_cluster_auto_vertices(&self, edges: &FastIterSet) -> bool { - self.find_valid_subgraph_auto_vertices(edges).is_some() - } - - pub fn is_vertex_defect(&self, vertex_index: VertexIndex) -> bool { - self.defect_vertices_hashset.contains(&vertex_index) - } - - pub fn get_edge_neighbors(&self, edge_index: EdgeIndex) -> &Vec { - self.model_graph.get_edge_neighbors(edge_index) - } - - pub fn get_vertex_neighbors(&self, vertex_index: VertexIndex) -> &Vec { - self.model_graph.get_vertex_neighbors(vertex_index) - } - - pub fn get_edges_neighbors(&self, edges: &FastIterSet) -> FastIterSet { - self.model_graph.get_edges_neighbors(edges) - } } impl MWPSVisualizer for DecodingHyperGraph { diff --git a/src/dual_module.rs b/src/dual_module.rs index 2439b5db..299cea07 100644 --- a/src/dual_module.rs +++ b/src/dual_module.rs @@ -10,11 +10,13 @@ use crate::model_hypergraph::*; use crate::num_traits::{FromPrimitive, One, Signed, ToPrimitive, Zero}; use crate::pointers::*; use crate::primal_module::Affinity; -use crate::primal_module_serial::PrimalClusterPtr; +use crate::primal_module_serial::{PrimalClusterPtr, PrimalModuleSerialNodeWeak}; use crate::relaxer_optimizer::OptimizerResult; use crate::util::*; use crate::visualize::*; +use crate::dual_module_pq::{EdgePtr, VertexPtr, EdgeWeak}; use hashbrown::{HashMap, HashSet}; +use crate::matrix::*; #[cfg(feature = "python_binding")] use pyo3::prelude::*; @@ -77,11 +79,13 @@ pub struct DualNode { /// the pointer to the global time /// Note: may employ some unsafe features while being sound in performance-critical cases /// and can remove option when removing dual_module_serial - global_time: Option>, + global_time: Option>, /// the last time this dual_node is synced/updated with the global time pub last_updated_time: Rational, /// dual variable's value at the last updated time pub dual_variable_at_last_updated_time: Rational, + /// the corresponding PrimalModuleSerialNode + pub primal_module_serial_node: Option, } impl DualNode { @@ -110,14 +114,14 @@ impl DualNode { } /// initialize the global time pointer and the last_updated_time - pub fn init_time(&mut self, global_time_ptr: ArcRwLock) { + pub fn init_time(&mut self, global_time_ptr: ArcManualSafeLock) { self.last_updated_time = global_time_ptr.read_recursive().clone(); self.global_time = Some(global_time_ptr); } } -pub type DualNodePtr = ArcRwLock; -pub type DualNodeWeak = WeakRwLock; +pub type DualNodePtr = ArcManualSafeLock; +pub type DualNodeWeak = WeakManualSafeLock; impl std::fmt::Debug for DualNodePtr { fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { @@ -126,9 +130,17 @@ impl std::fmt::Debug for DualNodePtr { .field("index", &dual_node.index) .field("dual_variable", &dual_node.get_dual_variable()) .field("grow_rate", &dual_node.grow_rate) - .field("hair", &dual_node.invalid_subgraph.hair) + .field( + "hair", + &dual_node + .invalid_subgraph + .hair + .iter() + .map(|e| e.read_recursive().edge_index) + .collect::>(), + ) .finish() - // let new = ArcRwLock::new_value(Rational::zero()); + // let new = ArcManualSafeLock::new_value(Rational::zero()); // let global_time = dual_node.global_time.as_ref().unwrap_or(&new).read_recursive(); // write!( // f, @@ -165,8 +177,8 @@ pub struct DualModuleInterface { pub decoding_graph: DecodingHyperGraph, } -pub type DualModuleInterfacePtr = ArcRwLock; -pub type DualModuleInterfaceWeak = WeakRwLock; +pub type DualModuleInterfacePtr = ArcManualSafeLock; +pub type DualModuleInterfaceWeak = WeakManualSafeLock; impl std::fmt::Debug for DualModuleInterfacePtr { fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { @@ -258,7 +270,7 @@ pub enum DualReport { /// common trait that must be implemented for each implementation of dual module pub trait DualModuleImpl { /// create a new dual module with empty syndrome - fn new_empty(initializer: &Arc) -> Self; + fn new_empty(initializer: &Arc, partition_id: usize) -> Self where Self: Sized; /// clear all growth and existing dual nodes, prepared for the next decoding fn clear(&mut self); @@ -286,13 +298,13 @@ pub trait DualModuleImpl { fn grow(&mut self, length: Rational); /// get all nodes contributing to the edge - fn get_edge_nodes(&self, edge_index: EdgeIndex) -> Vec; + fn get_edge_nodes(&self, edge_ptr: EdgePtr) -> Vec; /// get the slack on a specific edge (weight - growth) - fn get_edge_slack(&self, edge_index: EdgeIndex) -> Rational; + fn get_edge_slack(&self, edge_ptr: EdgePtr) -> Rational; /// check if the edge is tight - fn is_edge_tight(&self, edge_index: EdgeIndex) -> bool; + fn is_edge_tight(&self, edge_ptr: EdgePtr) -> bool; /* New tuning-related methods */ // mode managements @@ -340,20 +352,20 @@ pub trait DualModuleImpl { } /// grow a specific edge on the spot - fn grow_edge(&self, _edge_index: EdgeIndex, _amount: &Rational) { + fn grow_edge(&self, _edge_ptr: EdgePtr, _amount: &Rational) { panic!("this dual_module doesn't support edge growth"); } /// `is_edge_tight` but in tuning phase - fn is_edge_tight_tune(&self, edge_index: EdgeIndex) -> bool { + fn is_edge_tight_tune(&self, edge_ptr: EdgePtr) -> bool { eprintln!("this dual_module does not implement tuning"); - self.is_edge_tight(edge_index) + self.is_edge_tight(edge_ptr) } /// `get_edge_slack` but in tuning phase - fn get_edge_slack_tune(&self, edge_index: EdgeIndex) -> Rational { + fn get_edge_slack_tune(&self, edge_ptr: EdgePtr) -> Rational { eprintln!("this dual_module does not implement tuning"); - self.get_edge_slack(edge_index) + self.get_edge_slack(edge_ptr) } /* miscs */ @@ -378,7 +390,7 @@ pub trait DualModuleImpl { fn get_obstacles_tune( &self, optimizer_result: OptimizerResult, - dual_node_deltas: FastIterMap, + dual_node_deltas: FastIterMap, ) -> FastIterSet { let mut obstacles = FastIterSet::new(); match optimizer_result { @@ -391,9 +403,9 @@ pub trait DualModuleImpl { dual_node_ptr: dual_node_ptr.clone(), }); } - for edge_index in node_ptr_read.invalid_subgraph.hair.iter() { - if grow_rate.is_positive() && self.is_edge_tight_tune(*edge_index) { - obstacles.insert(Obstacle::Conflict { edge_index: *edge_index }); + for edge_ptr in node_ptr_read.invalid_subgraph.hair.iter() { + if grow_rate.is_positive() && self.is_edge_tight_tune(edge_ptr.clone()) { + obstacles.insert(Obstacle::Conflict { edge_ptr: edge_ptr.clone() }); } } } @@ -404,14 +416,14 @@ pub trait DualModuleImpl { // check if the single direction is growable let mut actual_grow_rate = Rational::from_usize(std::usize::MAX).unwrap(); let node_ptr_read = dual_node_ptr.ptr.read_recursive(); - for edge_index in node_ptr_read.invalid_subgraph.hair.iter() { - actual_grow_rate = std::cmp::min(actual_grow_rate, self.get_edge_slack_tune(*edge_index)); + for edge_ptr in node_ptr_read.invalid_subgraph.hair.iter() { + actual_grow_rate = std::cmp::min(actual_grow_rate, self.get_edge_slack_tune(edge_ptr.clone())); } if actual_grow_rate.is_zero() { // if not, return the current obstacles - for edge_index in node_ptr_read.invalid_subgraph.hair.iter() { - if grow_rate.is_positive() && self.is_edge_tight_tune(*edge_index) { - obstacles.insert(Obstacle::Conflict { edge_index: *edge_index }); + for edge_ptr in node_ptr_read.invalid_subgraph.hair.iter() { + if grow_rate.is_positive() && self.is_edge_tight_tune(edge_ptr.clone()) { + obstacles.insert(Obstacle::Conflict { edge_ptr: edge_ptr.clone() }); } } if grow_rate.is_negative() && node_ptr_read.dual_variable_at_last_updated_time.is_zero() { @@ -424,12 +436,12 @@ pub trait DualModuleImpl { // note: can grow directly here because this is guaranteed to only have a single direction drop(node_ptr_read); let mut node_ptr_write = dual_node_ptr.ptr.write(); - for edge_index in node_ptr_write.invalid_subgraph.hair.iter() { - self.grow_edge(*edge_index, &actual_grow_rate); + for edge_ptr in node_ptr_write.invalid_subgraph.hair.iter() { + self.grow_edge(edge_ptr.clone(), &actual_grow_rate); #[cfg(feature = "incr_lp")] - self.update_edge_cluster_weights(*edge_index, _cluster_index, actual_grow_rate.clone()); // note: comment out if not using cluster-based - if actual_grow_rate.is_positive() && self.is_edge_tight_tune(*edge_index) { - obstacles.insert(Obstacle::Conflict { edge_index: *edge_index }); + self.update_edge_cluster_weights(edge_ptr, _cluster_index, actual_grow_rate.clone()); // note: comment out if not using cluster-based + if actual_grow_rate.is_positive() && self.is_edge_tight_tune(edge_ptr.clone()) { + obstacles.insert(Obstacle::Conflict { edge_ptr: edge_ptr.clone() }); } } node_ptr_write.dual_variable_at_last_updated_time += actual_grow_rate.clone(); @@ -457,8 +469,8 @@ pub trait DualModuleImpl { } // calculate the total edge deltas - for edge_index in node_ptr_write.invalid_subgraph.hair.iter() { - match edge_deltas.entry(*edge_index) { + for edge_ptr in node_ptr_write.invalid_subgraph.hair.iter() { + match edge_deltas.entry(edge_ptr.clone()) { FastIterEntry::Vacant(v) => { v.insert(grow_rate.clone()); } @@ -468,19 +480,19 @@ pub trait DualModuleImpl { } #[cfg(feature = "incr_lp")] - self.update_edge_cluster_weights(*edge_index, _cluster_index, grow_rate.clone()); + self.update_edge_cluster_weights(edge_ptr.clone(), _cluster_index, grow_rate.clone()); // note: comment out if not using cluster-based } } // apply the edge deltas and check for obstacles - for (edge_index, grow_rate) in edge_deltas.into_iter() { + for (edge_ptr, grow_rate) in edge_deltas.into_iter() { if grow_rate.is_zero() { continue; } - self.grow_edge(edge_index, &grow_rate); - if grow_rate.is_positive() && self.is_edge_tight_tune(edge_index) { - obstacles.insert(Obstacle::Conflict { edge_index }); + self.grow_edge(edge_ptr.clone(), &grow_rate); + if grow_rate.is_positive() && self.is_edge_tight_tune(edge_ptr.clone()) { + obstacles.insert(Obstacle::Conflict { edge_ptr: edge_ptr.clone() }); } } } @@ -491,25 +503,25 @@ pub trait DualModuleImpl { /// get the edge free weight, for each edge what is the weight that are free to use by the given participating dual variables fn get_edge_free_weight( &self, - edge_index: EdgeIndex, + edge_ptr: EdgePtr, participating_dual_variables: &hashbrown::HashSet, ) -> Weight; - fn get_edge_weight(&self, edge_index: EdgeIndex) -> Weight; + fn get_edge_weight(&self, edge_weak: EdgeWeak) -> Weight; - fn get_subgraph_weight(&self, subgraph: &Subgraph) -> Weight { + fn get_subgraph_weight(&self, subgraph: &InternalSubgraph) -> Weight { let mut weight = Weight::zero(); - for &edge_index in subgraph { - weight += self.get_edge_weight(edge_index); + for edge_weak in subgraph { + weight += self.get_edge_weight(edge_weak.clone()); } weight } #[cfg(feature = "incr_lp")] - fn update_edge_cluster_weights(&self, edge_index: EdgeIndex, cluster_index: NodeIndex, grow_rate: Rational); + fn update_edge_cluster_weights(&self, edge_ptr: EdgePtr, cluster_index: NodeIndex, grow_rate: Rational); #[cfg(feature = "incr_lp")] - fn get_edge_free_weight_cluster(&self, edge_index: EdgeIndex, cluster_index: NodeIndex) -> Rational; + fn get_edge_free_weight_cluster(&self, edge_ptr: EdgePtr, cluster_index: NodeIndex) -> Rational; #[cfg(feature = "incr_lp")] fn update_edge_cluster_weights_union( @@ -519,6 +531,19 @@ pub trait DualModuleImpl { absorbing_cluster_index: NodeIndex, ); + fn get_vertex_ptr(&self, vertex_index: VertexIndex) -> VertexPtr; + + fn get_edge_ptr(&self, edge_index: EdgeIndex) -> EdgePtr; + + fn get_vertex_ptr_vec(&self, vertex_indices: &[VertexIndex]) -> Vec; + + fn get_edge_ptr_vec(&self, edge_indices: &[EdgeIndex]) -> Vec; + + fn get_vertex_num(&self) -> usize; + + fn get_edge_num(&self) -> usize; + + /// called invidually for each dual module, so we do not need EdgePtr here. fn adjust_weights_for_negative_edges(&mut self) { unimplemented!() } @@ -549,7 +574,7 @@ pub trait DualModuleImpl { #[derive(PartialEq, Eq, Debug, Clone, PartialOrd, Ord)] pub enum Obstacle { - Conflict { edge_index: EdgeIndex }, + Conflict { edge_ptr: EdgePtr }, ShrinkToZero { dual_node_ptr: OrderedDualNodePtr }, } @@ -557,11 +582,11 @@ pub enum Obstacle { impl std::hash::Hash for Obstacle { fn hash(&self, state: &mut H) { match self { - Obstacle::Conflict { edge_index } => { - (0, *edge_index as u64).hash(state); + Obstacle::Conflict { edge_ptr } => { + (0, edge_ptr).hash(state); } Obstacle::ShrinkToZero { dual_node_ptr } => { - (1, dual_node_ptr.index).hash(state); + (1, dual_node_ptr).hash(state); } } } @@ -620,25 +645,25 @@ impl DualReport { } impl DualModuleInterfacePtr { - pub fn new(model_graph: Arc) -> Self { + pub fn new(model_graph: Arc, partition_id: usize) -> Self { Self::new_value(DualModuleInterface { nodes: Vec::new(), hashmap: HashMap::new(), decoding_graph: DecodingHyperGraph::new(model_graph, Arc::new(SyndromePattern::new_empty())), - }) + }, (partition_id, 0)) } /// a dual module interface MUST be created given a concrete implementation of the dual module - pub fn new_load(decoding_graph: DecodingHyperGraph, dual_module_impl: &mut impl DualModuleImpl) -> Self { - let interface_ptr = Self::new(decoding_graph.model_graph.clone()); - interface_ptr.load(decoding_graph.syndrome_pattern, dual_module_impl); + pub fn new_load(decoding_graph: DecodingHyperGraph, dual_module_impl: &mut impl DualModuleImpl, partition_id: usize) -> Self { + let interface_ptr = Self::new(decoding_graph.model_graph.clone(), partition_id); + interface_ptr.load(decoding_graph.syndrome_pattern, dual_module_impl, partition_id); interface_ptr } - pub fn load(&self, syndrome_pattern: Arc, dual_module_impl: &mut impl DualModuleImpl) { + pub fn load(&self, syndrome_pattern: Arc, dual_module_impl: &mut impl DualModuleImpl, partition_id: usize) { self.write().decoding_graph.set_syndrome(syndrome_pattern.clone()); for vertex_idx in syndrome_pattern.defect_vertices.iter() { - self.create_defect_node(*vertex_idx, dual_module_impl); + self.create_defect_node(*vertex_idx, dual_module_impl, partition_id); } } @@ -664,12 +689,16 @@ impl DualModuleInterfacePtr { } /// make it private; use `load` instead - fn create_defect_node(&self, vertex_idx: VertexIndex, dual_module: &mut impl DualModuleImpl) -> DualNodePtr { + fn create_defect_node(&self, vertex_idx: VertexIndex, dual_module: &mut impl DualModuleImpl, partition_id: usize) -> DualNodePtr { let mut interface = self.write(); + let vertex_ptr = dual_module.get_vertex_ptr(vertex_idx); // this is okay because create_defect_node is only called upon local defect vertices, so we won't access index out of range + vertex_ptr.write().is_defect = true; + let mut vertices = FastIterSet::new(); + vertices.insert(vertex_ptr); let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete( - vec![vertex_idx].into_iter().collect(), + vertices, FastIterSet::new(), - &interface.decoding_graph, + dual_module, )); let node_index = interface.nodes.len() as NodeIndex; let node_ptr = DualNodePtr::new_value(DualNode { @@ -679,7 +708,8 @@ impl DualModuleInterfacePtr { dual_variable_at_last_updated_time: Rational::zero(), global_time: None, last_updated_time: Rational::zero(), - }); + primal_module_serial_node: None, + }, (partition_id, node_index)); interface.nodes.push(node_ptr.clone()); interface.hashmap.insert(invalid_subgraph, node_index); @@ -699,8 +729,8 @@ impl DualModuleInterfacePtr { .map(|index| interface.nodes[*index as usize].clone()) } - pub fn create_node(&self, invalid_subgraph: Arc, dual_module: &mut impl DualModuleImpl) -> DualNodePtr { - self.create_node_internal(invalid_subgraph, dual_module, Rational::one(), DualModuleImpl::add_dual_node) + pub fn create_node(&self, invalid_subgraph: Arc, dual_module: &mut impl DualModuleImpl, partition_id: usize) -> DualNodePtr { + self.create_node_internal(invalid_subgraph, dual_module, Rational::one(), DualModuleImpl::add_dual_node, partition_id) } /// `create_node` for tuning @@ -708,12 +738,14 @@ impl DualModuleInterfacePtr { &self, invalid_subgraph: Arc, dual_module: &mut impl DualModuleImpl, + partition_id: usize, ) -> DualNodePtr { self.create_node_internal( invalid_subgraph, dual_module, Rational::zero(), DualModuleImpl::add_dual_node_tune, + partition_id, ) } @@ -722,10 +754,11 @@ impl DualModuleInterfacePtr { &self, invalid_subgraph: &Arc, dual_module: &mut impl DualModuleImpl, + partition_id: usize, ) -> (bool, DualNodePtr) { match self.find_node(invalid_subgraph) { Some(node_ptr) => (true, node_ptr), - None => (false, self.create_node(invalid_subgraph.clone(), dual_module)), + None => (false, self.create_node(invalid_subgraph.clone(), dual_module, partition_id)), } } @@ -734,13 +767,18 @@ impl DualModuleInterfacePtr { &self, invalid_subgraph: &Arc, dual_module: &mut impl DualModuleImpl, + partition_id: usize ) -> Option<(bool, DualNodePtr)> { match self.find_node(invalid_subgraph) { Some(node_ptr) => Some((true, node_ptr)), - None => Some((false, self.create_node_tune(invalid_subgraph.clone(), dual_module))), + None => Some((false, self.create_node_tune(invalid_subgraph.clone(), dual_module, partition_id))), } } + pub fn get_nodes_num(&self) -> usize { + self.read_recursive().nodes.len() + } + /// internal function for creating a node, for D.R.Y. fn create_node_internal( &self, @@ -748,6 +786,7 @@ impl DualModuleInterfacePtr { dual_module: &mut D, grow_rate: Rational, add_dual_node_fn: fn(&mut D, &DualNodePtr), + partition_id: usize, ) -> DualNodePtr { debug_assert!( self.find_node(&invalid_subgraph).is_none(), @@ -765,7 +804,8 @@ impl DualModuleInterfacePtr { dual_variable_at_last_updated_time: Rational::zero(), global_time: None, last_updated_time: Rational::zero(), - }); + primal_module_serial_node: None, + }, (partition_id, node_index)); interface.nodes.push(node_ptr.clone()); drop(interface); @@ -774,29 +814,64 @@ impl DualModuleInterfacePtr { node_ptr } + + pub fn is_valid_cluster_auto_vertices(&self, edges: &FastIterSet) -> bool { + self.find_valid_subgraph_auto_vertices(edges).is_some() + } + + pub fn find_valid_subgraph_auto_vertices(&self, edges: &FastIterSet) -> Option> { + let mut vertices: FastIterSet = FastIterSet::new(); + for edge_ptr in edges.iter() { + let local_vertices = &edge_ptr.read_recursive().vertices; + for vertex in local_vertices { + vertices.insert(vertex.upgrade_force()); + } + } + + self.find_valid_subgraph(edges, &vertices) + } + + pub fn find_valid_subgraph(&self, edges: &FastIterSet, vertices: &FastIterSet) -> Option> { + let mut matrix = Echelon::::new(); + for edge_ptr in edges.iter() { + matrix.add_variable(edge_ptr.downgrade()); + } + + for vertex_ptr in vertices.iter() { + let vertex_weak = vertex_ptr.downgrade(); + let vertex = vertex_ptr.read_recursive(); + let incident_edges = &vertex.edges; + let parity = vertex.is_defect; + matrix.add_constraint(vertex_weak, incident_edges, parity); + } + matrix.get_solution() + } } // shortcuts for easier code writing at debugging impl DualModuleInterfacePtr { pub fn create_node_vec(&self, edges: &[EdgeIndex], dual_module: &mut impl DualModuleImpl) -> DualNodePtr { - let invalid_subgraph = Arc::new(InvalidSubgraph::new( - edges.iter().cloned().collect(), - &self.read_recursive().decoding_graph, - )); - self.create_node(invalid_subgraph, dual_module) + let edges_ptr = edges + .iter() + .map(|&idx| dual_module.get_edge_ptr(idx)) + .collect::>(); + let invalid_subgraph = Arc::new(InvalidSubgraph::new(edges_ptr, dual_module)); + self.create_node(invalid_subgraph, dual_module, 0) // since this function is only for debugging, we can just set the partition_id to 0 } + + // TODO: check if this function is needed pub fn create_node_complete_vec( &self, vertices: &[VertexIndex], edges: &[EdgeIndex], dual_module: &mut impl DualModuleImpl, ) -> DualNodePtr { - let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete( - vertices.iter().cloned().collect(), - edges.iter().cloned().collect(), - &self.read_recursive().decoding_graph, - )); - self.create_node(invalid_subgraph, dual_module) + unimplemented!() + // let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete( + // vertices.iter().cloned().collect(), + // edges.iter().cloned().collect(), + // )); + // self.create_node(invalid_subgraph, dual_module, 0) // since this function is only for debugging, we can just set the partition_id to 0 } } @@ -806,11 +881,29 @@ impl MWPSVisualizer for DualModuleInterfacePtr { let mut dual_nodes = Vec::::new(); for dual_node_ptr in interface.nodes.iter() { let dual_node = dual_node_ptr.read_recursive(); + let edges: Vec = dual_node + .invalid_subgraph + .edges + .iter() + .map(|e| e.read_recursive().edge_index) + .collect(); + let vertices: Vec = dual_node + .invalid_subgraph + .vertices + .iter() + .map(|e| e.read_recursive().vertex_index) + .collect(); + let hair: Vec = dual_node + .invalid_subgraph + .hair + .iter() + .map(|e| e.read_recursive().edge_index) + .collect(); #[cfg(feature = "fast_ds")] dual_nodes.push(json!({ - if abbrev { "e" } else { "edges" }: dual_node.invalid_subgraph.edges.iter().copied().collect::>(), - if abbrev { "v" } else { "vertices" }: dual_node.invalid_subgraph.vertices.iter().copied().collect::>(), - if abbrev { "h" } else { "hair" }: dual_node.invalid_subgraph.hair.iter().copied().collect::>(), + if abbrev { "e" } else { "edges" }: edges, + if abbrev { "v" } else { "vertices" }: vertices, + if abbrev { "h" } else { "hair" }: hair, if abbrev { "d" } else { "dual_variable" }: dual_node.get_dual_variable().to_f64(), if abbrev { "dn" } else { "dual_variable_numerator" }: numer_of(&dual_node.get_dual_variable()), if abbrev { "dd" } else { "dual_variable_denominator" }: denom_of(&dual_node.get_dual_variable()), @@ -820,9 +913,9 @@ impl MWPSVisualizer for DualModuleInterfacePtr { })); #[cfg(not(feature = "fast_ds"))] dual_nodes.push(json!({ - if abbrev { "e" } else { "edges" }: dual_node.invalid_subgraph.edges, - if abbrev { "v" } else { "vertices" }: dual_node.invalid_subgraph.vertices, - if abbrev { "h" } else { "hair" }: dual_node.invalid_subgraph.hair, + if abbrev { "e" } else { "edges" }: edges, + if abbrev { "v" } else { "vertices" }: vertices, + if abbrev { "h" } else { "hair" }: hair, if abbrev { "d" } else { "dual_variable" }: dual_node.get_dual_variable().to_f64(), if abbrev { "dn" } else { "dual_variable_numerator" }: numer_of(&dual_node.get_dual_variable()), if abbrev { "dd" } else { "dual_variable_denominator" }: denom_of(&dual_node.get_dual_variable()), diff --git a/src/dual_module_pq.rs b/src/dual_module_pq.rs index fcfc792c..cc95515d 100644 --- a/src/dual_module_pq.rs +++ b/src/dual_module_pq.rs @@ -144,8 +144,8 @@ impl Vertex { } } -pub type VertexPtr = ArcRwLock; -pub type VertexWeak = WeakRwLock; +pub type VertexPtr = ArcManualSafeLock; +pub type VertexWeak = WeakManualSafeLock; impl std::fmt::Debug for VertexPtr { fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { @@ -166,25 +166,25 @@ impl std::fmt::Debug for VertexWeak { #[derivative(Debug)] pub struct Edge { /// global edge index - edge_index: EdgeIndex, + pub edge_index: EdgeIndex, /// total weight of this edge - weight: Rational, + pub weight: Rational, #[derivative(Debug = "ignore")] - vertices: Vec, // note: consider using/constructing ordered vertex, this will speed up `adjust_weights_for_negative_edges` + pub vertices: Vec, // note: consider using/constructing ordered vertex, this will speed up `adjust_weights_for_negative_edges` /// the dual nodes that contributes to this edge - dual_nodes: Vec, + pub dual_nodes: Vec, /// the speed of growth, at the current time /// Note: changing this should cause the `growth_at_last_updated_time` and `last_updated_time` to update - grow_rate: Rational, + pub grow_rate: Rational, /// the last time this Edge is synced/updated with the global time - last_updated_time: Rational, + pub last_updated_time: Rational, /// growth value at the last updated time, also, growth_at_last_updated_time <= weight - growth_at_last_updated_time: Rational, + pub growth_at_last_updated_time: Rational, #[cfg(feature = "incr_lp")] /// storing the weights of the clusters that are currently contributing to this edge - cluster_weights: hashbrown::HashMap, + pub cluster_weights: hashbrown::HashMap, } impl Edge { @@ -198,8 +198,8 @@ impl Edge { } } -pub type EdgePtr = ArcRwLock; -pub type EdgeWeak = WeakRwLock; +pub type EdgePtr = ArcManualSafeLock; +pub type EdgeWeak = WeakManualSafeLock; impl std::fmt::Debug for EdgePtr { fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { @@ -251,7 +251,7 @@ where obstacle_queue: Queue, /// the global time of this dual module /// Note: Wrap-around edge case is not currently considered - global_time: ArcRwLock, + global_time: ArcManualSafeLock, /// the current mode of the dual module mode: DualModuleMode, @@ -262,8 +262,8 @@ where // negative weight handling negative_weight_sum: Rational, - negative_edges: HashSet, - flip_vertices: HashSet, + negative_edges: HashSet, // only used in weight_preprocessing and subgraph + flip_vertices: HashSet, // only used in weight_preprocessing // remember the initializer for original weights and heralded weighted edges pub initializer: Arc, @@ -274,7 +274,8 @@ where Queue: FutureQueueMethods + Default + std::fmt::Debug + Clone, { /// helper function to bring an edge update to speed with current time if needed - fn update_edge_if_necessary(&self, edge: &mut RwLockWriteGuard) { + /// the type of edge was set as `&mut RwLockWriteGuard` previously + fn update_edge_if_necessary(&self, edge: &mut Edge) { let global_time = self.global_time.read_recursive(); if global_time.eq(&edge.last_updated_time) { // the edge is not behind @@ -298,7 +299,8 @@ where } /// helper function to bring a dual node update to speed with current time if needed - fn update_dual_node_if_necessary(&mut self, node: &mut RwLockWriteGuard) { + /// the type of dual node was set as `&mut RwLockWriteGuard` previously + fn update_dual_node_if_necessary(&mut self, node: &mut DualNode) { let global_time = self.global_time.read_recursive(); if global_time.eq(&node.last_updated_time) { // the edge is not behind @@ -324,13 +326,18 @@ where fn debug_update_all(&mut self, dual_node_ptrs: &[DualNodePtr]) { // updating all edges for edge in self.edges.iter() { - let mut edge = edge.write(); - self.update_edge_if_necessary(&mut edge); + // SAFE MODE: returns RwLockWriteGuard + // UNSAFE MODE: returns &mut Edge + let mut edge_guard = edge.write(); + // Writing &mut *variable is the universal way to obtain a mutable reference to the data inside this thing. + self.update_edge_if_necessary(&mut *edge_guard); } // updating all dual nodes for dual_node_ptr in dual_node_ptrs.iter() { - let mut dual_node = dual_node_ptr.write(); - self.update_dual_node_if_necessary(&mut dual_node); + // SAFE MODE: returns RwLockWriteGuard + // UNSAFE MODE: returns &mut DualNode + let mut dual_node_guard = dual_node_ptr.write(); + self.update_dual_node_if_necessary(&mut *dual_node_guard); } } @@ -361,8 +368,8 @@ where ) -> bool { #[allow(clippy::unnecessary_cast)] return match obstacle { - Obstacle::Conflict { edge_index } => { - let edge = self.edges[*edge_index as usize].read_recursive(); + Obstacle::Conflict { edge_ptr } => { + let edge = edge_ptr.read_recursive(); // not changing, cannot have conflict if !edge.grow_rate.is_positive() { return false; @@ -389,8 +396,8 @@ where } } -pub type DualModulePQlPtr = ArcRwLock>; -pub type DualModulePQWeak = WeakRwLock>; +pub type DualModulePQlPtr = ArcManualSafeLock>; +pub type DualModulePQWeak = WeakManualSafeLock>; impl DualModuleImpl for DualModulePQGeneric where @@ -398,7 +405,7 @@ where { /// initialize the dual module, which is supposed to be reused for multiple decoding tasks with the same structure #[allow(clippy::unnecessary_cast)] - fn new_empty(initializer: &Arc) -> Self { + fn new_empty(initializer: &Arc, partition_id: usize) -> Self where Self: Sized { #[cfg(not(feature = "loose_sanity_check"))] initializer.sanity_check().unwrap(); @@ -409,14 +416,15 @@ where vertex_index, is_defect: false, edges: vec![], - }) + }, (partition_id, vertex_index)) }) .collect(); // set edges let mut edges = Vec::::new(); for hyperedge in initializer.weighted_edges.iter() { + let edge_id = edges.len() as EdgeIndex; let edge = Edge { - edge_index: edges.len() as EdgeIndex, + edge_index: edge_id, weight: hyperedge.weight.clone(), dual_nodes: vec![], vertices: hyperedge @@ -431,7 +439,7 @@ where cluster_weights: hashbrown::HashMap::new(), }; - let edge_ptr = EdgePtr::new_value(edge); + let edge_ptr = EdgePtr::new_value(edge, (partition_id, edge_id)); for &vertex_index in hyperedge.vertices.iter() { vertices[vertex_index as usize].write().edges.push(edge_ptr.downgrade()); @@ -443,7 +451,7 @@ where vertices, edges, obstacle_queue: Queue::default(), - global_time: ArcRwLock::new_value(Rational::zero()), + global_time: ArcManualSafeLock::new_value(Rational::zero(), (partition_id, 0)), mode: DualModuleMode::default(), tuning_start_time: None, total_tuning_time: None, @@ -481,6 +489,7 @@ where #[allow(clippy::unnecessary_cast)] /// Adding a defect node to the DualModule + /// TODO: Double check if we need to define the defect node here fn add_defect_node(&mut self, dual_node_ptr: &DualNodePtr) { let dual_node = dual_node_ptr.read_recursive(); debug_assert!(dual_node.invalid_subgraph.edges.is_empty()); @@ -488,12 +497,12 @@ where dual_node.invalid_subgraph.vertices.len() == 1, "defect node (without edges) should only work on a single vertex, for simplicity" ); - let vertex_index = dual_node.invalid_subgraph.vertices.iter().next().unwrap(); - let mut vertex = self.vertices[*vertex_index as usize].write(); - assert!(!vertex.is_defect, "defect should not be added twice"); - vertex.is_defect = true; + // let vertex_ptr = dual_node.invalid_subgraph.vertices.iter().next().unwrap(); + // let mut vertex = vertex_ptr.write(); + // assert!(!vertex.is_defect, "defect should not be added twice"); + // vertex.is_defect = true; drop(dual_node); - drop(vertex); + // drop(vertex); self.add_dual_node(dual_node_ptr); } @@ -515,8 +524,8 @@ where ); } - for &edge_index in dual_node.invalid_subgraph.hair.iter() { - let mut edge = self.edges[edge_index as usize].write(); + for edge_ptr in dual_node.invalid_subgraph.hair.iter() { + let mut edge = edge_ptr.write(); // should make sure the edge is up-to-speed before making its variables change self.update_edge_if_necessary(&mut edge); @@ -530,7 +539,7 @@ where // it is okay to use global_time now, as this must be up-to-speed (edge.weight.clone() - edge.growth_at_last_updated_time.clone()) / edge.grow_rate.clone() + global_time.clone(), - Obstacle::Conflict { edge_index }, + Obstacle::Conflict { edge_ptr: edge_ptr.clone() }, ); } } @@ -541,8 +550,8 @@ where let dual_node_weak = dual_node_ptr.downgrade(); let dual_node = dual_node_ptr.read_recursive(); - for &edge_index in dual_node.invalid_subgraph.hair.iter() { - let mut edge = self.edges[edge_index as usize].write(); + for edge_ptr in dual_node.invalid_subgraph.hair.iter() { + let mut edge = edge_ptr.write(); edge.dual_nodes .push(OrderedDualNodeWeak::new(dual_node.index, dual_node_weak.clone())); @@ -569,8 +578,8 @@ where } // don't reacquire the read guard - for &edge_index in dual_node.invalid_subgraph.hair.iter() { - let mut edge = self.edges[edge_index as usize].write(); + for edge_ptr in dual_node.invalid_subgraph.hair.iter() { + let mut edge = edge_ptr.write(); self.update_edge_if_necessary(&mut edge); edge.grow_rate += &grow_rate_diff; @@ -579,7 +588,7 @@ where // it is okay to use global_time now, as this must be up-to-speed (edge.weight.clone() - edge.growth_at_last_updated_time.clone()) / edge.grow_rate.clone() + global_time.clone(), - Obstacle::Conflict { edge_index }, + Obstacle::Conflict { edge_ptr: edge_ptr.clone() }, ); } } @@ -644,27 +653,23 @@ where /* identical with the dual_module_serial */ #[allow(clippy::unnecessary_cast)] - fn get_edge_nodes(&self, edge_index: EdgeIndex) -> Vec { - self.edges[edge_index as usize] - .read_recursive() - .dual_nodes - .iter() - .map(|x| x.upgrade_force().ptr) + fn get_edge_nodes(&self, edge_ptr: EdgePtr) -> Vec { + edge_ptr.read_recursive().dual_nodes.iter().map(|x| x.upgrade_force().ptr) .collect() } #[allow(clippy::unnecessary_cast)] /// how much away from saturated is the edge - fn get_edge_slack(&self, edge_index: EdgeIndex) -> Rational { - let edge = self.edges[edge_index as usize].read_recursive(); + fn get_edge_slack(&self, edge_ptr: EdgePtr) -> Rational { + let edge = edge_ptr.read_recursive(); edge.weight.clone() - (self.global_time.read_recursive().clone() - edge.last_updated_time.clone()) * edge.grow_rate.clone() - edge.growth_at_last_updated_time.clone() } /// is the edge saturated - fn is_edge_tight(&self, edge_index: EdgeIndex) -> bool { - self.get_edge_slack(edge_index).is_zero() + fn is_edge_tight(&self, edge_ptr: EdgePtr) -> bool { + self.get_edge_slack(edge_ptr).is_zero() } /* tuning mode related new methods */ @@ -673,13 +678,13 @@ where add_shared_methods!(); /// is the edge tight, but for tuning mode - fn is_edge_tight_tune(&self, edge_index: EdgeIndex) -> bool { - let edge = self.edges[edge_index].read_recursive(); + fn is_edge_tight_tune(&self, edge_ptr: EdgePtr) -> bool { + let edge = edge_ptr.read_recursive(); edge.weight == edge.growth_at_last_updated_time } - fn get_edge_slack_tune(&self, edge_index: EdgeIndex) -> Rational { - let edge = self.edges[edge_index].read_recursive(); + fn get_edge_slack_tune(&self, edge_ptr: EdgePtr) -> Rational { + let edge = edge_ptr.read_recursive(); edge.weight.clone() - edge.growth_at_last_updated_time.clone() } @@ -707,8 +712,8 @@ where } /// grow specific amount for a specific edge - fn grow_edge(&self, edge_index: EdgeIndex, amount: &Rational) { - let mut edge = self.edges[edge_index].write(); + fn grow_edge(&self, edge_ptr: EdgePtr, amount: &Rational) { + let mut edge = edge_ptr.write(); edge.growth_at_last_updated_time += amount; } @@ -758,7 +763,7 @@ where ); drop(node); - let mut node: RwLockWriteGuard = _dual_node_ptr.ptr.write(); + let mut node = _dual_node_ptr.ptr.write(); let dual_variable = node.get_dual_variable(); node.set_dual_variable(dual_variable); @@ -816,9 +821,10 @@ where + (&global_time - &dual_node_read_ptr.last_updated_time) * &dual_node_read_ptr.grow_rate; } if let Some(subgraph) = cluster.subgraph.as_ref() { - for &edge_index in subgraph.iter() { - let edge_ptr = self.edges[edge_index].read_recursive(); - primal_dual_gap -= &edge_ptr.weight; + for edge_weak in subgraph.iter() { + let edge_ptr = edge_weak.upgrade_force(); + let edge = edge_ptr.read_recursive(); + primal_dual_gap -= &edge.weight; } } if primal_dual_gap.is_zero() { @@ -830,10 +836,10 @@ where fn get_edge_free_weight( &self, - edge_index: EdgeIndex, + edge_ptr: EdgePtr, participating_dual_variables: &hashbrown::HashSet, ) -> Rational { - let edge = self.edges[edge_index].read_recursive(); + let edge = edge_ptr.read_recursive(); let mut free_weight = edge.weight.clone(); for dual_node in edge.dual_nodes.iter() { if participating_dual_variables.contains(&dual_node.index) { @@ -846,14 +852,15 @@ where free_weight } - fn get_edge_weight(&self, edge_index: EdgeIndex) -> Rational { - let edge = self.edges[edge_index].read_recursive(); + fn get_edge_weight(&self, edge_weak: EdgeWeak) -> Rational { + let binding = edge_weak.upgrade_force(); + let edge = binding.read_recursive(); edge.weight.clone() } #[cfg(feature = "incr_lp")] - fn get_edge_free_weight_cluster(&self, edge_index: EdgeIndex, cluster_index: NodeIndex) -> Rational { - let edge = self.edges[edge_index as usize].read_recursive(); + fn get_edge_free_weight_cluster(&self, edge_ptr: EdgePtr, cluster_index: NodeIndex) -> Rational { + let edge = edge_ptr.read_recursive(); edge.weight.clone() - edge .cluster_weights @@ -870,8 +877,9 @@ where absorbing_cluster_index: NodeIndex, ) { let dual_node = dual_node_ptr.read_recursive(); - for edge_index in dual_node.invalid_subgraph.hair.iter() { - let mut edge = self.edges[*edge_index as usize].write(); + for edge_weak in dual_node.invalid_subgraph.hair.iter() { + let edge_ptr = edge_weak.upgrade_force(); + let mut edge = edge_ptr.write(); if let Some(removed) = edge.cluster_weights.remove(&drained_cluster_index) { *edge .cluster_weights @@ -882,8 +890,8 @@ where } #[cfg(feature = "incr_lp")] - fn update_edge_cluster_weights(&self, edge_index: usize, cluster_index: usize, weight: Rational) { - match self.edges[edge_index].write().cluster_weights.entry(cluster_index) { + fn update_edge_cluster_weights(&self, edge_ptr: EdgePtr, cluster_index: usize, weight: Rational) { + match edge_ptr.write().cluster_weights.entry(cluster_index) { hashbrown::hash_map::Entry::Occupied(mut o) => { *o.get_mut() += weight; } @@ -893,19 +901,21 @@ where } } + /// called invidually for each dual module, so we do not need EdgePtr here. fn adjust_weights_for_negative_edges(&mut self) { - for edge in self.edges.iter() { - let mut edge = edge.write(); + for edge_ptr in self.edges.iter() { + let mut edge = edge_ptr.write(); if edge.weight.is_negative() { self.negative_edges.insert(edge.edge_index); self.negative_weight_sum += edge.weight.clone(); for vertex in edge.vertices.iter() { - let vertex = vertex.upgrade_force(); - if self.flip_vertices.contains(&vertex.read_recursive().vertex_index) { - self.flip_vertices.remove(&vertex.read_recursive().vertex_index); + let vertex_ptr = vertex.upgrade_force(); + let vertex_index = vertex_ptr.read_recursive().vertex_index; + if self.flip_vertices.contains(&vertex_index) { + self.flip_vertices.remove(&vertex_index); } else { - self.flip_vertices.insert(vertex.read_recursive().vertex_index); + self.flip_vertices.insert(vertex_index); } } @@ -928,7 +938,8 @@ where fn set_weights(&mut self, new_weights: FastIterMap) { for (edge_index, new_weight) in new_weights.into_iter() { - let mut edge = self.edges[edge_index].write(); + let edge_ptr = self.get_edge_ptr(edge_index); + let mut edge = edge_ptr.write(); edge.weight = new_weight; } } @@ -944,6 +955,38 @@ where fn get_flip_vertices(&self) -> HashSet { self.flip_vertices.clone() } + + fn get_vertex_ptr(&self, vertex_index: VertexIndex) -> VertexPtr { + self.vertices[vertex_index as usize].clone() + } + + fn get_edge_ptr(&self, edge_index: EdgeIndex) -> EdgePtr { + self.edges[edge_index as usize].clone() + } + + fn get_vertex_ptr_vec(&self, vertex_indices: &[VertexIndex]) -> Vec { + vertex_indices + .to_vec() + .iter() + .map(|&i| self.vertices[i as usize].clone()) + .collect() + } + + fn get_edge_ptr_vec(&self, edge_indices: &[EdgeIndex]) -> Vec { + edge_indices + .to_vec() + .iter() + .map(|&i| self.edges[i as usize].clone()) + .collect() + } + + fn get_vertex_num(&self) -> usize { + self.vertices.len() + } + + fn get_edge_num(&self) -> usize { + self.edges.len() + } } impl MWPSVisualizer for DualModulePQGeneric @@ -998,38 +1041,60 @@ mod tests { let mut future_obstacle_queue = _FutureObstacleQueue::::new(); assert_eq!(0, future_obstacle_queue.len()); macro_rules! ref_event { - ($index:expr) => { - Some((&$index, &Obstacle::Conflict { edge_index: $index })) + ($index:expr, $edges:expr) => { + Some((&$index, &Obstacle::Conflict { edge_ptr: $edges[$index as usize].clone() })) }; } macro_rules! value_event { - ($index:expr) => { - Some(($index, Obstacle::Conflict { edge_index: $index })) + ($index:expr, $edges:expr) => { + Some(($index, Obstacle::Conflict { edge_ptr: $edges[$index as usize].clone() })) }; } + // initialize edges + let edges: Vec = vec![0, 1, 2, 3] + .into_iter() + .map(|edge_index| { + EdgePtr::new_value( + Edge { + edge_index, + weight: Rational::zero(), + dual_nodes: vec![], + vertices: vec![], + last_updated_time: Rational::zero(), + growth_at_last_updated_time: Rational::zero(), + grow_rate: Rational::zero(), + // unit_index: None, + // connected_to_boundary_vertex: false, + #[cfg(feature = "incr_lp")] + cluster_weights: hashbrown::HashMap::new(), + }, + (0, edge_index), + ) + }) + .collect(); // test basic order - future_obstacle_queue.will_happen(2, Obstacle::Conflict { edge_index: 2 }); - future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_index: 1 }); - future_obstacle_queue.will_happen(3, Obstacle::Conflict { edge_index: 3 }); - assert_eq!(future_obstacle_queue.peek_event(), ref_event!(1)); - assert_eq!(future_obstacle_queue.peek_event(), ref_event!(1)); - assert_eq!(future_obstacle_queue.pop_event(), value_event!(1)); - assert_eq!(future_obstacle_queue.peek_event(), ref_event!(2)); - assert_eq!(future_obstacle_queue.pop_event(), value_event!(2)); - assert_eq!(future_obstacle_queue.pop_event(), value_event!(3)); + future_obstacle_queue.will_happen(2, Obstacle::Conflict { edge_ptr: edges[2].clone() }); + future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_ptr: edges[1].clone() }); + future_obstacle_queue.will_happen(3, Obstacle::Conflict { edge_ptr: edges[3].clone() }); + assert_eq!(future_obstacle_queue.peek_event(), ref_event!(1, edges)); + assert_eq!(future_obstacle_queue.peek_event(), ref_event!(1, edges)); + assert_eq!(future_obstacle_queue.pop_event(), value_event!(1, edges)); + assert_eq!(future_obstacle_queue.peek_event(), ref_event!(2, edges)); + assert_eq!(future_obstacle_queue.pop_event(), value_event!(2, edges)); + assert_eq!(future_obstacle_queue.pop_event(), value_event!(3, edges)); assert_eq!(future_obstacle_queue.peek_event(), None); // test duplicate elements, the queue must be able to hold all the duplicate events - future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_index: 1 }); - future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_index: 1 }); - future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_index: 1 }); - assert_eq!(future_obstacle_queue.pop_event(), value_event!(1)); - assert_eq!(future_obstacle_queue.pop_event(), value_event!(1)); - assert_eq!(future_obstacle_queue.pop_event(), value_event!(1)); + future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_ptr: edges[1].clone() }); + future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_ptr: edges[1].clone() }); + future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_ptr: edges[1].clone() }); + assert_eq!(future_obstacle_queue.pop_event(), value_event!(1, edges)); + assert_eq!(future_obstacle_queue.pop_event(), value_event!(1, edges)); + assert_eq!(future_obstacle_queue.pop_event(), value_event!(1, edges)); assert_eq!(future_obstacle_queue.peek_event(), None); // test order of events at the same time - future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_index: 2 }); - future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_index: 1 }); - future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_index: 3 }); + future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_ptr: edges[2].clone() }); + future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_ptr: edges[1].clone() }); + future_obstacle_queue.will_happen(1, Obstacle::Conflict { edge_ptr: edges[3].clone() }); let mut events = vec![]; while let Some((time, event)) = future_obstacle_queue.pop_event() { assert_eq!(time, 1); @@ -1053,10 +1118,10 @@ mod tests { .unwrap(); // create dual module let model_graph = code.get_model_graph(); - let mut dual_module = DualModulePQ::new_empty(&model_graph.initializer); + let mut dual_module = DualModulePQ::new_empty(&model_graph.initializer, 0); // try to work on a simple syndrome let decoding_graph = DecodingHyperGraph::new_defects(model_graph, vec![3, 12]); - let interface_ptr = DualModuleInterfacePtr::new_load(decoding_graph, &mut dual_module); + let interface_ptr = DualModuleInterfacePtr::new_load(decoding_graph, &mut dual_module, 0); visualizer .snapshot_combined("syndrome".to_string(), vec![&interface_ptr, &dual_module]) @@ -1105,10 +1170,10 @@ mod tests { .unwrap(); // create dual module let model_graph = code.get_model_graph(); - let mut dual_module = DualModulePQ::new_empty(&model_graph.initializer); + let mut dual_module = DualModulePQ::new_empty(&model_graph.initializer, 0); // try to work on a simple syndrome let decoding_graph = DecodingHyperGraph::new_defects(model_graph, vec![23, 24, 29, 30]); - let interface_ptr = DualModuleInterfacePtr::new_load(decoding_graph, &mut dual_module); + let interface_ptr = DualModuleInterfacePtr::new_load(decoding_graph, &mut dual_module, 0); visualizer .snapshot_combined("syndrome".to_string(), vec![&interface_ptr, &dual_module]) .unwrap(); @@ -1150,10 +1215,10 @@ mod tests { .unwrap(); // create dual module let model_graph = code.get_model_graph(); - let mut dual_module = DualModulePQ::new_empty(&model_graph.initializer); + let mut dual_module = DualModulePQ::new_empty(&model_graph.initializer, 0); // try to work on a simple syndrome let decoding_graph = DecodingHyperGraph::new_defects(model_graph, vec![17, 23, 29, 30]); - let interface_ptr = DualModuleInterfacePtr::new_load(decoding_graph, &mut dual_module); + let interface_ptr = DualModuleInterfacePtr::new_load(decoding_graph, &mut dual_module, 0); visualizer .snapshot_combined("syndrome".to_string(), vec![&interface_ptr, &dual_module]) .unwrap(); diff --git a/src/fast_ds.rs b/src/fast_ds.rs index ba7af30e..8b518f1f 100644 --- a/src/fast_ds.rs +++ b/src/fast_ds.rs @@ -29,7 +29,6 @@ pub struct MutValueGuard<'a, K: Hash + Clone + Eq, V: Hash> { hash: &'a mut u64, } -/// The guard will implement Deref/DerefMut so it can be used like a &mut V impl<'a, K: Hash + Clone + Eq, V: Hash> Deref for MutValueGuard<'a, K, V> { type Target = V; fn deref(&self) -> &Self::Target { @@ -42,7 +41,6 @@ impl<'a, K: Hash + Clone + Eq, V: Hash> DerefMut for MutValueGuard<'a, K, V> { } } -/// On drop, recalc new hash and update the map’s combined_hash impl<'a, K: Hash + Clone + Eq, V: Hash> Drop for MutValueGuard<'a, K, V> { fn drop(&mut self) { let new_hash = Map::::compute_hash(self.key, self.value); @@ -50,7 +48,8 @@ impl<'a, K: Hash + Clone + Eq, V: Hash> Drop for MutValueGuard<'a, K, V> { } } -impl Map { +// Basic methods need no bounds +impl Map { /// Creates a new empty map pub fn new() -> Self { Self { @@ -58,7 +57,41 @@ impl Map { combined_hash: 0, } } + + #[inline] + pub fn clear(&mut self) { + self.map.clear(); + self.combined_hash = 0; + } + + #[inline] + pub fn len(&self) -> usize { + self.map.len() + } + + #[inline] + pub fn is_empty(&self) -> bool { + self.map.is_empty() + } + + #[inline] + pub fn iter(&self) -> impl Iterator { + self.map.iter() + } + + #[inline] + pub fn keys(&self) -> impl Iterator { + self.map.keys() + } + + #[inline] + pub fn combined_hash(&self) -> u64 { + self.combined_hash + } +} +// Methods that require Hash/Eq but NOT Clone +impl Map { /// Computes the hash of a key-value pair fn compute_hash(key: &K, value: &V) -> u64 { let mut hasher = crate::util::DefaultHasher::default(); @@ -67,25 +100,11 @@ impl Map { hasher.finish() } - /// Inserts a key-value pair into the map, returning the old value if it exists - pub fn insert(&mut self, key: K, value: V) -> Option { - let hash = Self::compute_hash(&key, &value); - match self.map.entry(key.clone()) { - hashbrown::hash_map::Entry::Occupied(mut entry) => { - let old_value = entry.get_mut(); - let old_hash = Self::compute_hash(&key, old_value); - self.combined_hash = self.combined_hash.wrapping_sub(old_hash).wrapping_add(hash); - Some(std::mem::replace(old_value, value)) - } - hashbrown::hash_map::Entry::Vacant(entry) => { - self.combined_hash = self.combined_hash.wrapping_add(hash); - entry.insert(value); - None - } - } + #[inline] + pub fn contains_key(&self, key: &K) -> bool { + self.map.contains_key(key) } - /// Removes a key-value pair from the map, returning the value if it exists pub fn remove(&mut self, key: &K) -> Option { if let Some(old_val) = self.map.remove(key) { let old_hash = Self::compute_hash(key, &old_val); @@ -95,48 +114,36 @@ impl Map { None } } +} - /// Checks if the map contains a key - #[inline] - pub fn contains_key(&self, key: &K) -> bool { - self.map.contains_key(key) - } - - /// clear - #[inline] - pub fn clear(&mut self) { - self.map.clear(); - self.combined_hash = 0; - } - - /// iter - #[inline] - pub fn iter(&self) -> impl Iterator { - self.map.iter() - } - - /// len - #[inline] - pub fn len(&self) -> usize { - self.map.len() - } - - /// is_empty - #[inline] - pub fn is_empty(&self) -> bool { - self.map.is_empty() - } - - /// get +impl Map { #[inline] pub fn get(&self, key: &K) -> Option<&V> { self.map.get(key) } +} + +// Methods that require Clone (for inserting keys) +impl Map { + pub fn insert(&mut self, key: K, value: V) -> Option { + let hash = Self::compute_hash(&key, &value); + match self.map.entry(key.clone()) { + hashbrown::hash_map::Entry::Occupied(mut entry) => { + let old_value = entry.get_mut(); + let old_hash = Self::compute_hash(&key, old_value); + self.combined_hash = self.combined_hash.wrapping_sub(old_hash).wrapping_add(hash); + Some(std::mem::replace(old_value, value)) + } + hashbrown::hash_map::Entry::Vacant(entry) => { + self.combined_hash = self.combined_hash.wrapping_add(hash); + entry.insert(value); + None + } + } + } - /// The key method: returns a "guard" instead of `&mut V`. - /// When that guard is dropped, it will populate the hash with the re-hashed new value pub fn get_mut<'a>(&'a mut self, key: &'a K) -> Option> { - return if let Some(value) = self.map.get_mut(key) { + if let Some(value) = self.map.get_mut(key) { let old_hash = Self::compute_hash(key, value); self.combined_hash = self.combined_hash.wrapping_sub(old_hash); Some(MutValueGuard { @@ -146,30 +153,30 @@ impl Map { }) } else { None - }; - } - - /// Get keys in iterator form - #[inline] - pub fn keys(&self) -> impl Iterator { - self.map.keys() + } } - - /// Get combined hash value - #[inline] - pub fn combined_hash(&self) -> u64 { - self.combined_hash + + pub fn entry(&mut self, key: K) -> Entry { + match self.map.entry(key) { + hashbrown::hash_map::Entry::Occupied(entry) => Entry::Occupied(OccupiedEntry { + entry, + combined_hash: &mut self.combined_hash, + }), + hashbrown::hash_map::Entry::Vacant(entry) => Entry::Vacant(VacantEntry { + entry, + combined_hash: &mut self.combined_hash, + }), + } } } -// implement extend for owned values +// Implement Extend impl Extend<(K, V)> for Map { fn extend>(&mut self, iter: I) { let into_iter = iter.into_iter(); let reserve = if self.is_empty() { into_iter.size_hint().0 } else { - // consistent with std implementation (into_iter.size_hint().0 + 1) / 2 }; self.map.reserve(reserve); @@ -179,7 +186,37 @@ impl Extend<(K, V)> for Map { } } -// implement extend for references +impl std::ops::Index<&K> for Map { + type Output = V; + + fn index(&self, key: &K) -> &Self::Output { + self.get(key).expect("Key not found in Map") + } +} + +// ----------------------------------------------------------------------------- +// IntoIterator for References (allows `&map` and `&set` to be used in loops/cmp) +// ----------------------------------------------------------------------------- + +impl<'a, K, V> IntoIterator for &'a Map { + type Item = (&'a K, &'a V); + // Use the iterator type from the underlying hashbrown map + type IntoIter = hashbrown::hash_map::Iter<'a, K, V>; + + fn into_iter(self) -> Self::IntoIter { + self.map.iter() + } +} + +impl<'a, T> IntoIterator for &'a Set { + type Item = &'a T; + type IntoIter = hashbrown::hash_set::Iter<'a, T>; + + fn into_iter(self) -> Self::IntoIter { + self.set.iter() + } +} + impl<'a, K: Eq + Hash + Clone, V: Hash + Clone> Extend<(&'a K, &'a V)> for Map { fn extend>(&mut self, iter: I) { let into_iter = iter.into_iter(); @@ -195,13 +232,14 @@ impl<'a, K: Eq + Hash + Clone, V: Hash + Clone> Extend<(&'a K, &'a V)> for Map Hash for Map { +// Standard Trait Impls +impl Hash for Map { fn hash(&self, state: &mut H) { self.combined_hash.hash(state); } } -impl IntoIterator for Map { +impl IntoIterator for Map { type Item = (K, V); type IntoIter = hashbrown::hash_map::IntoIter; @@ -218,13 +256,13 @@ impl FromIterator<(K, V)> for Map { } } -impl Default for Map { +impl Default for Map { fn default() -> Self { Self::new() } } -impl PartialEq for Map { +impl PartialEq for Map { fn eq(&self, other: &Self) -> bool { if self.combined_hash != other.combined_hash { return false; @@ -233,40 +271,37 @@ impl PartialEq for Map { } } -impl PartialOrd for Map { +impl Eq for Map {} + +// ----------------------------------------------------------------------------- +// Ord and PartialOrd for Map +// ----------------------------------------------------------------------------- + +impl PartialOrd for Map { fn partial_cmp(&self, other: &Self) -> Option { Some(self.cmp(other)) } } -impl Eq for Map {} -impl Ord for Map { +impl Ord for Map { fn cmp(&self, other: &Self) -> Ordering { + // 1. Fast Path: Compare combined hashes first. let order = self.combined_hash.cmp(&other.combined_hash); - if !matches!(order, Ordering::Equal) { + if order != Ordering::Equal { return order; } - let self_sorted: BTreeMap<_, _> = self.map.iter().collect(); - let other_sorted: BTreeMap<_, _> = other.map.iter().collect(); + // 2. Slow Path: Deterministic comparison. + // We collect references into a BTreeMap to sort by Key. + // We use references (&K, &V) so we don't need K: Clone or V: Clone. + let self_sorted: BTreeMap<&K, &V> = self.map.iter().collect(); + let other_sorted: BTreeMap<&K, &V> = other.map.iter().collect(); + self_sorted.cmp(&other_sorted) } } -impl std::ops::Index<&K> for Map { - type Output = V; - - fn index(&self, key: &K) -> &Self::Output { - self.get(key).expect("Key not found in Map") - } -} -impl From<[(K, V); N]> for Map { - fn from(array: [(K, V); N]) -> Self { - array.into_iter().collect() - } -} - -/// An enum representing either an occupied or vacant entry in the map, consisten with std +// Entry Implementations pub enum Entry<'a, K, V> { Occupied(OccupiedEntry<'a, K, V>), Vacant(VacantEntry<'a, K, V>), @@ -283,13 +318,11 @@ pub struct VacantEntry<'a, K, V> { } impl<'a, K: Eq + Hash + Clone, V: Hash> OccupiedEntry<'a, K, V> { - /// Returns a reference to the key #[inline] pub fn get(&self) -> &V { self.entry.get() } - /// Replaces the value and returns the old value pub fn insert(&mut self, value: V) -> V { let key = self.entry.key(); let old_value = self.entry.get(); @@ -297,11 +330,9 @@ impl<'a, K: Eq + Hash + Clone, V: Hash> OccupiedEntry<'a, K, V> { let new_hash = Map::::compute_hash(key, &value); *self.combined_hash = self.combined_hash.wrapping_sub(old_hash).wrapping_add(new_hash); - self.entry.insert(value) } - /// Removes the entry and returns the value pub fn remove(self) -> V { let key = self.entry.key().clone(); let value = self.entry.remove(); @@ -320,34 +351,20 @@ impl<'a, K: Eq + Hash + Clone, V: Hash> VacantEntry<'a, K, V> { } } -impl Map { - pub fn entry(&mut self, key: K) -> Entry { - match self.map.entry(key) { - hashbrown::hash_map::Entry::Occupied(entry) => Entry::Occupied(OccupiedEntry { - entry, - combined_hash: &mut self.combined_hash, - }), - hashbrown::hash_map::Entry::Vacant(entry) => Entry::Vacant(VacantEntry { - entry, - combined_hash: &mut self.combined_hash, - }), - } - } -} - /* SET implementation */ + /// A `Set` that provides Ord and fast Hash #[derive(Debug, Clone, Derivative)] -pub struct Set { +pub struct Set { set: HashSet, combined_hash: u64, } #[cfg(feature = "python_binding")] impl<'py, T: Hash + Clone + Eq + IntoPyObject<'py>> IntoPyObject<'py> for Set { - type Target = PyAny; // the Python type - type Output = Bound<'py, Self::Target>; // in most cases this will be `Bound` - type Error = std::convert::Infallible; // the conversion error type, has to be convertable to `PyErr` + type Target = PyAny; + type Output = Bound<'py, Self::Target>; + type Error = std::convert::Infallible; fn into_pyobject(self, py: Python<'py>) -> Result { let set: std::collections::HashSet = self.set.iter().cloned().collect(); @@ -356,7 +373,8 @@ impl<'py, T: Hash + Clone + Eq + IntoPyObject<'py>> IntoPyObject<'py> for Set } } -impl Set { +// Base Implementation (No Bounds) +impl Set { /// Creates a new empty set pub fn new() -> Self { Self { @@ -365,14 +383,36 @@ impl Set { } } - /// Computes the hash of a value + #[inline] + pub fn clear(&mut self) { + self.set.clear(); + self.combined_hash = 0; + } + + #[inline] + pub fn iter(&self) -> impl Iterator { + self.set.iter() + } + + #[inline] + pub fn len(&self) -> usize { + self.set.len() + } + + #[inline] + pub fn is_empty(&self) -> bool { + self.set.is_empty() + } +} + +// Logic Implementation (Requires Hash + Eq, but NO Clone/Debug needed for basic ops) +impl Set { pub fn compute_hash(value: &T) -> u64 { let mut hasher = crate::util::DefaultHasher::default(); value.hash(&mut hasher); hasher.finish() } - /// Inserts an element, returning `true` if it was newly inserted pub fn insert(&mut self, value: T) -> bool { let hash = Self::compute_hash(&value); let inserted = self.set.insert(value); @@ -382,7 +422,6 @@ impl Set { inserted } - /// Removes an element, returning `true` if it was present pub fn remove(&mut self, value: &T) -> bool { let hash = Self::compute_hash(value); let removed = self.set.remove(value); @@ -392,59 +431,32 @@ impl Set { removed } - /// Checks if an element exists in the set #[inline] pub fn contains(&self, value: &T) -> bool { self.set.contains(value) } - /// clear #[inline] - pub fn clear(&mut self) { - self.set.clear(); - self.combined_hash = 0; + pub fn is_disjoint(&self, other: &Self) -> bool { + self.set.is_disjoint(&other.set) } - /// iter #[inline] - pub fn iter(&self) -> impl Iterator { - self.set.iter() + pub fn intersection<'a>(&'a self, other: &'a Self) -> impl Iterator { + self.set.intersection(&other.set) } - - /// Appends elements from `other` into `self`, consuming `other`. + + // Only append requires mutable access to other's hash, keeping it here is fine pub fn append(&mut self, other: &mut Self) { + // Note: drain() requires T to be moved, so no extra bounds needed self.set.extend(other.set.drain()); self.combined_hash = self.combined_hash.wrapping_add(other.combined_hash); other.combined_hash = 0; } - - /// len - #[inline] - pub fn len(&self) -> usize { - self.set.len() - } - - /// is_empty - #[inline] - pub fn is_empty(&self) -> bool { - self.set.is_empty() - } - - /// Checks if two sets have no elements in common - #[inline] - pub fn is_disjoint(&self, other: &Self) -> bool { - self.set.is_disjoint(&other.set) - } - - /// Returns a new set containing only elements found in both sets - #[inline] - pub fn intersection<'a>(&'a self, other: &'a Self) -> impl Iterator { - self.set.intersection(&other.set) - } } -// implement extend -impl Extend for Set { +// Extend +impl Extend for Set { fn extend>(&mut self, iter: I) { let into_iter = iter.into_iter(); let reserve = if self.is_empty() { @@ -459,8 +471,8 @@ impl Extend for Set { } } -// implement extend for references -impl<'a, T: Eq + Hash + Clone + Debug> Extend<&'a T> for Set { +// Extend for References (Requires Clone) +impl<'a, T: Eq + Hash + Clone> Extend<&'a T> for Set { fn extend>(&mut self, iter: I) { let into_iter = iter.into_iter(); let reserve = if self.is_empty() { @@ -475,7 +487,7 @@ impl<'a, T: Eq + Hash + Clone + Debug> Extend<&'a T> for Set { } } -impl IntoIterator for Set { +impl IntoIterator for Set { type Item = T; type IntoIter = hashbrown::hash_set::IntoIter; @@ -483,17 +495,16 @@ impl IntoIterator for Set { self.set.into_iter() } } -impl FromIterator for Set { + +// Removed "Debug" and "Clone" requirements from FromIterator +impl FromIterator for Set { fn from_iter>(iter: I) -> Self { let mut set = Set::new(); - iter.into_iter().for_each(|x| { - set.insert(x); - }); + set.extend(iter); set } } -// implement `PartialEq` and `Eq` for `Set` impl PartialEq for Set { fn eq(&self, other: &Self) -> bool { if self.combined_hash != other.combined_hash { @@ -509,33 +520,39 @@ impl PartialOrd for Set { Some(self.cmp(other)) } } + +// Ord requires BTreeSet collection, so T must be Clone (to be collected) or we must iterate refs. +// Standard trick: collect references to sort. impl Ord for Set { fn cmp(&self, other: &Self) -> Ordering { if self.combined_hash != other.combined_hash { - return self.combined_hash.cmp(&other.combined_hash); // ✅ Compare hash first + return self.combined_hash.cmp(&other.combined_hash); } + // Collect references to avoid requiring T: Clone let self_sorted: BTreeSet<_> = self.set.iter().collect(); let other_sorted: BTreeSet<_> = other.set.iter().collect(); self_sorted.cmp(&other_sorted) } } -impl Default for Set { +impl Default for Set { fn default() -> Self { Self::new() } } -unsafe impl Send for Set {} -unsafe impl Sync for Set {} +// Corrected Send/Sync: We don't need Hash/Eq to move the Set between threads. +// We only need T to be Send/Sync. +unsafe impl Send for Set {} +unsafe impl Sync for Set {} -impl Hash for Set { +impl Hash for Set { fn hash(&self, state: &mut H) { self.combined_hash.hash(state); } } -impl From<[T; N]> for Set { +impl From<[T; N]> for Set { fn from(array: [T; N]) -> Self { array.into_iter().collect() } @@ -548,81 +565,9 @@ mod tests { #[test] fn test_insert_and_contains() { let mut set = Set::new(); - // Inserting a new element should return true. assert!(set.insert(1)); assert!(set.contains(&1)); - // Re-inserting the same element should return false. assert!(!set.insert(1)); assert_eq!(set.len(), 1); } - - #[test] - fn test_removal() { - let mut set = Set::new(); - set.insert(2); - set.insert(3); - // Remove existing element. - assert!(set.remove(&2)); - assert!(!set.contains(&2)); - assert_eq!(set.len(), 1); - // Removing a non-existent element should return false. - assert!(!set.remove(&2)); - } - - // #[test] - fn _test_iteration_order() { - let mut set = Set::new(); - set.insert(10); - set.insert(20); - set.insert(30); - // Expect the iteration order to match insertion order. - let elements: Vec<_> = set.iter().cloned().collect(); - assert_eq!(elements, vec![10, 20, 30]); - } - - #[test] - fn test_extend_and_append() { - let mut set1 = Set::new(); - set1.insert(1); - set1.insert(2); - - let mut set2 = Set::new(); - set2.insert(3); - set2.insert(4); - - // Append set2 into set1. - set1.append(&mut set2); - assert_eq!(set1.len(), 4); - assert!(set1.contains(&3)); - assert!(set1.contains(&4)); - // After appending, set2 should be empty. - assert!(set2.is_empty()); - } - - #[test] - fn test_intersection() { - let mut set1 = Set::new(); - set1.insert(1); - set1.insert(2); - set1.insert(3); - - let mut set2 = Set::new(); - set2.insert(2); - set2.insert(4); - - // The intersection should only contain the common element. - let inter: Vec<_> = set1.intersection(&set2).cloned().collect(); - assert_eq!(inter, vec![2]); - } - - #[test] - fn test_into_iter() { - let mut set = Set::new(); - set.insert(100); - set.insert(200); - // Collect the elements by consuming the set. - let mut collected: Vec<_> = set.into_iter().collect(); - collected.sort(); - assert_eq!(collected, vec![100, 200]); - } -} +} \ No newline at end of file diff --git a/src/invalid_subgraph.rs b/src/invalid_subgraph.rs index c4906d06..d91ec26c 100644 --- a/src/invalid_subgraph.rs +++ b/src/invalid_subgraph.rs @@ -1,8 +1,9 @@ -use crate::decoding_hypergraph::*; use crate::derivative::Derivative; use crate::matrix::*; use crate::plugin::EchelonMatrix; use crate::util::*; +use crate::dual_module_pq::{VertexPtr, EdgePtr}; +use crate::dual_module::DualModuleImpl; use std::cmp::Ordering; // use std::collections::hash_map::DefaultHasher; use std::hash::{Hash, Hasher}; @@ -10,17 +11,15 @@ use std::sync::Arc; /// an invalid subgraph $S = (V_S, E_S)$, also store the hair $\delta(S)$ #[derive(Clone, PartialEq, Eq, Derivative, Default)] -#[derivative(Debug)] pub struct InvalidSubgraph { /// the hash value calculated by other fields - #[derivative(Debug = "ignore")] pub hash_value: u64, /// subset of vertices - pub vertices: FastIterSet, + pub vertices: FastIterSet, /// subset of edges - pub edges: FastIterSet, + pub edges: FastIterSet, /// the hair of the invalid subgraph, to avoid repeated computation - pub hair: FastIterSet, + pub hair: FastIterSet, } impl Hash for InvalidSubgraph { @@ -48,43 +47,62 @@ impl PartialOrd for InvalidSubgraph { } } +impl std::fmt::Debug for InvalidSubgraph { + fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { + write!( + f, + "InvalidSubgraph:\nVertices: {:?}\nEdges: {:?}\nHair: {:?}\n", + self.vertices + .iter() + .map(|v| v.read_recursive().vertex_index) + .collect::>(), + self.edges + .iter() + .map(|e| e.read_recursive().edge_index) + .collect::>(), + self.hair + .iter() + .map(|e| e.read_recursive().edge_index) + .collect::>(), + ) + } +} + impl InvalidSubgraph { /// construct an invalid subgraph using only $E_S$, and constructing the $V_S$ by $\cup E_S$ #[allow(clippy::unnecessary_cast)] - pub fn new(edges: FastIterSet, decoding_graph: &DecodingHyperGraph) -> Self { + pub fn new(edges: FastIterSet, dual_module: &mut (impl DualModuleImpl + ?Sized)) -> Self { let mut vertices = FastIterSet::new(); - for &edge_index in edges.iter() { - let hyperedge = &decoding_graph.model_graph.initializer.weighted_edges[edge_index as usize]; - for &vertex_index in hyperedge.vertices.iter() { - vertices.insert(vertex_index); + for edge_ptr in edges.iter() { + for vertex_weak in edge_ptr.read_recursive().vertices.iter() { + vertices.insert(vertex_weak.upgrade_force()); } } - Self::new_complete(vertices, edges, decoding_graph) + Self::new_complete(vertices, edges, dual_module) } /// complete definition of invalid subgraph $S = (V_S, E_S)$ #[allow(clippy::unnecessary_cast)] pub fn new_complete( - vertices: FastIterSet, - edges: FastIterSet, - decoding_graph: &DecodingHyperGraph, + vertices: FastIterSet, + edges: FastIterSet, + dual_module: &mut (impl DualModuleImpl + ?Sized), ) -> Self { let mut hair = FastIterSet::new(); - for &vertex_index in vertices.iter() { - let vertex = &decoding_graph.model_graph.vertices[vertex_index as usize]; - for &edge_index in vertex.edges.iter() { - if !edges.contains(&edge_index) { - hair.insert(edge_index); + for vertex_ptr in vertices.iter() { + for edge_weak in vertex_ptr.read_recursive().edges.iter() { + if !edges.contains(&edge_weak.upgrade_force()) { + hair.insert(edge_weak.upgrade_force().clone()); } } } let invalid_subgraph = Self::new_raw(vertices, edges, hair); - debug_assert_eq!(invalid_subgraph.sanity_check(decoding_graph), Ok(())); + debug_assert_eq!(invalid_subgraph.sanity_check(dual_module), Ok(())); invalid_subgraph } /// create $S = (V_S, E_S)$ and $\delta(S)$ directly, without any checks - pub fn new_raw(vertices: FastIterSet, edges: FastIterSet, hair: FastIterSet) -> Self { + pub fn new_raw(vertices: FastIterSet, edges: FastIterSet, hair: FastIterSet) -> Self { let mut invalid_subgraph = Self { hash_value: 0, vertices, @@ -103,43 +121,84 @@ impl InvalidSubgraph { self.hash_value = hasher.finish(); } + pub fn new_from_indices(edges: FastIterSet, dual_module: &mut impl DualModuleImpl) -> Self { + let edges_ptr = edges.iter().map(|i| dual_module.get_edge_ptr(*i)).collect::>(); + let invalid_subgraph = Self::new(edges_ptr, dual_module); + debug_assert_eq!(invalid_subgraph.sanity_check(dual_module), Ok(())); + invalid_subgraph + } + + pub fn new_complete_from_indices( + vertices: FastIterSet, + edges: FastIterSet, + dual_module: &mut impl DualModuleImpl, + ) -> Self { + let vertices_ptr = vertices + .iter() + .map(|i| dual_module.get_vertex_ptr(*i)) + .collect::>(); + let edges_ptr = edges.iter().map(|i| dual_module.get_edge_ptr(*i)).collect::>(); + let invalid_subgraph = Self::new_complete(vertices_ptr, edges_ptr, dual_module); + debug_assert_eq!(invalid_subgraph.sanity_check(dual_module), Ok(())); + invalid_subgraph + } + + pub fn new_raw_from_indices( + vertices: FastIterSet, + edges: FastIterSet, + hair: FastIterSet, + dual_module: &mut impl DualModuleImpl, + ) -> Self { + let vertices_ptr = vertices + .iter() + .map(|i| dual_module.get_vertex_ptr(*i)) + .collect::>(); + let edges_ptr = edges.iter().map(|i| dual_module.get_edge_ptr(*i)).collect::>(); + let hair_ptr = hair.iter().map(|i| dual_module.get_edge_ptr(*i)).collect::>(); + Self::new_raw(vertices_ptr, edges_ptr, hair_ptr) + } + // check whether this invalid subgraph is indeed invalid, this is costly and should be disabled in release runs #[allow(clippy::unnecessary_cast)] - pub fn sanity_check(&self, decoding_graph: &DecodingHyperGraph) -> Result<(), String> { + pub fn sanity_check(&self, dual_module: &mut (impl DualModuleImpl + ?Sized)) -> Result<(), String> { if self.vertices.is_empty() { return Err("an invalid subgraph must contain at least one vertex".to_string()); } // check if all vertices are valid - for &vertex_index in self.vertices.iter() { - if vertex_index >= decoding_graph.model_graph.initializer.vertex_num { + for vertex_ptr in self.vertices.iter() { + let vertex_index = vertex_ptr.read_recursive().vertex_index; + if vertex_index >= dual_module.get_vertex_num() { return Err(format!("vertex {vertex_index} is not a vertex in the model graph")); } } // check if every edge is subset of its vertices - for &edge_index in self.edges.iter() { - if edge_index as usize >= decoding_graph.model_graph.initializer.weighted_edges.len() { + for edge_ptr in self.edges.iter() { + let edge = edge_ptr.read_recursive(); + let edge_index = edge.edge_index; + if edge_index as usize >= dual_module.get_edge_num() { return Err(format!("edge {edge_index} is not an edge in the model graph")); } - let hyperedge = &decoding_graph.model_graph.initializer.weighted_edges[edge_index as usize]; - for &vertex_index in hyperedge.vertices.iter() { - if !self.vertices.contains(&vertex_index) { + for vertex_ptr in edge.vertices.iter().map(|v| v.upgrade_force()) { + if !self.vertices.contains(&vertex_ptr) { + let vertex_index = vertex_ptr.read_recursive().vertex_index; return Err(format!( "hyperedge {edge_index} connects vertices {:?}, \ but vertex {vertex_index} is not in the invalid subgraph vertices {:?}", - hyperedge.vertices, self.vertices + edge.vertices, self.vertices )); } } } // check the edges indeed cannot satisfy the requirement of the vertices let mut matrix = Echelon::::new(); - for &edge_index in self.edges.iter() { - matrix.add_variable(edge_index); + for edge_ptr in self.edges.iter() { + matrix.add_variable(edge_ptr.downgrade().clone()); } - for &vertex_index in self.vertices.iter() { - let incident_edges = decoding_graph.get_vertex_neighbors(vertex_index); - let parity = decoding_graph.is_vertex_defect(vertex_index); - matrix.add_constraint(vertex_index, incident_edges, parity); + for vertex_ptr in self.vertices.iter() { + let vertex = vertex_ptr.read_recursive(); + let incident_edges = &vertex.edges; + let parity = vertex.is_defect; + matrix.add_constraint(vertex_ptr.downgrade().clone(), incident_edges, parity); } if matrix.get_echelon_info().satisfiable { return Err(format!( @@ -152,15 +211,16 @@ impl InvalidSubgraph { Ok(()) } - pub fn generate_matrix(&self, decoding_graph: &DecodingHyperGraph) -> EchelonMatrix { + pub fn generate_matrix(&self) -> EchelonMatrix { let mut matrix = EchelonMatrix::new(); - for &edge_index in self.hair.iter() { - matrix.add_variable(edge_index); + for edge_ptr in self.hair.iter() { + matrix.add_variable(edge_ptr.downgrade().clone()); } - for &vertex_index in self.vertices.iter() { - let incident_edges = decoding_graph.get_vertex_neighbors(vertex_index); - let parity = decoding_graph.is_vertex_defect(vertex_index); - matrix.add_constraint(vertex_index, incident_edges, parity); + for vertex_ptr in self.vertices.iter() { + let vertex = vertex_ptr.read_recursive(); + let incident_edges = &vertex.edges; + let parity = vertex.is_defect; + matrix.add_constraint(vertex_ptr.downgrade().clone(), incident_edges, parity); } matrix } @@ -168,28 +228,28 @@ impl InvalidSubgraph { // shortcuts for easier code writing at debugging impl InvalidSubgraph { - pub fn new_ptr(edges: FastIterSet, decoding_graph: &DecodingHyperGraph) -> Arc { - Arc::new(Self::new(edges, decoding_graph)) + pub fn new_ptr(edges: FastIterSet, dual_module: &mut impl DualModuleImpl) -> Arc { + Arc::new(Self::new(edges, dual_module)) } - pub fn new_vec_ptr(edges: &[EdgeIndex], decoding_graph: &DecodingHyperGraph) -> Arc { - Self::new_ptr(edges.iter().cloned().collect(), decoding_graph) + pub fn new_vec_ptr(edges: &[EdgePtr], dual_module: &mut impl DualModuleImpl) -> Arc { + Self::new_ptr(edges.iter().cloned().collect(), dual_module) } pub fn new_complete_ptr( - vertices: FastIterSet, - edges: FastIterSet, - decoding_graph: &DecodingHyperGraph, + vertices: FastIterSet, + edges: FastIterSet, + dual_module: &mut (impl DualModuleImpl + ?Sized), ) -> Arc { - Arc::new(Self::new_complete(vertices, edges, decoding_graph)) + Arc::new(Self::new_complete(vertices, edges, dual_module)) } pub fn new_complete_vec_ptr( - vertices: FastIterSet, - edges: &[EdgeIndex], - decoding_graph: &DecodingHyperGraph, + vertices: FastIterSet, + edges: &[EdgePtr], + dual_module: &mut impl DualModuleImpl, ) -> Arc { Self::new_complete_ptr( vertices.iter().cloned().collect(), edges.iter().cloned().collect(), - decoding_graph, + dual_module, ) } } @@ -198,19 +258,45 @@ impl InvalidSubgraph { pub mod tests { use super::*; use crate::decoding_hypergraph::tests::*; - + use crate::dual_module::DualModuleInterfacePtr; + use crate::dual_module_pq::DualModulePQ; + use crate::dual_module_pq::{EdgePtr, Vertex, VertexPtr}; + use std::collections::HashSet; + #[test] fn invalid_subgraph_good() { // cargo test invalid_subgraph_good -- --nocapture let visualize_filename = "invalid_subgraph_good.json".to_string(); let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); - let invalid_subgraph_1 = InvalidSubgraph::new(vec![13].into_iter().collect(), decoding_graph.as_ref()); + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); + interface_ptr.load(decoding_graph.syndrome_pattern.clone(), &mut dual_module, 0); // this is needed to load the defect vertices + let invalid_subgraph_1 = InvalidSubgraph::new_from_indices(fast_iter_set! {13}, &mut dual_module); println!("invalid_subgraph_1: {invalid_subgraph_1:?}"); - assert_eq!(sorted_vec(invalid_subgraph_1.vertices.into_iter().collect()), vec![2, 6, 7]); - assert_eq!(sorted_vec(invalid_subgraph_1.edges.into_iter().collect()), vec![13]); assert_eq!( - sorted_vec(invalid_subgraph_1.hair.into_iter().collect()), - vec![5, 6, 9, 10, 11, 12, 14, 15, 16, 17] + invalid_subgraph_1 + .vertices + .iter() + .map(|v| v.read_recursive().vertex_index) + .collect::>(), + vec![2, 6, 7].into_iter().collect::>() + ); + assert_eq!( + invalid_subgraph_1 + .edges + .iter() + .map(|e| e.read_recursive().edge_index) + .collect::>(), + vec![13].into_iter().collect::>() + ); + assert_eq!( + invalid_subgraph_1 + .hair + .iter() + .map(|e| e.read_recursive().edge_index) + .collect::>(), + vec![5, 6, 9, 10, 11, 12, 14, 15, 16, 17].into_iter().collect::>() ); } @@ -220,7 +306,13 @@ pub mod tests { // cargo test invalid_subgraph_bad -- --nocapture let visualize_filename = "invalid_subgraph_bad.json".to_string(); let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); - let invalid_subgraph = InvalidSubgraph::new(vec![6, 10].into_iter().collect(), decoding_graph.as_ref()); + println!("hello1!"); + let initializer = decoding_graph.model_graph.initializer.clone(); + println!("hello2!"); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + println!("hello3!"); + let invalid_subgraph = InvalidSubgraph::new_from_indices(fast_iter_set! {6, 10}, &mut dual_module); + println!("hello4!"); println!("invalid_subgraph: {invalid_subgraph:?}"); // should not print because it panics } @@ -233,11 +325,17 @@ pub mod tests { #[test] fn invalid_subgraph_hash() { // cargo test invalid_subgraph_hash -- --nocapture + let visualize_filename = "invalid_subgraph_hash.json".to_string(); + // we use an arbitrary decoding graph, the defect vertices here are not loaded + let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let vertices: FastIterSet = [1, 2, 3].into(); let edges: FastIterSet = [4, 5].into(); let hair: FastIterSet = [6, 7, 8].into(); - let invalid_subgraph_1 = InvalidSubgraph::new_raw(vertices.clone(), edges.clone(), hair.clone()); - let invalid_subgraph_2 = InvalidSubgraph::new_raw(vertices.clone(), edges.clone(), hair.clone()); + let invalid_subgraph_1 = InvalidSubgraph::new_raw_from_indices(vertices.clone(), edges.clone(), hair.clone(), &mut dual_module); + let invalid_subgraph_2 = InvalidSubgraph::new_raw_from_indices(vertices.clone(), edges.clone(), hair.clone(), &mut dual_module); assert_eq!(invalid_subgraph_1, invalid_subgraph_2); // they should have the same hash value assert_eq!( @@ -256,15 +354,15 @@ pub mod tests { // any different value would generate a different invalid subgraph assert_ne!( invalid_subgraph_1, - InvalidSubgraph::new_raw([1, 2].into(), edges.clone(), hair.clone()) + InvalidSubgraph::new_raw_from_indices([1, 2].into(), edges.clone(), hair.clone(), &mut dual_module) ); assert_ne!( invalid_subgraph_1, - InvalidSubgraph::new_raw(vertices.clone(), [4, 5, 6].into(), hair.clone()) + InvalidSubgraph::new_raw_from_indices(vertices.clone(), [4, 5, 6].into(), hair.clone(), &mut dual_module) ); assert_ne!( invalid_subgraph_1, - InvalidSubgraph::new_raw(vertices.clone(), edges.clone(), [6, 7].into()) + InvalidSubgraph::new_raw_from_indices(vertices.clone(), edges.clone(), [6, 7].into(), &mut dual_module) ); } } diff --git a/src/lib.rs b/src/lib.rs index d6150088..17e8ea33 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -96,13 +96,13 @@ pub fn get_version() -> String { let code = CodeCapacityTailoredCode::new(7, 0., 0.01); // create dual module let model_graph = code.get_model_graph(); - let mut dual_module = DualModulePQ::new_empty(&model_graph.initializer); + let mut dual_module = DualModulePQ::new_empty(&model_graph.initializer, 0); // create primal module let mut primal_module = PrimalModuleSerial::new_empty(&model_graph.initializer); primal_module.plugins = std::sync::Arc::new(vec![]); // try to work on a simple syndrome let decoding_graph = DecodingHyperGraph::new_defects(model_graph, defect_vertices.clone()); - let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone()); + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); primal_module.solve_visualizer( &interface_ptr, decoding_graph.syndrome_pattern.clone(), diff --git a/src/matrix/basic.rs b/src/matrix/basic.rs index b76c417b..e84f471c 100644 --- a/src/matrix/basic.rs +++ b/src/matrix/basic.rs @@ -3,53 +3,54 @@ use super::row::*; use super::visualize::*; use crate::util::*; use derivative::Derivative; +use crate::dual_module_pq::{EdgeWeak, VertexWeak}; #[derive(Clone, Derivative, PartialEq, Eq)] #[derivative(Default(new = "true"))] pub struct BasicMatrix { /// the vertices already maintained by this parity check - pub vertices: FastIterSet, + pub vertices: FastIterSet, /// the edges maintained by this parity check, mapping to the local indices - pub edges: FastIterMap, + pub edges: FastIterMap, /// variable index map to edge index - pub variables: Vec, + pub variables: Vec, pub constraints: Vec, } impl MatrixBasic for BasicMatrix { - fn add_variable(&mut self, edge_index: EdgeIndex) -> Option { - if self.edges.contains_key(&edge_index) { + fn add_variable(&mut self, edge_weak: EdgeWeak) -> Option { + if self.edges.contains_key(&edge_weak.clone()) { // variable already exists return None; } let var_index = self.variables.len(); - self.edges.insert(edge_index, var_index); - self.variables.push(edge_index); + self.edges.insert(edge_weak.clone(), var_index); + self.variables.push(edge_weak.clone()); ParityRow::add_one_variable(&mut self.constraints, self.variables.len()); Some(var_index) } fn add_constraint( &mut self, - vertex_index: VertexIndex, - incident_edges: &[EdgeIndex], + vertex_weak: VertexWeak, + incident_edges: &[EdgeWeak], parity: bool, ) -> Option> { - if self.vertices.contains(&vertex_index) { + if self.vertices.contains(&vertex_weak) { // no need to add repeat constraint return None; } let mut var_indices = None; - self.vertices.insert(vertex_index); - for &edge_index in incident_edges.iter() { - if let Some(var_index) = self.add_variable(edge_index) { + self.vertices.insert(vertex_weak.clone()); + for edge_weak in incident_edges.iter() { + if let Some(var_index) = self.add_variable(edge_weak.clone()) { // this is a newly added edge var_indices.get_or_insert_with(Vec::new).push(var_index); } } let mut row = ParityRow::new_length(self.variables.len()); - for &edge_index in incident_edges.iter() { - let var_index = self.edges[&edge_index]; + for edge_weak in incident_edges.iter() { + let var_index = self.edges[&edge_weak.clone()]; row.set_left(var_index, true); } row.set_right(parity); @@ -74,19 +75,19 @@ impl MatrixBasic for BasicMatrix { self.constraints[row].get_right() } - fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeIndex { - self.variables[var_index] + fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeWeak { + self.variables[var_index].clone() } - fn edge_to_var_index(&self, edge_index: EdgeIndex) -> Option { - self.edges.get(&edge_index).cloned() + fn edge_to_var_index(&self, edge_weak: EdgeWeak) -> Option { + self.edges.get(&edge_weak.clone()).cloned() } - fn get_vertices(&self) -> FastIterSet { + fn get_vertices(&self) -> FastIterSet { self.vertices.clone() } - fn get_edges(&self) -> FastIterSet { + fn get_edges(&self) -> FastIterSet { self.edges.keys().cloned().collect() } } @@ -114,164 +115,156 @@ impl VizTrait for BasicMatrix { #[cfg(test)] pub mod tests { use super::*; + use crate::dual_module_pq::{Edge, EdgePtr, Vertex, VertexPtr}; + use crate::num_traits::Zero; + use std::collections::HashSet; + + /// Helper to create mock pointers for testing + pub fn initialize_vertex_edges_for_matrix_testing( + vertex_indices: Vec, + edge_indices: Vec, + ) -> (Vec, Vec) { + let edges: Vec = edge_indices + .into_iter() + .map(|edge_index| { + EdgePtr::new_value( + Edge { + edge_index, + weight: Rational::zero(), + dual_nodes: vec![], + vertices: vec![], + last_updated_time: Rational::zero(), + growth_at_last_updated_time: Rational::zero(), + grow_rate: Rational::zero(), + // unit_index: Some(0), + // connected_to_boundary_vertex: false, + #[cfg(feature = "incr_lp")] + cluster_weights: hashbrown::HashMap::new(), + }, + (0, edge_index), + ) + }) + .collect(); + + let vertices: Vec = vertex_indices + .into_iter() + .map(|vertex_index| { + VertexPtr::new_value( + Vertex { + vertex_index, + is_defect: false, + edges: vec![], + // mirrored_vertices: vec![], + }, + (0, vertex_index), + ) + }) + .collect(); + + (vertices, edges) + } + + pub fn edge_vec_from_indices(edge_sequences: &[usize], edges: &Vec) -> Vec { + edge_sequences + .iter() + .map(|&i| edges[i].downgrade()) + .collect() + } #[test] fn basic_matrix_1() { - // cargo test --features=colorful basic_matrix_1 -- --nocapture let mut matrix = BasicMatrix::new(); - matrix.printstd(); - assert_eq!( - matrix.printstd_str(), - "\ -┌┬───┐ -┊┊ = ┊ -╞╪═══╡ -└┴───┘ -" - ); - matrix.add_variable(1); - matrix.add_variable(4); - matrix.add_variable(12); - matrix.add_variable(345); - matrix.printstd(); - assert_eq!( - matrix.printstd_str(), - "\ -┌┬─┬─┬─┬─┬───┐ -┊┊1┊4┊1┊3┊ = ┊ -┊┊ ┊ ┊2┊4┊ ┊ -┊┊ ┊ ┊ ┊5┊ ┊ -╞╪═╪═╪═╪═╪═══╡ -└┴─┴─┴─┴─┴───┘ -" + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing( + vec![0, 1, 2], + vec![1, 4, 12, 345] ); - matrix.add_constraint(0, &[1, 4, 12], true); - matrix.add_constraint(1, &[4, 345], false); - matrix.add_constraint(2, &[1, 345], true); + + // Add variables + for edge in edges.iter() { + matrix.add_variable(edge.downgrade()); + } + + // Add constraints using sequence indices from the 'edges' vector + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&[0, 1, 2], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&[1, 3], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&[0, 3], &edges), true); + matrix.printstd(); - assert_eq!( - matrix.clone().printstd_str(), - "\ -┌─┬─┬─┬─┬─┬───┐ -┊ ┊1┊4┊1┊3┊ = ┊ -┊ ┊ ┊ ┊2┊4┊ ┊ -┊ ┊ ┊ ┊ ┊5┊ ┊ -╞═╪═╪═╪═╪═╪═══╡ -┊0┊1┊1┊1┊ ┊ 1 ┊ -├─┼─┼─┼─┼─┼───┤ -┊1┊ ┊1┊ ┊1┊ ┊ -├─┼─┼─┼─┼─┼───┤ -┊2┊1┊ ┊ ┊1┊ 1 ┊ -└─┴─┴─┴─┴─┴───┘ -" - ); - assert_eq!(matrix.get_vertices(), [0, 1, 2].into()); - assert_eq!(matrix.get_view_edges(), [1, 4, 12, 345]); + + let vertex_indices: HashSet<_> = matrix.get_vertices().iter() + .map(|v| v.upgrade_force().read_recursive().vertex_index).collect(); + assert_eq!(vertex_indices, [0, 1, 2].into_iter().collect()); } #[test] fn basic_matrix_should_not_add_repeated_constraint() { - // cargo test --features=colorful basic_matrix_should_not_add_repeated_constraint -- --nocapture let mut matrix = BasicMatrix::new(); - assert_eq!(matrix.add_constraint(0, &[1, 4, 8], false), Some(vec![0, 1, 2])); - assert_eq!(matrix.add_constraint(1, &[4, 8], true), None); - assert_eq!(matrix.add_constraint(0, &[4], true), None); // repeated - matrix.printstd(); + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing( + vec![0, 1], + vec![1, 4, 8] + ); + + // First add: Success (returns indices of newly created variables 0, 1, 2) + assert_eq!( + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&[0, 1, 2], &edges), false), + Some(vec![0, 1, 2]) + ); + + // Second add (new vertex): Success (returns None because no NEW variables were created) + assert_eq!( + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&[1, 2], &edges), true), + None + ); + + // Third add (repeat vertex): Success is None because vertex already exists assert_eq!( - matrix.clone().printstd_str(), - "\ -┌─┬─┬─┬─┬───┐ -┊ ┊1┊4┊8┊ = ┊ -╞═╪═╪═╪═╪═══╡ -┊0┊1┊1┊1┊ ┊ -├─┼─┼─┼─┼───┤ -┊1┊ ┊1┊1┊ 1 ┊ -└─┴─┴─┴─┴───┘ -" + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&[1], &edges), true), + None ); } #[test] fn basic_matrix_row_operations() { - // cargo test --features=colorful basic_matrix_row_operations -- --nocapture let mut matrix = BasicMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); - matrix.printstd(); - assert_eq!( - matrix.clone().printstd_str(), - "\ -┌─┬─┬─┬─┬─┬───┐ -┊ ┊1┊4┊6┊9┊ = ┊ -╞═╪═╪═╪═╪═╪═══╡ -┊0┊1┊1┊1┊ ┊ 1 ┊ -├─┼─┼─┼─┼─┼───┤ -┊1┊ ┊1┊ ┊1┊ ┊ -├─┼─┼─┼─┼─┼───┤ -┊2┊1┊ ┊ ┊1┊ 1 ┊ -└─┴─┴─┴─┴─┴───┘ -" + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing( + vec![0, 1, 2], + vec![1, 4, 6, 9] ); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&[0, 1, 2], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&[1, 3], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&[0, 3], &edges), true); + + // Test Swap matrix.swap_row(2, 1); - matrix.printstd(); - assert_eq!( - matrix.clone().printstd_str(), - "\ -┌─┬─┬─┬─┬─┬───┐ -┊ ┊1┊4┊6┊9┊ = ┊ -╞═╪═╪═╪═╪═╪═══╡ -┊0┊1┊1┊1┊ ┊ 1 ┊ -├─┼─┼─┼─┼─┼───┤ -┊1┊1┊ ┊ ┊1┊ 1 ┊ -├─┼─┼─┼─┼─┼───┤ -┊2┊ ┊1┊ ┊1┊ ┊ -└─┴─┴─┴─┴─┴───┘ -" - ); + assert_eq!(matrix.get_rhs(1), true); // Row 2 was true, now at index 1 + + // Test XOR matrix.xor_row(0, 1); - matrix.printstd(); - assert_eq!( - matrix.clone().printstd_str(), - "\ -┌─┬─┬─┬─┬─┬───┐ -┊ ┊1┊4┊6┊9┊ = ┊ -╞═╪═╪═╪═╪═╪═══╡ -┊0┊ ┊1┊1┊1┊ ┊ -├─┼─┼─┼─┼─┼───┤ -┊1┊1┊ ┊ ┊1┊ 1 ┊ -├─┼─┼─┼─┼─┼───┤ -┊2┊ ┊1┊ ┊1┊ ┊ -└─┴─┴─┴─┴─┴───┘ -" - ); + // Column 0: Row 0 (1) XOR Row 1 (1) = 0 + assert_eq!(matrix.get_lhs(0, 0), false); } #[test] fn basic_matrix_manual_echelon() { - // cargo test --features=colorful basic_matrix_manual_echelon -- --nocapture let mut matrix = BasicMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing( + vec![0, 1, 2], + vec![1, 4, 6, 9] + ); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&[0, 1, 2], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&[1, 3], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&[0, 3], &edges), true); + matrix.xor_row(2, 0); matrix.xor_row(0, 1); matrix.xor_row(2, 1); matrix.xor_row(0, 2); - matrix.printstd(); - assert_eq!( - matrix.clone().printstd_str(), - "\ -┌─┬─┬─┬─┬─┬───┐ -┊ ┊1┊4┊6┊9┊ = ┊ -╞═╪═╪═╪═╪═╪═══╡ -┊0┊1┊ ┊ ┊1┊ 1 ┊ -├─┼─┼─┼─┼─┼───┤ -┊1┊ ┊1┊ ┊1┊ ┊ -├─┼─┼─┼─┼─┼───┤ -┊2┊ ┊ ┊1┊ ┊ ┊ -└─┴─┴─┴─┴─┴───┘ -" - ); + + // Verify specific cell after operations (Row 0, Var 0 should be 1) + assert_eq!(matrix.get_lhs(0, 0), true); + // Row 2, Var 2 (edge 6) should be 1 + assert_eq!(matrix.get_lhs(2, 2), true); } -} +} \ No newline at end of file diff --git a/src/matrix/complete.rs b/src/matrix/complete.rs index 4fc85222..ad0c379d 100644 --- a/src/matrix/complete.rs +++ b/src/matrix/complete.rs @@ -3,23 +3,24 @@ use super::row::*; use super::visualize::*; use crate::util::*; use derivative::Derivative; +use crate::dual_module_pq::{EdgeWeak, VertexWeak}; /// complete matrix considers a predefined set of edges and won't consider any other edges #[derive(Clone, Derivative)] #[derivative(Default(new = "true"))] pub struct CompleteMatrix { /// the vertices already maintained by this parity check - vertices: FastIterSet, + vertices: FastIterSet, /// the edges maintained by this parity check, mapping to the local indices - edges: FastIterMap, + edges: FastIterMap, /// variable index map to edge index - variables: Vec, + variables: Vec, constraints: Vec, } impl MatrixBasic for CompleteMatrix { - fn add_variable(&mut self, edge_index: EdgeIndex) -> Option { - if self.edges.contains_key(&edge_index) { + fn add_variable(&mut self, edge_weak: EdgeWeak) -> Option { + if self.edges.contains_key(&edge_weak) { // variable already exists return None; } @@ -27,26 +28,26 @@ impl MatrixBasic for CompleteMatrix { panic!("complete matrix doesn't allow dynamic edges, please insert all edges at the beginning") } let var_index = self.variables.len(); - self.edges.insert(edge_index, var_index); - self.variables.push(edge_index); + self.edges.insert(edge_weak.clone(), var_index); + self.variables.push(edge_weak.clone()); Some(var_index) } fn add_constraint( &mut self, - vertex_index: VertexIndex, - incident_edges: &[EdgeIndex], + vertex_weak: VertexWeak, + incident_edges: &[EdgeWeak], parity: bool, ) -> Option> { - if self.vertices.contains(&vertex_index) { + if self.vertices.contains(&vertex_weak) { // no need to add repeat constraint return None; } - self.vertices.insert(vertex_index); + self.vertices.insert(vertex_weak.clone()); let mut row = ParityRow::new_length(self.variables.len()); - for &edge_index in incident_edges.iter() { - if self.exists_edge(edge_index) { - let var_index = self.edges[&edge_index]; + for edge_weak in incident_edges.iter() { + if self.exists_edge(edge_weak.clone()) { + let var_index = self.edges[&edge_weak.clone()]; row.set_left(var_index, true); } } @@ -73,19 +74,19 @@ impl MatrixBasic for CompleteMatrix { self.constraints[row].get_right() } - fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeIndex { - self.variables[var_index] + fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeWeak { + self.variables[var_index].clone() } - fn edge_to_var_index(&self, edge_index: EdgeIndex) -> Option { - self.edges.get(&edge_index).cloned() + fn edge_to_var_index(&self, edge_weak: EdgeWeak) -> Option { + self.edges.get(&edge_weak.clone()).cloned() } - fn get_vertices(&self) -> FastIterSet { + fn get_vertices(&self) -> FastIterSet { self.vertices.clone() } - fn get_edges(&self) -> FastIterSet { + fn get_edges(&self) -> FastIterSet { self.edges.keys().cloned().collect() } } @@ -113,6 +114,8 @@ impl VizTrait for CompleteMatrix { #[cfg(test)] pub mod tests { use crate::matrix::Echelon; + use crate::matrix::basic::tests::{initialize_vertex_edges_for_matrix_testing, edge_vec_from_indices}; + use std::collections::HashSet; use super::*; @@ -120,8 +123,17 @@ pub mod tests { fn complete_matrix_1() { // cargo test --features=colorful complete_matrix_1 -- --nocapture let mut matrix = CompleteMatrix::new(); - for edge_index in [1, 4, 12, 345] { - matrix.add_variable(edge_index); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 12, 345]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + for edge_ptr in edges.iter() { + matrix.add_variable(edge_ptr.downgrade()); } matrix.printstd(); assert_eq!( @@ -135,9 +147,9 @@ pub mod tests { └┴─┴─┴─┴─┴───┘ " ); - matrix.add_constraint(0, &[1, 4, 12], true); - matrix.add_constraint(1, &[4, 345], false); - matrix.add_constraint(2, &[1, 345], true); + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -155,20 +167,34 @@ pub mod tests { └─┴─┴─┴─┴─┴───┘ " ); - assert_eq!(matrix.get_vertices(), [0, 1, 2].into()); - assert_eq!(matrix.get_view_edges(), [1, 4, 12, 345]); + assert_eq!( + matrix.get_vertices().iter().map(|v| v.upgrade_force().read_recursive().vertex_index).collect::>(), + [0, 1, 2].into_iter().collect::>()); + assert_eq!( + matrix.get_view_edges().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + [1, 4, 12, 345].into_iter().collect::>()); } #[test] fn complete_matrix_should_not_add_repeated_constraint() { // cargo test --features=colorful complete_matrix_should_not_add_repeated_constraint -- --nocapture let mut matrix = CompleteMatrix::new(); - for edge_index in [1, 4, 8] { - matrix.add_variable(edge_index); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 8]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 2], + vec![1], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + for edge_ptr in edges.iter() { + matrix.add_variable(edge_ptr.downgrade()); } - assert_eq!(matrix.add_constraint(0, &[1, 4, 8], false), None); - assert_eq!(matrix.add_constraint(1, &[4, 8], true), None); - assert_eq!(matrix.add_constraint(0, &[4], true), None); // repeated + + assert_eq!(matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), false), None); + assert_eq!(matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), true), None); + assert_eq!(matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true), None); // repeated matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -188,12 +214,21 @@ pub mod tests { fn complete_matrix_row_operations() { // cargo test --features=colorful complete_matrix_row_operations -- --nocapture let mut matrix = CompleteMatrix::new(); - for edge_index in [1, 4, 6, 9] { - matrix.add_variable(edge_index); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + for edge_ptr in edges.iter() { + matrix.add_variable(edge_ptr.downgrade()); } - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -247,12 +282,25 @@ pub mod tests { fn complete_matrix_manual_echelon() { // cargo test --features=colorful complete_matrix_manual_echelon -- --nocapture let mut matrix = CompleteMatrix::new(); - for edge_index in [1, 4, 6, 9, 9, 6, 4, 1] { - matrix.add_variable(edge_index); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + // add variables [1, 4, 6, 9, 9, 6, 4, 1] + for edge_ptr in edges.iter() { + matrix.add_variable(edge_ptr.downgrade()); } - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); + for edge_ptr in edges.clone().into_iter().rev() { + matrix.add_variable(edge_ptr.downgrade()); + } + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); matrix.xor_row(2, 0); matrix.xor_row(0, 1); matrix.xor_row(2, 1); @@ -278,12 +326,21 @@ pub mod tests { fn complete_matrix_automatic_echelon() { // cargo test --features=colorful complete_matrix_automatic_echelon -- --nocapture let mut matrix = Echelon::::new(); - for edge_index in [1, 4, 6, 9] { - matrix.add_variable(edge_index); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9, 11, 12, 23]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2, 4, 5], + vec![1, 3, 6, 5], + vec![0, 3, 4], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + for edge_index in 0..4 { + matrix.add_variable(edges[edge_index].downgrade()); } - matrix.add_constraint(0, &[1, 4, 6, 11, 12], true); - matrix.add_constraint(1, &[4, 9, 23, 12], false); - matrix.add_constraint(2, &[1, 9, 11], true); + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -308,12 +365,21 @@ pub mod tests { fn complete_matrix_dynamic_variables_forbidden() { // cargo test complete_matrix_dynamic_variables_forbidden -- --nocapture let mut matrix = Echelon::::new(); - for edge_index in [1, 4, 6, 9] { - matrix.add_variable(edge_index); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9, 2]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + for edge_index in 0..4 { + matrix.add_variable(edges[edge_index].downgrade()); } - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); - matrix.add_variable(2); + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + matrix.add_variable(edges[4].downgrade()); } } diff --git a/src/matrix/echelon.rs b/src/matrix/echelon.rs index 2b09c364..170991e3 100644 --- a/src/matrix/echelon.rs +++ b/src/matrix/echelon.rs @@ -4,6 +4,7 @@ use crate::util::*; use core::panic; use derivative::Derivative; use prettytable::*; +use crate::dual_module_pq::{EdgeWeak, VertexWeak}; #[derive(Clone, Derivative)] #[derivative(Default(new = "true"))] @@ -32,42 +33,42 @@ impl Echelon { } impl MatrixTail for Echelon { - fn get_tail_edges(&self) -> &FastIterSet { + fn get_tail_edges(&self) -> &FastIterSet { self.base.get_tail_edges() } - fn get_tail_edges_mut(&mut self) -> &mut FastIterSet { + fn get_tail_edges_mut(&mut self) -> &mut FastIterSet { self.is_info_outdated = true; self.base.get_tail_edges_mut() } } impl MatrixTight for Echelon { - fn update_edge_tightness(&mut self, edge_index: EdgeIndex, is_tight: bool) { + fn update_edge_tightness(&mut self, edge_weak: EdgeWeak, is_tight: bool) { self.is_info_outdated = true; - self.base.update_edge_tightness(edge_index, is_tight) + self.base.update_edge_tightness(edge_weak.clone(), is_tight) } - fn is_tight(&self, edge_index: usize) -> bool { - self.base.is_tight(edge_index) + fn is_tight(&self, edge_weak: EdgeWeak) -> bool { + self.base.is_tight(edge_weak.clone()) } - fn get_tight_edges(&self) -> &FastIterSet { + fn get_tight_edges(&self) -> &FastIterSet { self.base.get_tight_edges() } } impl MatrixBasic for Echelon { - fn add_variable(&mut self, edge_index: EdgeIndex) -> Option { + fn add_variable(&mut self, edge_weak: EdgeWeak) -> Option { self.is_info_outdated = true; - self.base.add_variable(edge_index) + self.base.add_variable(edge_weak.clone()) } fn add_constraint( &mut self, - vertex_index: VertexIndex, - incident_edges: &[EdgeIndex], + vertex_weak: VertexWeak, + incident_edges: &[EdgeWeak], parity: bool, ) -> Option> { self.is_info_outdated = true; - self.base.add_constraint(vertex_index, incident_edges, parity) + self.base.add_constraint(vertex_weak.clone(), incident_edges, parity) } fn xor_row(&mut self, _target: RowIndex, _source: RowIndex) { @@ -82,16 +83,16 @@ impl MatrixBasic for Echelon { fn get_rhs(&self, row: RowIndex) -> bool { self.get_base().get_rhs(row) } - fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeIndex { + fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeWeak { self.get_base().var_to_edge_index(var_index) } - fn edge_to_var_index(&self, edge_index: EdgeIndex) -> Option { - self.get_base().edge_to_var_index(edge_index) + fn edge_to_var_index(&self, edge_weak: EdgeWeak) -> Option { + self.get_base().edge_to_var_index(edge_weak.clone()) } - fn get_vertices(&self) -> FastIterSet { + fn get_vertices(&self) -> FastIterSet { self.get_base().get_vertices() } - fn get_edges(&self) -> FastIterSet { + fn get_edges(&self) -> FastIterSet { self.get_base().get_edges() } } @@ -274,7 +275,7 @@ impl VizTrait for Echelon { table.title.add_cell(Cell::new("\u{25BC}")); for (row, row_info) in info.rows.iter().enumerate() { let cell = if row_info.has_leading() { - Cell::new(self.column_to_edge_index(row_info.column).to_string().as_str()).style_spec("irFm") + Cell::new(self.column_to_edge_index(row_info.column).upgrade_force().read_recursive().edge_index.to_string().as_str()).style_spec("irFm") } else { Cell::new("*").style_spec("rFr") }; @@ -305,6 +306,9 @@ pub mod tests { use super::super::tight::*; use super::*; use crate::rand::{Rng, SeedableRng}; + use crate::matrix::basic::tests::{initialize_vertex_edges_for_matrix_testing, edge_vec_from_indices}; + use std::collections::HashSet; + use crate::dual_module_pq::{EdgePtr, VertexPtr}; type EchelonMatrix = Echelon>>; @@ -312,12 +316,22 @@ pub mod tests { fn echelon_matrix_simple() { // cargo test --features=colorful echelon_matrix_simple -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); - assert_eq!(matrix.edge_to_var_index(4), Some(1)); - for edge_index in [1, 4, 6, 9] { - matrix.update_edge_tightness(edge_index, true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + + assert_eq!(matrix.edge_to_var_index(edges[1].downgrade()), Some(1)); + for edge_index in 0..4 { + matrix.update_edge_tightness(edges[edge_index].downgrade(), true); } matrix.printstd(); assert_eq!( @@ -336,8 +350,10 @@ pub mod tests { └──┴─┴─┴─┴─┴───┴─┘ " ); - matrix.set_tail_edges([6, 1].into_iter()); - assert_eq!(matrix.get_tail_edges_vec(), [1, 6]); + matrix.set_tail_edges(edge_vec_from_indices(&vec![2, 0], &edges).into_iter()); + assert_eq!( + matrix.get_tail_edges_vec().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + [1, 6].into_iter().collect::>()); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -355,7 +371,7 @@ pub mod tests { └──┴─┴─┴─┴─┴───┴─┘ " ); - matrix.set_tail_edges([4].into_iter()); + matrix.set_tail_edges(edge_vec_from_indices(&vec![1], &edges).into_iter()); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -373,7 +389,7 @@ pub mod tests { └──┴─┴─┴─┴─┴───┴─┘ " ); - matrix.update_edge_tightness(6, false); + matrix.update_edge_tightness(edges[2].downgrade(), false); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -389,8 +405,8 @@ pub mod tests { └──┴─┴─┴─┴───┴─┘ " ); - matrix.update_edge_tightness(1, false); - matrix.update_edge_tightness(9, false); + matrix.update_edge_tightness(edges[0].downgrade(), false); + matrix.update_edge_tightness(edges[3].downgrade(), false); matrix.printstd(); } @@ -399,8 +415,16 @@ pub mod tests { fn echelon_matrix_should_not_xor() { // cargo test echelon_matrix_should_not_xor -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); matrix.xor_row(0, 1); } @@ -409,8 +433,16 @@ pub mod tests { fn echelon_matrix_should_not_swap() { // cargo test echelon_matrix_should_not_swap -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); matrix.swap_row(0, 1); } @@ -418,12 +450,22 @@ pub mod tests { fn echelon_matrix_basic_trait() { // cargo test --features=colorful echelon_matrix_basic_trait -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_variable(3); // un-tight edges will not show - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); - for edge_index in [1, 4, 6, 9] { - matrix.update_edge_tightness(edge_index, true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9, 3]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_variable(edges[4].downgrade()); // un-tight edges will not show + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + + for edge_index in 0..4 { + matrix.update_edge_tightness(edges[edge_index].downgrade(), true); } matrix.printstd(); assert_eq!( @@ -442,8 +484,8 @@ pub mod tests { └──┴─┴─┴─┴─┴───┴─┘ " ); - assert!(matrix.is_tight(1)); - assert_eq!(matrix.edge_to_var_index(4), Some(2)); + assert!(matrix.is_tight(edges[0].downgrade())); + assert_eq!(matrix.edge_to_var_index(edges[1].downgrade()), Some(2)); } #[test] @@ -451,8 +493,15 @@ pub mod tests { fn echelon_matrix_cannot_call_dirty_column() { // cargo test echelon_matrix_cannot_call_dirty_column -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.update_edge_tightness(1, true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.update_edge_tightness(edges[0].downgrade(), true); // even though there is indeed such a column, we forbid such dangerous calls // always call `columns()` before accessing any column matrix.column_to_var_index(0); @@ -463,8 +512,15 @@ pub mod tests { fn echelon_matrix_cannot_call_dirty_echelon_info() { // cargo test echelon_matrix_cannot_call_dirty_echelon_info -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.update_edge_tightness(1, true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.update_edge_tightness(edges[0].downgrade(), true); // even though there is indeed such a column, we forbid such dangerous calls // always call `columns()` before accessing any column matrix.get_echelon_info_immutable(); @@ -496,7 +552,14 @@ pub mod tests { fn echelon_matrix_no_variable_satisfiable() { // cargo test --features=colorful echelon_matrix_no_variable_satisfiable -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], false); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), false); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -519,7 +582,14 @@ pub mod tests { fn echelon_matrix_no_variable_unsatisfiable() { // cargo test --features=colorful echelon_matrix_no_variable_unsatisfiable -- --nocapture let mut matrix: Echelon>> = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -544,12 +614,23 @@ pub mod tests { fn echelon_matrix_no_more_variable_satisfiable() { // cargo test --features=colorful echelon_matrix_no_more_variable_satisfiable -- --nocapture let mut matrix: Echelon>> = EchelonMatrix::new(); - matrix.add_constraint(0, &[0, 1], true); - matrix.add_constraint(1, &[1, 2], true); - matrix.add_constraint(2, &[2, 3], true); - matrix.add_constraint(3, &[3, 1], false); + let vertex_indices = vec![0, 1, 2, 3]; + let edge_indices = vec![0, 1, 2, 3]; + let vertex_incident_edges_vec = vec![ + vec![0, 1], + vec![1, 2], + vec![2, 3], + vec![3, 1], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), true); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + matrix.add_constraint(vertices[3].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[3], &edges), false); + for edge_index in [0, 1, 2, 3] { - matrix.update_edge_tightness(edge_index, true); + matrix.update_edge_tightness(edges[edge_index].downgrade(), true); } matrix.printstd(); assert_eq!( @@ -572,14 +653,26 @@ pub mod tests { #[test] fn echelon_matrix_no_more_variable_unsatisfiable() { - // cargo test --features=colorful echelon_matrix_no_more_variable_satisfiable -- --nocapture + // cargo test --features=colorful echelon_matrix_no_more_variable_unsatisfiable -- --nocapture let mut matrix: Echelon>> = EchelonMatrix::new(); - matrix.add_constraint(0, &[0, 1], true); - matrix.add_constraint(1, &[1, 2], true); - matrix.add_constraint(2, &[2, 3], true); - matrix.add_constraint(3, &[3, 1], true); + let vertex_indices = vec![0, 1, 2, 3]; + let edge_indices = vec![0, 1, 2, 3]; + let vertex_incident_edges_vec = vec![ + vec![0, 1], + vec![1, 2], + vec![2, 3], + vec![3, 1], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), true); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + matrix.add_constraint(vertices[3].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[3], &edges), true); + + for edge_index in [0, 1, 2, 3] { - matrix.update_edge_tightness(edge_index, true); + matrix.update_edge_tightness(edges[edge_index].downgrade(), true); } matrix.printstd(); assert_eq!( @@ -775,15 +868,28 @@ pub mod tests { fn echelon_matrix_another_echelon_simple() { // cargo test --features=colorful echelon_matrix_another_echelon_simple -- --nocapture let mut echelon = EchelonMatrix::new(); + let vertex_indices = vec![0, 1, 2, 3, 4, 5]; + let edge_indices = vec![0, 1, 2, 3, 4, 5, 6]; + let vertex_incident_edges_vec = vec![ + vec![0, 1], + vec![0, 2], + vec![2, 3, 5], + vec![1, 3, 4], + vec![4, 6], + vec![5, 6], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + for edge_index in 0..7 { - echelon.add_tight_variable(edge_index); + echelon.add_tight_variable(edges[edge_index].downgrade()); } - echelon.add_constraint(0, &[0, 1], true); - echelon.add_constraint(1, &[0, 2], false); - echelon.add_constraint(2, &[2, 3, 5], false); - echelon.add_constraint(3, &[1, 3, 4], false); - echelon.add_constraint(4, &[4, 6], false); - echelon.add_constraint(5, &[5, 6], true); + + echelon.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + echelon.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + echelon.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), false); + echelon.add_constraint(vertices[3].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[3], &edges), false); + echelon.add_constraint(vertices[4].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[4], &edges), false); + echelon.add_constraint(vertices[5].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[5], &edges), true); let mut another = YetAnotherRowEchelon::new(&echelon); another.print(); // both go to echelon form @@ -803,13 +909,18 @@ pub mod tests { for constraint_count in 0..31 { for _ in 0..repeat { let mut echelon = EchelonMatrix::new(); + let parity_checks = generate_random_parity_checks(&mut rng, variable_count, constraint_count); + let vertex_indices: Vec = (0..parity_checks.len()).collect(); + let edge_indices: Vec = (0..variable_count).collect(); + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); for edge_index in 0..variable_count { - echelon.add_tight_variable(edge_index); + echelon.add_tight_variable(edges[edge_index].downgrade()); + } + if variable_count == 9 { + println!("variable_count: {variable_count}, parity_checks: {parity_checks:?}"); } - let parity_checks = generate_random_parity_checks(&mut rng, variable_count, constraint_count); - // println!("variable_count: {variable_count}, parity_checks: {parity_checks:?}"); for (vertex_index, (incident_edges, parity)) in parity_checks.iter().enumerate() { - echelon.add_constraint(vertex_index, incident_edges, *parity); + echelon.add_constraint(vertices[vertex_index].downgrade(), &edge_vec_from_indices(incident_edges, &edges), *parity); } let mut another = YetAnotherRowEchelon::new(&echelon); // echelon.printstd(); @@ -820,18 +931,20 @@ pub mod tests { // another.print(); another.assert_eq(&echelon); } + drop(vertices); + drop(edges); } } } } - fn debug_echelon_matrix_case(variable_count: usize, parity_checks: Vec<(Vec, bool)>) -> EchelonMatrix { + fn debug_echelon_matrix_case(variable_count: usize, parity_checks: Vec<(Vec, bool)>, edges: &Vec, vertices: &Vec) -> EchelonMatrix { let mut echelon = EchelonMatrix::new(); for edge_index in 0..variable_count { - echelon.add_tight_variable(edge_index); + echelon.add_tight_variable(edges[edge_index].downgrade()); } for (vertex_index, (incident_edges, parity)) in parity_checks.iter().enumerate() { - echelon.add_constraint(vertex_index, incident_edges, *parity); + echelon.add_constraint(vertices[vertex_index].downgrade(), &edge_vec_from_indices(incident_edges, edges), *parity); } echelon } @@ -841,7 +954,11 @@ pub mod tests { fn echelon_matrix_debug_1() { // cargo test --features=colorful echelon_matrix_debug_1 -- --nocapture let parity_checks = vec![(vec![0], true), (vec![0, 1], true), (vec![], true)]; - let mut echelon = debug_echelon_matrix_case(2, parity_checks); + let vertex_indices: Vec = (0..parity_checks.len()).collect(); + let edge_indices: Vec = (0..2).collect(); + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut echelon = debug_echelon_matrix_case(2, parity_checks, &edges, &vertices); echelon.printstd(); assert_eq!( echelon.printstd_str(), @@ -865,7 +982,11 @@ pub mod tests { fn echelon_matrix_debug_2() { // cargo test --features=colorful echelon_matrix_debug_2 -- --nocapture let parity_checks = vec![]; - let mut echelon = debug_echelon_matrix_case(1, parity_checks); + let vertex_indices: Vec = (0..parity_checks.len()).collect(); + let edge_indices: Vec = (0..1).collect(); + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut echelon = debug_echelon_matrix_case(1, parity_checks, &edges, &vertices); echelon.printstd(); assert_eq!( echelon.printstd_str(), diff --git a/src/matrix/hair.rs b/src/matrix/hair.rs index 7ad78339..aff59c97 100644 --- a/src/matrix/hair.rs +++ b/src/matrix/hair.rs @@ -7,6 +7,7 @@ use super::interface::*; use super::visualize::*; use crate::util::*; use prettytable::*; +use crate::dual_module_pq::{EdgeWeak, VertexWeak}; pub struct HairView<'a, M: MatrixTail + MatrixEchelon> { base: &'a mut M, @@ -18,7 +19,7 @@ impl<'a, M: MatrixTail + MatrixEchelon> HairView<'a, M> { pub fn get_base(&self) -> &M { self.base } - pub fn get_base_view_edges(&mut self) -> Vec { + pub fn get_base_view_edges(&mut self) -> Vec { self.base.get_view_edges() } } @@ -26,7 +27,7 @@ impl<'a, M: MatrixTail + MatrixEchelon> HairView<'a, M> { impl<'a, M: MatrixTail + MatrixEchelon> HairView<'a, M> { pub fn new(matrix: &'a mut M, hair: EdgeIter) -> Self where - EdgeIter: Iterator, + EdgeIter: Iterator, { matrix.set_tail_edges(hair); let columns = matrix.columns(); @@ -34,8 +35,8 @@ impl<'a, M: MatrixTail + MatrixEchelon> HairView<'a, M> { let mut column_bias = columns; let mut row_bias = rows; for column in (0..columns).rev() { - let edge_index = matrix.column_to_edge_index(column); - if matrix.get_tail_edges().contains(&edge_index) { + let edge_ptr = matrix.column_to_edge_index(column); + if matrix.get_tail_edges().contains(&edge_ptr) { column_bias = column; } else { break; @@ -70,10 +71,10 @@ impl<'a, M: MatrixTail + MatrixEchelon> HairView<'a, M> { } impl<'a, M: MatrixTail + MatrixEchelon> MatrixTail for HairView<'a, M> { - fn get_tail_edges(&self) -> &FastIterSet { + fn get_tail_edges(&self) -> &FastIterSet { self.get_base().get_tail_edges() } - fn get_tail_edges_mut(&mut self) -> &mut FastIterSet { + fn get_tail_edges_mut(&mut self) -> &mut FastIterSet { panic!("cannot mutate a hair view"); } } @@ -88,26 +89,26 @@ impl<'a, M: MatrixTail + MatrixEchelon> MatrixEchelon for HairView<'a, M> { } impl<'a, M: MatrixTight + MatrixTail + MatrixEchelon> MatrixTight for HairView<'a, M> { - fn update_edge_tightness(&mut self, _edge_index: EdgeIndex, _is_tight: bool) { + fn update_edge_tightness(&mut self, _edge_weak: EdgeWeak, _is_tight: bool) { panic!("cannot mutate a hair view"); } - fn is_tight(&self, edge_index: usize) -> bool { - self.get_base().is_tight(edge_index) + fn is_tight(&self, edge_weak: EdgeWeak) -> bool { + self.get_base().is_tight(edge_weak.clone()) } - fn get_tight_edges(&self) -> &FastIterSet { + fn get_tight_edges(&self) -> &FastIterSet { self.base.get_tight_edges() } } impl<'a, M: MatrixTail + MatrixEchelon> MatrixBasic for HairView<'a, M> { - fn add_variable(&mut self, _edge_index: EdgeIndex) -> Option { + fn add_variable(&mut self, _edge_weak: EdgeWeak) -> Option { panic!("cannot mutate a hair view"); } fn add_constraint( &mut self, - _vertex_index: VertexIndex, - _incident_edges: &[EdgeIndex], + _vertex_weak: VertexWeak, + _incident_edges: &[EdgeWeak], _parity: bool, ) -> Option> { panic!("cannot mutate a hair view"); @@ -125,16 +126,16 @@ impl<'a, M: MatrixTail + MatrixEchelon> MatrixBasic for HairView<'a, M> { fn get_rhs(&self, row: RowIndex) -> bool { self.get_base().get_rhs(row + self.row_bias) } - fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeIndex { + fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeWeak { self.get_base().var_to_edge_index(var_index) } - fn edge_to_var_index(&self, edge_index: EdgeIndex) -> Option { - self.get_base().edge_to_var_index(edge_index) + fn edge_to_var_index(&self, edge_weak: EdgeWeak) -> Option { + self.get_base().edge_to_var_index(edge_weak.clone()) } - fn get_vertices(&self) -> FastIterSet { + fn get_vertices(&self) -> FastIterSet { self.get_base().get_vertices() } - fn get_edges(&self) -> FastIterSet { + fn get_edges(&self) -> FastIterSet { self.get_base().get_edges() } } @@ -170,7 +171,7 @@ impl<'a, M: MatrixTail + MatrixEchelon> VizTrait for HairView<'a, M> { let row_info = self.get_echelon_row_info(row); let cell = if row_info.has_leading() { Cell::new( - self.column_to_edge_index(row_info.column - self.column_bias) + self.column_to_edge_index(row_info.column - self.column_bias).upgrade_force().read_recursive().edge_index .to_string() .as_str(), ) @@ -201,6 +202,7 @@ impl<'a, M: MatrixTail + MatrixEchelon> VizTrait for HairView<'a, M> { } } + #[cfg(test)] pub mod tests { use super::super::basic::*; @@ -208,6 +210,9 @@ pub mod tests { use super::super::tail::*; use super::super::tight::*; use super::*; + use crate::matrix::basic::tests::{initialize_vertex_edges_for_matrix_testing, edge_vec_from_indices}; + use std::collections::HashSet; + use crate::dual_module_pq::{EdgePtr, VertexPtr}; type EchelonMatrix = Echelon>>; @@ -215,12 +220,21 @@ pub mod tests { fn hair_view_simple() { // cargo test --features=colorful hair_view_simple -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); - assert_eq!(matrix.edge_to_var_index(4), Some(1)); - for edge_index in [1, 4, 6, 9] { - matrix.update_edge_tightness(edge_index, true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + assert_eq!(matrix.edge_to_var_index(edges[1].downgrade()), Some(1)); + for edge_ptr in edges.iter() { + matrix.update_edge_tightness(edge_ptr.downgrade(), true); } matrix.printstd(); assert_eq!( @@ -239,8 +253,8 @@ pub mod tests { └──┴─┴─┴─┴─┴───┴─┘ " ); - let mut hair_view = HairView::new(&mut matrix, [6, 9].into_iter()); - assert_eq!(hair_view.edge_to_var_index(4), Some(1)); + let mut hair_view = HairView::new(&mut matrix, [edges[2].downgrade(), edges[3].downgrade()].into_iter()); + assert_eq!(hair_view.edge_to_var_index(edges[1].downgrade()), Some(1)); hair_view.printstd(); assert_eq!( hair_view.printstd_str(), @@ -254,7 +268,7 @@ pub mod tests { └──┴─┴─┴───┴─┘ " ); - let mut hair_view = HairView::new(&mut matrix, [1, 6].into_iter()); + let mut hair_view = HairView::new(&mut matrix, [edges[0].downgrade(), edges[2].downgrade()].into_iter()); hair_view.base.printstd(); assert_eq!( hair_view.base.printstd_str(), @@ -285,19 +299,25 @@ pub mod tests { └──┴─┴─┴───┴─┘ " ); - assert_eq!(hair_view.get_tail_edges_vec(), [1, 6]); - assert!(hair_view.is_tight(1)); + assert_eq!( + hair_view.get_tail_edges_vec().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + [1, 6].into_iter().collect::>()); + assert!(hair_view.is_tight(edges[0].downgrade())); assert!(hair_view.get_echelon_satisfiable()); - assert_eq!(hair_view.get_vertices(), [0, 1, 2].into()); - assert_eq!(hair_view.get_base_view_edges(), [4, 9, 1, 6]); + assert_eq!( + hair_view.get_vertices().iter().map(|v| v.upgrade_force().read_recursive().vertex_index).collect::>(), + [0, 1, 2].into_iter().collect::>()); + assert_eq!( + hair_view.get_base_view_edges().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + [4, 9, 1, 6].into_iter().collect::>()); } - fn generate_demo_matrix() -> EchelonMatrix { + fn generate_demo_matrix(edges: &Vec, vertices: &Vec) -> EchelonMatrix { let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - for edge_index in [1, 4, 6, 9] { - matrix.update_edge_tightness(edge_index, true); + matrix.add_constraint(vertices[0].downgrade(), &[edges[0].downgrade(), edges[1].downgrade(), edges[2].downgrade()], true); + matrix.add_constraint(vertices[1].downgrade(), &[edges[1].downgrade(), edges[3].downgrade()], false); + for edge_index in 0..4 { + matrix.update_edge_tightness(edges[edge_index].downgrade(), true); } matrix } @@ -306,7 +326,11 @@ pub mod tests { #[should_panic] fn hair_view_should_not_modify_tail_edges() { // cargo test hair_view_should_not_modify_tail_edges -- --nocapture - let mut matrix = generate_demo_matrix(); + let vertex_indices = vec![0, 1]; + let edge_indices = vec![1, 4, 6, 9]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut matrix = generate_demo_matrix(&edges, &vertices); let mut hair_view = HairView::new(&mut matrix, [].into_iter()); hair_view.get_tail_edges_mut(); } @@ -315,34 +339,50 @@ pub mod tests { #[should_panic] fn hair_view_should_not_update_edge_tightness() { // cargo test hair_view_should_not_update_edge_tightness -- --nocapture - let mut matrix = generate_demo_matrix(); + let vertex_indices = vec![0, 1]; + let edge_indices = vec![1, 4, 6, 9]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut matrix = generate_demo_matrix(&edges, &vertices); let mut hair_view = HairView::new(&mut matrix, [].into_iter()); - hair_view.update_edge_tightness(1, false); + hair_view.update_edge_tightness(edges[0].downgrade(), false); } #[test] #[should_panic] fn hair_view_should_not_add_variable() { // cargo test hair_view_should_not_add_variable -- --nocapture - let mut matrix = generate_demo_matrix(); + let vertex_indices = vec![0, 1]; + let edge_indices = vec![1, 4, 6, 9, 100]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut matrix = generate_demo_matrix(&edges, &vertices); let mut hair_view = HairView::new(&mut matrix, [].into_iter()); - hair_view.add_variable(100); + hair_view.add_variable(edges[4].downgrade()); } #[test] #[should_panic] fn hair_view_should_not_add_constraint() { // cargo test hair_view_should_not_add_constraint -- --nocapture - let mut matrix = generate_demo_matrix(); + let vertex_indices = vec![0, 1, 5]; + let edge_indices = vec![1, 4, 6, 9, 2, 3]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut matrix = generate_demo_matrix(&edges, &vertices); let mut hair_view = HairView::new(&mut matrix, [].into_iter()); - hair_view.add_constraint(5, &[1, 2, 3], false); + hair_view.add_constraint(vertices[2].downgrade(), &[edges[0].downgrade(), edges[4].downgrade(), edges[5].downgrade()], false); } #[test] #[should_panic] fn hair_view_should_not_xor_row() { // cargo test hair_view_should_not_xor_row -- --nocapture - let mut matrix = generate_demo_matrix(); + let vertex_indices = vec![0, 1, 5]; + let edge_indices = vec![1, 4, 6, 9, 2, 3]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut matrix = generate_demo_matrix(&edges, &vertices); let mut hair_view = HairView::new(&mut matrix, [].into_iter()); hair_view.xor_row(0, 1); } @@ -351,7 +391,11 @@ pub mod tests { #[should_panic] fn hair_view_should_not_swap_row() { // cargo test hair_view_should_not_swap_row -- --nocapture - let mut matrix = generate_demo_matrix(); + let vertex_indices = vec![0, 1, 5]; + let edge_indices = vec![1, 4, 6, 9, 2, 3]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut matrix = generate_demo_matrix(&edges, &vertices); let mut hair_view = HairView::new(&mut matrix, [].into_iter()); hair_view.swap_row(0, 1); } @@ -360,7 +404,11 @@ pub mod tests { #[should_panic] fn hair_view_should_not_get_echelon_info() { // cargo test hair_view_should_not_get_echelon_info -- --nocapture - let mut matrix = generate_demo_matrix(); + let vertex_indices = vec![0, 1, 5]; + let edge_indices = vec![1, 4, 6, 9, 2, 3]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + let mut matrix = generate_demo_matrix(&edges, &vertices); let mut hair_view = HairView::new(&mut matrix, [].into_iter()); hair_view.get_echelon_info(); } @@ -369,7 +417,11 @@ pub mod tests { #[should_panic] fn hair_view_should_not_get_echelon_info_immutable() { // cargo test hair_view_should_not_get_echelon_info_immutable -- --nocapture - let mut matrix = generate_demo_matrix(); + let vertex_indices = vec![0, 1, 5]; + let edge_indices = vec![1, 4, 6, 9, 2, 3]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + let mut matrix = generate_demo_matrix(&edges, &vertices); + let hair_view = HairView::new(&mut matrix, [].into_iter()); hair_view.get_echelon_info_immutable(); } @@ -378,12 +430,22 @@ pub mod tests { fn hair_view_unsatisfiable() { // cargo test --features=colorful hair_view_unsatisfiable -- --nocapture let mut matrix = EchelonMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); - matrix.add_constraint(3, &[1, 9], false); - for edge_index in [1, 4, 6, 9] { - matrix.update_edge_tightness(edge_index, true); + let vertex_indices = vec![0, 1, 2, 3]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + matrix.add_constraint(vertices[3].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[3], &edges), false); + + for edge_ptr in edges.iter() { + matrix.update_edge_tightness(edge_ptr.downgrade(), true); } matrix.printstd(); assert_eq!( @@ -404,7 +466,7 @@ pub mod tests { └──┴─┴─┴─┴─┴───┴─┘ " ); - let mut hair_view = HairView::new(&mut matrix, [6, 9].into_iter()); + let mut hair_view = HairView::new(&mut matrix, [edges[2].downgrade(), edges[3].downgrade()].into_iter()); hair_view.printstd(); assert_eq!( hair_view.printstd_str(), diff --git a/src/matrix/interface.rs b/src/matrix/interface.rs index 296be01c..f160692f 100644 --- a/src/matrix/interface.rs +++ b/src/matrix/interface.rs @@ -22,6 +22,7 @@ use crate::util::*; use derivative::Derivative; use num_traits::{One, Zero}; +use crate::dual_module_pq::{EdgeWeak, VertexWeak}; pub type VarIndex = usize; pub type RowIndex = usize; @@ -29,13 +30,13 @@ pub type ColumnIndex = usize; pub trait MatrixBasic { /// add an edge to the basic matrix, return the `var_index` if newly created - fn add_variable(&mut self, edge_index: EdgeIndex) -> Option; + fn add_variable(&mut self, edge_weak: EdgeWeak) -> Option; /// add constraint will implicitly call `add_variable` if the edge is not added and return the indices of them fn add_constraint( &mut self, - vertex_index: VertexIndex, - incident_edges: &[EdgeIndex], + vertex_weak: VertexWeak, + incident_edges: &[EdgeWeak], parity: bool, ) -> Option>; @@ -48,16 +49,17 @@ pub trait MatrixBasic { fn get_rhs(&self, row: RowIndex) -> bool; /// get edge index from the var_index - fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeIndex; + fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeWeak; - fn edge_to_var_index(&self, edge_index: EdgeIndex) -> Option; + fn edge_to_var_index(&self, edge_weak: EdgeWeak) -> Option; - fn exists_edge(&self, edge_index: EdgeIndex) -> bool { - self.edge_to_var_index(edge_index).is_some() + fn exists_edge(&self, edge_weak: EdgeWeak) -> bool { + self.edge_to_var_index(edge_weak.clone()).is_some() } - fn get_edges(&self) -> FastIterSet; - fn get_vertices(&self) -> FastIterSet; + // TODO: look into these to see if we can pass reference instead of cloning + fn get_edges(&self) -> FastIterSet; + fn get_vertices(&self) -> FastIterSet; } pub trait MatrixView: MatrixBasic { @@ -69,7 +71,7 @@ pub trait MatrixView: MatrixBasic { /// get the `var_index` in the basic matrix fn column_to_var_index(&self, column: ColumnIndex) -> VarIndex; - fn column_to_edge_index(&self, column: ColumnIndex) -> EdgeIndex { + fn column_to_edge_index(&self, column: ColumnIndex) -> EdgeWeak { let var_index = self.column_to_var_index(column); self.var_to_edge_index(var_index) } @@ -77,7 +79,7 @@ pub trait MatrixView: MatrixBasic { /// the number of rows: rows always have indices 0..rows fn rows(&mut self) -> usize; - fn get_view_edges(&mut self) -> Vec { + fn get_view_edges(&mut self) -> Vec { (0..self.columns()) .map(|column: usize| self.column_to_edge_index(column)) .collect() @@ -87,44 +89,44 @@ pub trait MatrixView: MatrixBasic { (0..self.columns()).find(|&column| self.column_to_var_index(column) == var_index) } - fn edge_to_column_index(&mut self, edge_index: EdgeIndex) -> Option { - let var_index = self.edge_to_var_index(edge_index)?; + fn edge_to_column_index(&mut self, edge_weak: EdgeWeak) -> Option { + let var_index = self.edge_to_var_index(edge_weak)?; self.var_to_column_index(var_index) } } pub trait MatrixTight: MatrixView { - fn update_edge_tightness(&mut self, edge_index: EdgeIndex, is_tight: bool); - fn is_tight(&self, edge_index: usize) -> bool; - fn get_tight_edges(&self) -> &FastIterSet; + fn update_edge_tightness(&mut self, edge_weak: EdgeWeak, is_tight: bool); + fn is_tight(&self, edge_weak: EdgeWeak) -> bool; + fn get_tight_edges(&self) -> &FastIterSet; - fn add_variable_with_tightness(&mut self, edge_index: EdgeIndex, is_tight: bool) { - self.add_variable(edge_index); - self.update_edge_tightness(edge_index, is_tight); + fn add_variable_with_tightness(&mut self, edge_weak: EdgeWeak, is_tight: bool) { + self.add_variable(edge_weak.clone()); + self.update_edge_tightness(edge_weak.clone(), is_tight); } - fn add_tight_variable(&mut self, edge_index: EdgeIndex) { - self.add_variable_with_tightness(edge_index, true) + fn add_tight_variable(&mut self, edge_weak: EdgeWeak) { + self.add_variable_with_tightness(edge_weak, true) } } pub trait MatrixTail { - fn get_tail_edges(&self) -> &FastIterSet; - fn get_tail_edges_mut(&mut self) -> &mut FastIterSet; + fn get_tail_edges(&self) -> &FastIterSet; + fn get_tail_edges_mut(&mut self) -> &mut FastIterSet; fn set_tail_edges(&mut self, edges: EdgeIter) where - EdgeIter: Iterator, + EdgeIter: Iterator, { let tail_edges = self.get_tail_edges_mut(); tail_edges.clear(); - for edge_index in edges { - tail_edges.insert(edge_index); + for edge_weak in edges { + tail_edges.insert(edge_weak.clone()); } } - fn get_tail_edges_vec(&self) -> Vec { - let mut edges: Vec = self.get_tail_edges().iter().cloned().collect(); + fn get_tail_edges_vec(&self) -> Vec { + let mut edges: Vec = self.get_tail_edges().iter().cloned().collect(); edges.sort(); edges } @@ -139,7 +141,7 @@ pub trait MatrixEchelon: MatrixView { fn get_echelon_info(&mut self) -> &EchelonInfo; fn get_echelon_info_immutable(&self) -> &EchelonInfo; - fn get_solution(&mut self) -> Option { + fn get_solution(&mut self) -> Option { self.get_echelon_info(); // make sure it's in echelon form let info = self.get_echelon_info_immutable(); if !info.satisfiable { @@ -150,17 +152,17 @@ pub trait MatrixEchelon: MatrixView { debug_assert!(row_info.has_leading()); if self.get_rhs(row) { let column = row_info.column; - let edge_index = self.column_to_edge_index(column); - solution.push(edge_index); + let edge_weak = self.column_to_edge_index(column); + solution.push(edge_weak.clone()); } } Some(solution) } /// try every independent variables and try to minimize the total weight of the solution - fn get_solution_local_minimum(&mut self, weight_of: F) -> Option + fn get_solution_local_minimum(&mut self, weight_of: F) -> Option where - F: Fn(EdgeIndex) -> Weight, + F: Fn(EdgeWeak) -> Weight, { self.get_echelon_info(); // make sure it's in echelon form let info = self.get_echelon_info_immutable(); @@ -172,8 +174,8 @@ pub trait MatrixEchelon: MatrixView { debug_assert!(row_info.has_leading()); if self.get_rhs(row) { let column = row_info.column; - let edge_index = self.column_to_edge_index(column); - solution.insert(edge_index); + let edge_weak = self.column_to_edge_index(column); + solution.insert(edge_weak.clone()); } } let mut independent_columns = vec![]; @@ -183,8 +185,8 @@ pub trait MatrixEchelon: MatrixView { } } let mut total_weight = Rational::zero(); - for &edge_index in solution.iter() { - total_weight += weight_of(edge_index); + for edge_weak in solution.iter() { + total_weight += weight_of(edge_weak.clone()); } let mut pending_flip_edge_indices = vec![]; let mut is_local_minimum = false; @@ -194,37 +196,37 @@ pub trait MatrixEchelon: MatrixView { for &column in independent_columns.iter() { pending_flip_edge_indices.clear(); let var_index = self.column_to_var_index(column); - let edge_index = self.var_to_edge_index(var_index); - let mut primal_delta = (weight_of(edge_index)) - * if solution.contains(&edge_index) { + let edge_ptr = self.var_to_edge_index(var_index); + let mut primal_delta = (weight_of(edge_ptr.clone())) + * if solution.contains(&edge_ptr) { -Rational::one() } else { Rational::one() }; - pending_flip_edge_indices.push(edge_index); + pending_flip_edge_indices.push(edge_ptr); for row in 0..info.rows.len() { if self.get_lhs(row, var_index) { debug_assert!(info.rows[row].has_leading()); let flip_column = info.rows[row].column; debug_assert!(flip_column < column); - let flip_edge_index = self.column_to_edge_index(flip_column); - primal_delta += (weight_of(flip_edge_index)) - * if solution.contains(&flip_edge_index) { + let flip_edge_ptr = self.column_to_edge_index(flip_column); + primal_delta += (weight_of(flip_edge_ptr.clone())) + * if solution.contains(&flip_edge_ptr) { -Rational::one() } else { Rational::one() }; - pending_flip_edge_indices.push(flip_edge_index); + pending_flip_edge_indices.push(flip_edge_ptr); } } // warning: has to be this form (instead of .is_negative) to use the tolerance of OrderedFloat if primal_delta < Rational::zero() { total_weight = total_weight + primal_delta; - for &edge_index in pending_flip_edge_indices.iter() { - if solution.contains(&edge_index) { - solution.remove(&edge_index); + for edge_weak in pending_flip_edge_indices.iter() { + if solution.contains(edge_weak) { + solution.remove(edge_weak); } else { - solution.insert(edge_index); + solution.insert(edge_weak.clone()); } } is_local_minimum = false; @@ -355,10 +357,15 @@ impl std::fmt::Debug for RowInfo { } } + #[cfg(test)] pub mod tests { use super::super::*; use super::*; + use std::collections::BTreeMap; + use crate::matrix::basic::tests::{initialize_vertex_edges_for_matrix_testing, edge_vec_from_indices}; + use std::collections::HashSet; + use crate::dual_module_pq::{EdgePtr, VertexPtr}; type TightMatrix = Tight; @@ -366,21 +373,27 @@ pub mod tests { fn matrix_interface_simple() { // cargo test --features=colorful matrix_interface_simple -- --nocapture let mut matrix = TightMatrix::new(); - matrix.add_tight_variable(233); - matrix.add_tight_variable(14); - matrix.add_variable(68); - matrix.add_tight_variable(75); + let vertex_indices = vec![0, 1, 2, 3]; + let edge_indices = vec![233, 14, 68, 75, 666]; + let (_vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_tight_variable(edges[0].downgrade()); + matrix.add_tight_variable(edges[1].downgrade()); + matrix.add_variable(edges[2].downgrade()); + matrix.add_tight_variable(edges[3].downgrade()); matrix.printstd(); - assert_eq!(matrix.get_view_edges(), [233, 14, 75]); + assert_eq!( + matrix.get_view_edges().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + [233, 14, 75].into_iter().collect::>()); assert_eq!(matrix.var_to_column_index(0), Some(0)); assert_eq!(matrix.var_to_column_index(1), Some(1)); assert_eq!(matrix.var_to_column_index(2), None); assert_eq!(matrix.var_to_column_index(3), Some(2)); - assert_eq!(matrix.edge_to_column_index(233), Some(0)); - assert_eq!(matrix.edge_to_column_index(14), Some(1)); - assert_eq!(matrix.edge_to_column_index(68), None); - assert_eq!(matrix.edge_to_column_index(75), Some(2)); - assert_eq!(matrix.edge_to_column_index(666), None); + assert_eq!(matrix.edge_to_column_index(edges[0].downgrade()), Some(0)); + assert_eq!(matrix.edge_to_column_index(edges[1].downgrade()), Some(1)); + assert_eq!(matrix.edge_to_column_index(edges[2].downgrade()), None); + assert_eq!(matrix.edge_to_column_index(edges[3].downgrade()), Some(2)); + assert_eq!(matrix.edge_to_column_index(edges[4].downgrade()), None); } #[test] @@ -404,23 +417,23 @@ pub mod tests { #[derive(Default)] struct TestEdgeWeights { - pub weights: FastIterMap, + pub weights: BTreeMap, } impl TestEdgeWeights { - fn new(weights: &[(EdgeIndex, Weight)]) -> Self { + fn new(weights: &[(EdgeWeak, Weight)]) -> Self { let mut result: TestEdgeWeights = Default::default(); - for (edge_index, weight) in weights { - result.weights.insert(edge_index.clone(), weight.clone()); + for (edge_weak, weight) in weights { + result.weights.insert(edge_weak.clone(), weight.clone()); } result } - fn get_solution_local_minimum(&self, matrix: &mut Echelon>) -> Option { - matrix.get_solution_local_minimum(|edge_index| { - if let Some(weight) = self.weights.get(&edge_index) { + fn get_solution_local_minimum(&self, matrix: &mut Echelon>) -> Option> { + matrix.get_solution_local_minimum(|edge_weak| { + if let Some(weight) = self.weights.get(&edge_weak) { weight.clone() } else { - Rational::from_float(1.).unwrap() + Rational::from(1.) } }) } @@ -447,45 +460,44 @@ pub mod tests { (vec![6, 9], false), (vec![0, 8, 9], true), ]; + let vertex_indices: Vec = (0..parity_checks.len()).collect(); + let edge_indices: Vec = (0..10).collect(); + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + for (vertex_index, (incident_edges, parity)) in parity_checks.iter().enumerate() { - matrix.add_constraint(vertex_index, incident_edges, *parity); + matrix.add_constraint(vertices[vertex_index].downgrade(), &edge_vec_from_indices(incident_edges, &edges),*parity); + // matrix.printstd(); } matrix.printstd(); - assert_eq!(matrix.get_solution(), Some(vec![0, 1, 2, 3, 4])); - let weights = TestEdgeWeights::new(&[ - (3, Rational::from_float(10.).unwrap()), - (9, Rational::from_float(10.).unwrap()), - ]); assert_eq!( - sorted_vec_option(weights.get_solution_local_minimum(&mut matrix)), - Some(vec![5, 7, 8]) - ); - let weights = TestEdgeWeights::new(&[ - (7, Rational::from_float(10.).unwrap()), - (9, Rational::from_float(10.).unwrap()), - ]); + matrix.get_solution().unwrap().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + vec![0, 1, 2, 3, 4].into_iter().collect::>()); + let weights = TestEdgeWeights::new(&[(edges[3].downgrade(), Rational::from(10.)), (edges[9].downgrade(), Rational::from(10.))]); assert_eq!( - sorted_vec_option(weights.get_solution_local_minimum(&mut matrix)), - Some(vec![3, 4, 8]) - ); - let weights = TestEdgeWeights::new(&[ - (3, Rational::from_float(10.).unwrap()), - (4, Rational::from_float(10.).unwrap()), - (7, Rational::from_float(10.).unwrap()), - ]); + weights.get_solution_local_minimum(&mut matrix).unwrap().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + vec![5, 7, 8].into_iter().collect::>()); + let weights = TestEdgeWeights::new(&[(edges[7].downgrade(), Rational::from(10.)), (edges[9].downgrade(), Rational::from(10.))]); assert_eq!( - sorted_vec_option(weights.get_solution_local_minimum(&mut matrix)), - Some(vec![5, 6, 9]) - ); + weights.get_solution_local_minimum(&mut matrix).unwrap().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + vec![3, 4, 8].into_iter().collect::>()); + let weights = TestEdgeWeights::new(&[(edges[3].downgrade(), Rational::from(10.)), (edges[4].downgrade(), Rational::from(10.)), (edges[7].downgrade(), Rational::from(10.))]); + assert_eq!( + weights.get_solution_local_minimum(&mut matrix).unwrap().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + vec![5, 6, 9].into_iter().collect::>()); } #[test] fn matrix_interface_echelon_no_solution() { // cargo test matrix_interface_echelon_no_solution -- --nocapture let mut matrix = Echelon::>::new(); - let parity_checks = [(vec![0, 1], false), (vec![0, 1], true)]; + let parity_checks = vec![(vec![0, 1], false), (vec![0, 1], true)]; + let vertex_indices: Vec = (0..parity_checks.len()).collect(); + let edge_indices: Vec = (0..10).collect(); + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + for (vertex_index, (incident_edges, parity)) in parity_checks.iter().enumerate() { - matrix.add_constraint(vertex_index, incident_edges, *parity); + let incident_edges_weak: Vec = incident_edges.iter().map(|&i| edges[i].downgrade()).collect(); + matrix.add_constraint(vertices[vertex_index].downgrade(), &incident_edges_weak, *parity); } assert_eq!(matrix.get_solution(), None); let weights = TestEdgeWeights::new(&[]); diff --git a/src/matrix/tail.rs b/src/matrix/tail.rs index d9ec1152..83aed918 100644 --- a/src/matrix/tail.rs +++ b/src/matrix/tail.rs @@ -2,13 +2,14 @@ use super::interface::*; use super::visualize::*; use crate::util::*; use derivative::Derivative; +use crate::dual_module_pq::{EdgeWeak, VertexWeak}; #[derive(Clone, Derivative)] #[derivative(Default(new = "true"))] pub struct Tail { base: M, /// the set of edges that should be placed at the end, if any - tail_edges: FastIterSet, + tail_edges: FastIterSet, /// var indices are outdated on any changes to the underlying matrix #[derivative(Default(value = "true"))] is_var_indices_outdated: bool, @@ -36,41 +37,41 @@ impl Tail { } impl MatrixTail for Tail { - fn get_tail_edges(&self) -> &FastIterSet { + fn get_tail_edges(&self) -> &FastIterSet { &self.tail_edges } - fn get_tail_edges_mut(&mut self) -> &mut FastIterSet { + fn get_tail_edges_mut(&mut self) -> &mut FastIterSet { self.is_var_indices_outdated = true; &mut self.tail_edges } } impl MatrixTight for Tail { - fn update_edge_tightness(&mut self, edge_index: EdgeIndex, is_tight: bool) { + fn update_edge_tightness(&mut self, edge_weak: EdgeWeak, is_tight: bool) { self.is_var_indices_outdated = true; - self.base.update_edge_tightness(edge_index, is_tight) + self.base.update_edge_tightness(edge_weak, is_tight) } - fn is_tight(&self, edge_index: usize) -> bool { - self.base.is_tight(edge_index) + fn is_tight(&self, edge_weak: EdgeWeak) -> bool { + self.base.is_tight(edge_weak) } - fn get_tight_edges(&self) -> &FastIterSet { + fn get_tight_edges(&self) -> &FastIterSet { self.base.get_tight_edges() } } impl MatrixBasic for Tail { - fn add_variable(&mut self, edge_index: EdgeIndex) -> Option { + fn add_variable(&mut self, edge_weak: EdgeWeak) -> Option { self.is_var_indices_outdated = true; - self.base.add_variable(edge_index) + self.base.add_variable(edge_weak.clone()) } fn add_constraint( &mut self, - vertex_index: VertexIndex, - incident_edges: &[EdgeIndex], + vertex_weak: VertexWeak, + incident_edges: &[EdgeWeak], parity: bool, ) -> Option> { - self.base.add_constraint(vertex_index, incident_edges, parity) + self.base.add_constraint(vertex_weak.clone(), incident_edges, parity) } fn xor_row(&mut self, target: RowIndex, source: RowIndex) { @@ -85,16 +86,16 @@ impl MatrixBasic for Tail { fn get_rhs(&self, row: RowIndex) -> bool { self.get_base().get_rhs(row) } - fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeIndex { + fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeWeak { self.get_base().var_to_edge_index(var_index) } - fn edge_to_var_index(&self, edge_index: EdgeIndex) -> Option { - self.get_base().edge_to_var_index(edge_index) + fn edge_to_var_index(&self, edge_weak: EdgeWeak) -> Option { + self.get_base().edge_to_var_index(edge_weak) } - fn get_vertices(&self) -> FastIterSet { + fn get_vertices(&self) -> FastIterSet { self.get_base().get_vertices() } - fn get_edges(&self) -> FastIterSet { + fn get_edges(&self) -> FastIterSet { self.get_base().get_edges() } } @@ -150,6 +151,9 @@ pub mod tests { use super::super::basic::*; use super::super::tight::*; use super::*; + use crate::matrix::basic::tests::{initialize_vertex_edges_for_matrix_testing, edge_vec_from_indices}; + use std::collections::HashSet; + use crate::dual_module_pq::{EdgePtr, VertexPtr}; type TailMatrix = Tail>; @@ -157,10 +161,19 @@ pub mod tests { fn tail_matrix_1() { // cargo test --features=colorful tail_matrix_1 -- --nocapture let mut matrix = TailMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); - assert_eq!(matrix.edge_to_var_index(4), Some(1)); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + assert_eq!(matrix.edge_to_var_index(edges[1].downgrade()), Some(1)); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -176,8 +189,8 @@ pub mod tests { └─┴───┘ " ); - for edge_index in [1, 4, 6, 9] { - matrix.update_edge_tightness(edge_index, true); + for edge_ptr in edges.iter() { + matrix.update_edge_tightness(edge_ptr.downgrade(), true); } matrix.printstd(); assert_eq!( @@ -194,7 +207,7 @@ pub mod tests { └─┴─┴─┴─┴─┴───┘ " ); - matrix.set_tail_edges([1, 6].into_iter()); + matrix.set_tail_edges([edges[0].downgrade(), edges[2].downgrade()].into_iter()); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -210,8 +223,10 @@ pub mod tests { └─┴─┴─┴─┴─┴───┘ " ); - assert_eq!(matrix.get_tail_edges_vec(), [1, 6]); - assert_eq!(matrix.edge_to_var_index(4), Some(1)); + assert_eq!( + matrix.get_tail_edges_vec().iter().map(|e| e.upgrade_force().read_recursive().edge_index).collect::>(), + [1, 6].into_iter().collect::>()); + assert_eq!(matrix.edge_to_var_index(edges[1].downgrade()), Some(1)); } #[test] @@ -219,8 +234,15 @@ pub mod tests { fn tail_matrix_cannot_call_dirty_column() { // cargo test tail_matrix_cannot_call_dirty_column -- --nocapture let mut matrix = TailMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.update_edge_tightness(1, true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + + matrix.update_edge_tightness(edges[0].downgrade(), true); // even though there is indeed such a column, we forbid such dangerous calls // always call `columns()` before accessing any column matrix.column_to_var_index(0); @@ -230,14 +252,23 @@ pub mod tests { fn tail_matrix_basic_trait() { // cargo test --features=colorful tail_matrix_basic_trait -- --nocapture let mut matrix = TailMatrix::new(); - matrix.add_variable(3); // untight edges will not show - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9, 3]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_variable(edges[4].downgrade()); // untight edges will not show + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); matrix.swap_row(2, 1); matrix.xor_row(0, 1); - for edge_index in [1, 4, 6, 9] { - matrix.update_edge_tightness(edge_index, true); + for edge_index in 0..4 { + matrix.update_edge_tightness(edges[edge_index].downgrade(), true); } matrix.printstd(); assert_eq!( @@ -254,7 +285,7 @@ pub mod tests { └─┴─┴─┴─┴─┴───┘ " ); - assert!(matrix.is_tight(1)); - assert_eq!(matrix.edge_to_var_index(4), Some(2)); + assert!(matrix.is_tight(edges[0].downgrade())); + assert_eq!(matrix.edge_to_var_index(edges[1].downgrade()), Some(2)); } } diff --git a/src/matrix/tight.rs b/src/matrix/tight.rs index c583a337..4f582d3c 100644 --- a/src/matrix/tight.rs +++ b/src/matrix/tight.rs @@ -2,13 +2,14 @@ use super::interface::*; use super::visualize::*; use crate::util::*; use derivative::Derivative; +use crate::dual_module_pq::{EdgeWeak, VertexWeak}; #[derive(Clone, Derivative)] #[derivative(Default(new = "true"))] pub struct Tight { base: M, /// the set of tight edges: should be a relatively small set - tight_edges: FastIterSet, + tight_edges: FastIterSet, /// tight matrix gives a view of only tight edges, with sorted indices #[derivative(Default(value = "true"))] is_var_indices_outdated: bool, @@ -33,38 +34,38 @@ impl Tight { } impl MatrixTight for Tight { - fn update_edge_tightness(&mut self, edge_index: EdgeIndex, is_tight: bool) { - debug_assert!(self.exists_edge(edge_index)); + fn update_edge_tightness(&mut self, edge_weak: EdgeWeak, is_tight: bool) { + debug_assert!(self.exists_edge(edge_weak.clone())); self.is_var_indices_outdated = true; if is_tight { - self.tight_edges.insert(edge_index); + self.tight_edges.insert(edge_weak.clone()); } else { - self.tight_edges.remove(&edge_index); + self.tight_edges.remove(&edge_weak); } } - fn is_tight(&self, edge_index: usize) -> bool { - debug_assert!(self.exists_edge(edge_index)); - self.tight_edges.contains(&edge_index) + fn is_tight(&self, edge_weak: EdgeWeak) -> bool { + debug_assert!(self.exists_edge(edge_weak.clone())); + self.tight_edges.contains(&edge_weak) } - fn get_tight_edges(&self) -> &FastIterSet { + fn get_tight_edges(&self) -> &FastIterSet { &self.tight_edges } } impl MatrixBasic for Tight { - fn add_variable(&mut self, edge_index: EdgeIndex) -> Option { - self.base.add_variable(edge_index) + fn add_variable(&mut self, edge_weak: EdgeWeak) -> Option { + self.base.add_variable(edge_weak) } fn add_constraint( &mut self, - vertex_index: VertexIndex, - incident_edges: &[EdgeIndex], + vertex_weak: VertexWeak, + incident_edges: &[EdgeWeak], parity: bool, ) -> Option> { - self.base.add_constraint(vertex_index, incident_edges, parity) + self.base.add_constraint(vertex_weak.clone(), incident_edges, parity) } fn xor_row(&mut self, target: RowIndex, source: RowIndex) { @@ -79,16 +80,16 @@ impl MatrixBasic for Tight { fn get_rhs(&self, row: RowIndex) -> bool { self.get_base().get_rhs(row) } - fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeIndex { + fn var_to_edge_index(&self, var_index: VarIndex) -> EdgeWeak { self.get_base().var_to_edge_index(var_index) } - fn edge_to_var_index(&self, edge_index: EdgeIndex) -> Option { - self.get_base().edge_to_var_index(edge_index) + fn edge_to_var_index(&self, edge_weak: EdgeWeak) -> Option { + self.get_base().edge_to_var_index(edge_weak) } - fn get_vertices(&self) -> FastIterSet { + fn get_vertices(&self) -> FastIterSet { self.get_base().get_vertices() } - fn get_edges(&self) -> FastIterSet { + fn get_edges(&self) -> FastIterSet { self.get_base().get_edges() } } @@ -98,8 +99,8 @@ impl Tight { self.var_indices.clear(); for column in 0..self.base.columns() { let var_index = self.base.column_to_var_index(column); - let edge_index = self.base.var_to_edge_index(var_index); - if self.is_tight(edge_index) { + let edge_ptr = self.base.var_to_edge_index(var_index); + if self.is_tight(edge_ptr) { self.var_indices.push(var_index); } } @@ -139,6 +140,9 @@ impl VizTrait for Tight { pub mod tests { use super::super::basic::*; use super::*; + use crate::matrix::basic::tests::{initialize_vertex_edges_for_matrix_testing, edge_vec_from_indices}; + use std::collections::HashSet; + use crate::dual_module_pq::{EdgePtr, VertexPtr}; type TightMatrix = Tight; @@ -146,9 +150,18 @@ pub mod tests { fn tight_matrix_1() { // cargo test --features=colorful tight_matrix_1 -- --nocapture let mut matrix = TightMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); matrix.printstd(); // this is because by default all edges are not tight assert_eq!( @@ -165,8 +178,8 @@ pub mod tests { └─┴───┘ " ); - matrix.update_edge_tightness(4, true); - matrix.update_edge_tightness(9, true); + matrix.update_edge_tightness(edges[1].downgrade(), true); + matrix.update_edge_tightness(edges[3].downgrade(), true); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -182,7 +195,7 @@ pub mod tests { └─┴─┴─┴───┘ " ); - matrix.update_edge_tightness(9, false); + matrix.update_edge_tightness(edges[3].downgrade(), false); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), @@ -205,8 +218,15 @@ pub mod tests { fn tight_matrix_cannot_set_nonexistent_edge() { // cargo test tight_matrix_cannot_set_nonexistent_edge -- --nocapture let mut matrix = TightMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.update_edge_tightness(2, true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9, 2]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.update_edge_tightness(edges[4].downgrade(), true); } #[test] @@ -214,22 +234,37 @@ pub mod tests { fn tight_matrix_cannot_read_nonexistent_edge() { // cargo test tight_matrix_cannot_read_nonexistent_edge -- --nocapture let mut matrix = TightMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.is_tight(2); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9, 2]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.is_tight(edges[4].downgrade()); } #[test] fn tight_matrix_basic_trait() { // cargo test --features=colorful tight_matrix_basic_trait -- --nocapture let mut matrix = TightMatrix::new(); - matrix.add_variable(3); // untight edges will not show - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9, 3]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_variable(edges[4].downgrade()); // untight edges will not show + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); matrix.swap_row(2, 1); matrix.xor_row(0, 1); - for edge_index in [1, 4, 6, 9] { - matrix.update_edge_tightness(edge_index, true); + for edge_index in 0..4 { + matrix.update_edge_tightness(edges[edge_index].downgrade(), true); } matrix.printstd(); assert_eq!( @@ -252,19 +287,29 @@ pub mod tests { fn tight_matrix_rebuild_var_indices() { // cargo test --features=colorful tight_matrix_rebuild_var_indices -- --nocapture let mut matrix = TightMatrix::new(); - matrix.add_variable(3); // untight edges will not show - matrix.add_constraint(0, &[1, 4, 6], true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 6, 9, 3]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + + matrix.add_variable(edges[4].downgrade()); // untight edges will not show + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); assert_eq!(matrix.columns(), 0); - for edge_index in [1, 4, 6] { - matrix.update_edge_tightness(edge_index, true); + for edge_index in 0..3 { + matrix.update_edge_tightness(edges[edge_index].downgrade(), true); } assert_eq!(matrix.columns(), 3); assert_eq!(matrix.columns(), 3); // should only update var_indices_once - matrix.add_constraint(1, &[4, 9], false); - matrix.add_constraint(2, &[1, 9], true); - matrix.update_edge_tightness(9, true); - matrix.update_edge_tightness(4, false); - matrix.update_edge_tightness(6, false); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); + + matrix.update_edge_tightness(edges[3].downgrade(), true); + matrix.update_edge_tightness(edges[1].downgrade(), false); + matrix.update_edge_tightness(edges[2].downgrade(), false); assert_eq!(matrix.columns(), 2); matrix.printstd(); assert_eq!( @@ -288,8 +333,14 @@ pub mod tests { fn tight_matrix_cannot_call_dirty_column() { // cargo test tight_matrix_cannot_call_dirty_column -- --nocapture let mut matrix = TightMatrix::new(); - matrix.add_constraint(0, &[1, 4, 6], true); - matrix.update_edge_tightness(1, true); + let vertex_indices = vec![0]; + let edge_indices = vec![1, 4, 6, 9]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.update_edge_tightness(edges[0].downgrade(), true); // even though there is indeed such a column, we forbid such dangerous calls // always call `columns()` before accessing any column matrix.column_to_var_index(0); diff --git a/src/matrix/visualize.rs b/src/matrix/visualize.rs index e6fa5663..2e35fd85 100644 --- a/src/matrix/visualize.rs +++ b/src/matrix/visualize.rs @@ -57,7 +57,8 @@ impl From<&mut M> for VizTable { let mut edges = vec![]; for column in 0..matrix.columns() { let var_index = matrix.column_to_var_index(column); - let edge_index = matrix.var_to_edge_index(var_index); + let edge_weak = matrix.var_to_edge_index(var_index); + let edge_index = edge_weak.upgrade_force().read_recursive().edge_index; edges.push(edge_index); let edge_index_str = Self::force_single_column(edge_index.to_string().as_str()); title.add_cell(Cell::new(edge_index_str.as_str()).style_spec("brFm")); @@ -142,14 +143,25 @@ impl VizTable { #[cfg(test)] pub mod tests { use super::super::*; + use crate::matrix::basic::tests::{initialize_vertex_edges_for_matrix_testing, edge_vec_from_indices}; + use std::collections::HashSet; + use crate::dual_module_pq::{EdgePtr, VertexPtr}; #[test] fn viz_table_1() { // cargo test --features=colorful viz_table_1 -- --nocapture let mut matrix = BasicMatrix::new(); - matrix.add_constraint(0, &[1, 4, 16], true); - matrix.add_constraint(1, &[4, 23], false); - matrix.add_constraint(2, &[1, 23], true); + let vertex_indices = vec![0, 1, 2]; + let edge_indices = vec![1, 4, 16, 23]; + let vertex_incident_edges_vec = vec![ + vec![0, 1, 2], + vec![1, 3], + vec![0, 3], + ]; + let (vertices, edges) = initialize_vertex_edges_for_matrix_testing(vertex_indices, edge_indices); + matrix.add_constraint(vertices[0].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[0], &edges), true); + matrix.add_constraint(vertices[1].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[1], &edges), false); + matrix.add_constraint(vertices[2].downgrade(), &edge_vec_from_indices(&vertex_incident_edges_vec[2], &edges), true); matrix.printstd(); assert_eq!( matrix.clone().printstd_str(), diff --git a/src/mwpf_solver.rs b/src/mwpf_solver.rs index 154259ed..dece190d 100644 --- a/src/mwpf_solver.rs +++ b/src/mwpf_solver.rs @@ -17,6 +17,8 @@ use crate::primal_module::*; use crate::primal_module_serial::*; use crate::util::*; use crate::visualize::*; +use crate::pointers::*; +use crate::dual_module_pq::VertexPtr; use bp::bp::BpDecoder; @@ -295,8 +297,8 @@ macro_rules! bind_trait_to_python { macro_rules! inherit_solver_plugin_methods { ($struct_name:ident) => { impl $struct_name { - pub fn get_cluster(&self, vertex_index: VertexIndex) -> Cluster { - self.0.get_cluster(vertex_index) + pub fn get_cluster(&self, vertex_ptr: VertexPtr) -> Cluster { + self.0.get_cluster(vertex_ptr) } } }; @@ -313,7 +315,7 @@ pub struct SolverSerialPluginsConfig { #[derive(Clone)] pub struct SolverSerialPlugins { - dual_module: DualModulePQ, + pub dual_module: DualModulePQ, primal_module: PrimalModuleSerial, interface_ptr: DualModuleInterfacePtr, model_graph: Arc, @@ -338,9 +340,9 @@ impl SolverSerialPlugins { primal_module.plugins = plugins; primal_module.config = config.primal.as_ref().unwrap_or(&config.flatten_primal).clone(); Self { - dual_module: DualModulePQ::new_empty(initializer), + dual_module: DualModulePQ::new_empty(initializer, 0), // TODO: double check the partition id here primal_module, - interface_ptr: DualModuleInterfacePtr::new(model_graph.clone()), + interface_ptr: DualModuleInterfacePtr::new(model_graph.clone(), 0), // TODO: double check the partition id here model_graph, config, syndrome_loaded: false, @@ -361,8 +363,8 @@ impl SolverSerialPlugins { if !skip_initial_duals { self.interface_ptr - .load(Arc::new(syndrome_pattern.clone()), &mut self.dual_module); - self.primal_module.load(&self.interface_ptr, &mut self.dual_module); + .load(Arc::new(syndrome_pattern.clone()), &mut self.dual_module, 0); // TODO: double check the partition id here + self.primal_module.load(&self.interface_ptr, &mut self.dual_module, 0); // TODO: double check the partition id here } else { self.interface_ptr .write() @@ -384,43 +386,47 @@ impl SolverSerialPlugins { } /// get the cluster information of a vertex - pub fn get_cluster(&self, vertex_index: VertexIndex) -> Cluster { + pub fn get_cluster(&self, vertex_ptr: VertexPtr) -> Cluster { let mut cluster = Cluster::new(); // visit the graph via tight edges let mut current_vertices = FastIterSet::new(); - current_vertices.insert(vertex_index); + current_vertices.insert(vertex_ptr); while !current_vertices.is_empty() { let mut next_vertices = FastIterSet::new(); - for &vertex_index in current_vertices.iter() { - cluster.add_vertex(vertex_index); - for &edge_index in self.model_graph.get_vertex_neighbors(vertex_index).iter() { - if self.dual_module.is_edge_tight(edge_index) { - cluster.add_edge(edge_index); - cluster.parity_matrix.add_tight_variable(edge_index); - for &next_vertex_index in self.model_graph.get_edge_neighbors(edge_index).iter() { - if !cluster.vertices.contains(&next_vertex_index) { - next_vertices.insert(next_vertex_index); + for vertex_ptr0 in current_vertices.iter() { + cluster.add_vertex(vertex_ptr0.clone()); + for edge_weak in vertex_ptr0.read_recursive().edges.iter() { + let edge_ptr = edge_weak.upgrade_force(); + if self.dual_module.is_edge_tight(edge_ptr.clone()) { + cluster.add_edge(edge_ptr.clone()); + cluster.parity_matrix.add_tight_variable(edge_weak.clone()); + for next_vertex_weak in edge_ptr.read_recursive().vertices.iter() { + let next_vertex_ptr = next_vertex_weak.upgrade_force(); + if !cluster.vertices.contains(&next_vertex_ptr) { + next_vertices.insert(next_vertex_ptr); } } } else { - cluster.add_hair(edge_index); + cluster.add_hair(edge_ptr.clone()); } } } current_vertices = next_vertices; } // add dual variables - for &edge_index in cluster.edges.iter() { - for node_ptr in self.dual_module.get_edge_nodes(edge_index).iter() { + for edge_ptr in cluster.edges.iter() { + for node_ptr in self.dual_module.get_edge_nodes(edge_ptr.clone()).iter() { cluster.nodes.insert(node_ptr.clone().into()); } } // construct the parity matrix - let interface = self.interface_ptr.read(); - for &vertex_index in cluster.vertices.iter() { - let incident_edges = self.model_graph.get_vertex_neighbors(vertex_index); - let parity = interface.decoding_graph.is_vertex_defect(vertex_index); - cluster.parity_matrix.add_constraint(vertex_index, incident_edges, parity); + let interface = self.interface_ptr.read_recursive(); + for vertex_ptr in cluster.vertices.iter() { + let vertex_weak = vertex_ptr.downgrade(); + let vertex = vertex_ptr.read_recursive(); + let incident_edges = &vertex.edges; + let parity = vertex.is_defect; + cluster.parity_matrix.add_constraint(vertex_weak, incident_edges, parity); } cluster } @@ -629,6 +635,10 @@ impl SolverSerialJointSingleHair { config, )) } + + pub fn get_dual_module(&self) -> &DualModulePQ { + &self.0.dual_module + } } #[cfg(feature = "python_binding")] diff --git a/src/plugin.rs b/src/plugin.rs index 5f5aa6fa..98379dd0 100644 --- a/src/plugin.rs +++ b/src/plugin.rs @@ -5,7 +5,6 @@ //! A plugin must implement Clone trait, because it will be cloned multiple times for each cluster //! -use crate::decoding_hypergraph::*; use crate::derivative::Derivative; use crate::dual_module::*; use crate::matrix::*; @@ -23,7 +22,7 @@ pub trait PluginImpl: std::fmt::Debug { /// given the tight edges and parity constraints, find relaxers fn find_relaxers( &self, - decoding_graph: &DecodingHyperGraph, + dual_module: &mut dyn DualModuleImpl, matrix: &mut EchelonMatrix, positive_dual_nodes: &[DualNodePtr], ) -> RelaxerVec; @@ -75,7 +74,7 @@ pub struct PluginEntry { impl PluginEntry { pub fn execute( &self, - decoding_graph: &DecodingHyperGraph, + dual_module: &mut impl DualModuleImpl, matrix: &mut EchelonMatrix, positive_dual_nodes: &[DualNodePtr], relaxer_forest: &mut RelaxerForest, @@ -84,13 +83,13 @@ impl PluginEntry { let mut repeat_count = 0; while repeat { // execute the plugin - let relaxers = self.plugin.find_relaxers(decoding_graph, &mut *matrix, positive_dual_nodes); + let relaxers = self.plugin.find_relaxers(dual_module, &mut *matrix, positive_dual_nodes); if relaxers.is_empty() { repeat = false; } for relaxer in relaxers.into_iter() { - for edge_index in relaxer.get_untighten_edges().keys() { - matrix.update_edge_tightness(*edge_index, false); + for edge_ptr in relaxer.get_untighten_edges().keys() { + matrix.update_edge_tightness(edge_ptr.downgrade(), false); } let relaxer = Arc::new(relaxer); let sum_speed = relaxer.get_sum_speed(); @@ -137,7 +136,7 @@ impl PluginManager { pub fn find_relaxer( &mut self, - decoding_graph: &DecodingHyperGraph, + dual_module: &mut impl DualModuleImpl, matrix: &mut EchelonMatrix, positive_dual_nodes: &[DualNodePtr], ) -> Option { @@ -148,11 +147,11 @@ impl PluginManager { .map(|ptr| ptr.read_recursive().invalid_subgraph.clone()), ); for plugin_entry in self.plugins.iter().take(*self.plugin_count.read_recursive()) { - if let Some(relaxer) = plugin_entry.execute(decoding_graph, matrix, positive_dual_nodes, &mut relaxer_forest) { + if let Some(relaxer) = plugin_entry.execute(dual_module, matrix, positive_dual_nodes, &mut relaxer_forest) { return Some(relaxer); } } // add a union find relaxer finder as the last resort if nothing is reported - PluginUnionFind::entry().execute(decoding_graph, matrix, positive_dual_nodes, &mut relaxer_forest) + PluginUnionFind::entry().execute(dual_module, matrix, positive_dual_nodes, &mut relaxer_forest) } } diff --git a/src/plugin_single_hair.rs b/src/plugin_single_hair.rs index 86a2a2e8..c52c8000 100644 --- a/src/plugin_single_hair.rs +++ b/src/plugin_single_hair.rs @@ -5,7 +5,6 @@ //! A plugin must implement Clone trait, because it will be cloned multiple times for each cluster //! -use crate::decoding_hypergraph::*; use crate::dual_module::*; use crate::invalid_subgraph::InvalidSubgraph; use crate::matrix::*; @@ -13,7 +12,9 @@ use crate::plugin::*; use crate::plugin_union_find::*; use crate::relaxer::*; use crate::util::*; +use crate::pointers::*; use num_traits::One; +use crate::dual_module_pq::{EdgePtr, VertexPtr}; use std::sync::Arc; @@ -23,19 +24,19 @@ pub struct PluginSingleHair {} impl PluginImpl for PluginSingleHair { fn find_relaxers( &self, - decoding_graph: &DecodingHyperGraph, + dual_module: &mut dyn DualModuleImpl, matrix: &mut EchelonMatrix, positive_dual_nodes: &[DualNodePtr], ) -> Vec { // single hair requires the matrix to have at least one feasible solution - if let Some(relaxer) = PluginUnionFind::find_single_relaxer(decoding_graph, matrix) { - return vec![relaxer]; + if let Some(relaxer) = PluginUnionFind::find_single_relaxer(dual_module, matrix) { + return vec![relaxer] } // then try to find more relaxers let mut relaxers = vec![]; for dual_node_ptr in positive_dual_nodes.iter() { let dual_node = dual_node_ptr.read_recursive(); - let mut hair_view = HairView::new(matrix, dual_node.invalid_subgraph.hair.iter().cloned()); + let mut hair_view = HairView::new(matrix, dual_node.invalid_subgraph.hair.iter().map(|e| e.downgrade())); debug_assert!(hair_view.get_echelon_satisfiable()); // hair_view.printstd(); // optimization: check if there exists a single-hair solution, if not, clear the previous relaxers @@ -65,22 +66,25 @@ impl PluginImpl for PluginSingleHair { if !unnecessary_edges.is_empty() { // we can construct a relaxer here, by growing a new invalid subgraph that // removes those unnecessary edges and shrinking the existing one - let mut vertices: FastIterSet = hair_view.get_vertices(); - let mut edges: FastIterSet = FastIterSet::from_iter(hair_view.get_base_view_edges()); - for &edge_index in dual_node.invalid_subgraph.hair.iter() { - edges.remove(&edge_index); + let mut vertices: FastIterSet = hair_view.get_vertices().iter().map(|v| v.upgrade_force()).collect::>(); + let mut edges: FastIterSet = hair_view.get_base_view_edges().iter().map(|e| e.upgrade_force()).collect::>(); + for edge_ptr in dual_node.invalid_subgraph.hair.iter() { + edges.remove(&edge_ptr); } - for &edge_index in unnecessary_edges.iter() { - edges.insert(edge_index); - vertices.extend(decoding_graph.get_edge_neighbors(edge_index)); + for edge_weak in unnecessary_edges.iter() { + let edge_ptr = edge_weak.upgrade_force(); + edges.insert(edge_ptr.clone()); + for vertex_weak in edge_ptr.read_recursive().vertices.iter() { + vertices.insert(vertex_weak.upgrade_force()); + } } - let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete(vertices, edges, decoding_graph)); + let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete(vertices, edges, dual_module)); let relaxer = Relaxer::new( [ (invalid_subgraph, Rational::one()), (dual_node.invalid_subgraph.clone(), -Rational::one()), ] - .into(), + .into_iter().collect(), ); relaxers.push(relaxer); } diff --git a/src/plugin_union_find.rs b/src/plugin_union_find.rs index 4a7b3571..5a328167 100644 --- a/src/plugin_union_find.rs +++ b/src/plugin_union_find.rs @@ -13,33 +13,35 @@ use crate::num_traits::One; use crate::plugin::*; use crate::relaxer::*; use crate::util::*; +use crate::dual_module_pq::EdgePtr; #[derive(Debug, Clone, Default)] pub struct PluginUnionFind {} impl PluginUnionFind { /// check if the cluster is valid (hypergraph union-find decoder) - pub fn find_single_relaxer(decoding_graph: &DecodingHyperGraph, matrix: &mut EchelonMatrix) -> Option { + pub fn find_single_relaxer(dual_module: &mut dyn DualModuleImpl, matrix: &mut EchelonMatrix) -> Option { if matrix.get_echelon_info().satisfiable { return None; // cannot find any relaxer } + let local_edges: FastIterSet = matrix.get_view_edges().iter().map(|e| e.upgrade_force()).collect::>(); let invalid_subgraph = InvalidSubgraph::new_complete_ptr( - matrix.get_vertices(), - FastIterSet::from_iter(matrix.get_view_edges()), - decoding_graph, + matrix.get_vertices().iter().map(|e| e.upgrade_force()).collect::>(), + local_edges, + dual_module ); - Some(Relaxer::new([(invalid_subgraph, Rational::one())].into())) + Some(Relaxer::new([(invalid_subgraph, Rational::one())].into_iter().collect())) } } impl PluginImpl for PluginUnionFind { fn find_relaxers( &self, - decoding_graph: &DecodingHyperGraph, + dual_module: &mut dyn DualModuleImpl, matrix: &mut EchelonMatrix, _positive_dual_nodes: &[DualNodePtr], ) -> Vec { - if let Some(relaxer) = Self::find_single_relaxer(decoding_graph, matrix) { + if let Some(relaxer) = Self::find_single_relaxer(dual_module, matrix) { vec![relaxer] } else { vec![] diff --git a/src/pointers.rs b/src/pointers.rs index 456c34e5..1398218d 100644 --- a/src/pointers.rs +++ b/src/pointers.rs @@ -1,118 +1,599 @@ //! Pointer Types //! +//! This module provides a unified interface for reference-counted objects with interior mutability. +//! +//! Feature Flags: +//! - Default: Uses `OrdArcRwLock` (Safe, Thread-safe, Ordered). +//! - `unsafe_pointer`: Uses `OrdArcUnsafe` (Unsafe, High-performance, Ordered). + +#![cfg_attr(feature = "unsafe_pointer", allow(dropping_references))] -use crate::parking_lot::lock_api::{RwLockReadGuard, RwLockWriteGuard}; -use crate::parking_lot::{RawRwLock, RwLock}; use std::sync::{Arc, Weak}; +use std::hash::Hash; -pub trait RwLockPtr { - fn new_ptr(ptr: Arc>) -> Self; +// ========================================================================================= +// UNSAFE IMPLEMENTATION (feature = "unsafe_pointer") +// Uses UnsafeCell for maximum performance, bypassing locks. +// Includes Standard (ArcUnsafe) and Ordered (OrdArcUnsafe) variants. +// ========================================================================================= +cfg_if::cfg_if! { + if #[cfg(feature="unsafe_pointer")] { + use std::cell::UnsafeCell; - fn new_value(obj: ObjType) -> Self; + /// Trait for unsafe pointer operations. + /// WARNING: The user must ensure no data races occur. + pub trait UnsafePtr { + fn new_ptr(ptr: Arc>) -> Self; + fn new_value(obj: ObjType) -> Self; + + fn ptr(&self) -> &Arc>; + fn ptr_mut(&mut self) -> &mut Arc>; - fn ptr(&self) -> &Arc>; + #[inline(always)] + fn read_recursive(&self) -> &ObjType { + // SAFETY: User promises no concurrent mutable access. + unsafe { &*self.ptr().get() } + } - fn ptr_mut(&mut self) -> &mut Arc>; + #[inline(always)] + fn write(&self) -> &mut ObjType { + // SAFETY: User promises uniqueness. + unsafe { &mut *self.ptr().get() } + } + + #[inline(always)] + fn try_write(&self) -> Option<&mut ObjType> { + Some(self.write()) + } - #[inline(always)] - fn read_recursive(&self) -> RwLockReadGuard { - let ret = self.ptr().read_recursive(); - ret - } + fn ptr_eq(&self, other: &Self) -> bool { + Arc::ptr_eq(self.ptr(), other.ptr()) + } + } - #[inline(always)] - fn write(&self) -> RwLockWriteGuard { - let ret = self.ptr().write(); - ret - } + // --- STRUCT DEFINITIONS --- - fn ptr_eq(&self, other: &Self) -> bool { - Arc::ptr_eq(self.ptr(), other.ptr()) - } -} + // 1. Standard ArcUnsafe (Identity based) + pub struct ArcUnsafe { + ptr: Arc>, + } -pub struct ArcRwLock { - ptr: Arc>, -} + pub struct WeakUnsafe { + ptr: Weak>, + } -pub struct WeakRwLock { - ptr: Weak>, -} + // 2. Ordered ArcUnsafe (Sort key based) + pub struct OrdArcUnsafe { + ord: U, + ptr: Arc>, + } -impl ArcRwLock { - pub fn downgrade(&self) -> WeakRwLock { - WeakRwLock:: { - ptr: Arc::downgrade(&self.ptr), + pub struct OrdWeakUnsafe { + ord: U, + ptr: Weak>, } - } -} -impl WeakRwLock { - pub fn upgrade_force(&self) -> ArcRwLock { - ArcRwLock:: { - ptr: self.ptr.upgrade().unwrap(), + // --- MANUAL SEND/SYNC --- + // SAFETY: We mimic RwLock behavior. The user is responsible for race conditions. + unsafe impl Sync for ArcUnsafe {} + unsafe impl Send for ArcUnsafe {} + unsafe impl Sync for WeakUnsafe {} + unsafe impl Send for WeakUnsafe {} + + unsafe impl Sync for OrdArcUnsafe {} + unsafe impl Send for OrdArcUnsafe {} + unsafe impl Sync for OrdWeakUnsafe {} + unsafe impl Send for OrdWeakUnsafe {} + + // --- IMPLEMENTATIONS: ArcUnsafe (Standard) --- + impl ArcUnsafe { + pub fn downgrade(&self) -> WeakUnsafe { + WeakUnsafe:: { + ptr: Arc::downgrade(&self.ptr) + } + } } - } - pub fn upgrade(&self) -> Option> { - self.ptr.upgrade().map(|x| ArcRwLock:: { ptr: x }) - } - pub fn ptr_eq(&self, other: &Self) -> bool { - Weak::ptr_eq(&self.ptr, &other.ptr) - } -} -impl Clone for ArcRwLock { - fn clone(&self) -> Self { - Self::new_ptr(Arc::clone(self.ptr())) - } -} + impl WeakUnsafe { + pub fn upgrade_force(&self) -> ArcUnsafe { + ArcUnsafe:: { + ptr: self.ptr.upgrade().unwrap() + } + } + pub fn upgrade(&self) -> Option> { + self.ptr.upgrade().map(|x| ArcUnsafe:: { ptr: x }) + } + pub fn ptr_eq(&self, other: &Self) -> bool { + Weak::ptr_eq(&self.ptr, &other.ptr) + } + } -impl RwLockPtr for ArcRwLock { - fn new_ptr(ptr: Arc>) -> Self { - Self { ptr } - } - fn new_value(obj: T) -> Self { - Self::new_ptr(Arc::new(RwLock::new(obj))) - } - #[inline(always)] - fn ptr(&self) -> &Arc> { - &self.ptr - } - #[inline(always)] - fn ptr_mut(&mut self) -> &mut Arc> { - &mut self.ptr - } -} + impl Clone for ArcUnsafe { + fn clone(&self) -> Self { + Self { ptr: self.ptr.clone() } + } + } -impl PartialEq for ArcRwLock { - fn eq(&self, other: &Self) -> bool { - self.ptr_eq(other) - } -} + impl UnsafePtr for ArcUnsafe { + fn new_ptr(ptr: Arc>) -> Self { Self { ptr } } + fn new_value(obj: T) -> Self { Self::new_ptr(Arc::new(UnsafeCell::new(obj))) } + #[inline(always)] fn ptr(&self) -> &Arc> { &self.ptr } + #[inline(always)] fn ptr_mut(&mut self) -> &mut Arc> { &mut self.ptr } + } -impl Eq for ArcRwLock {} + impl WeakUnsafe { + #[inline(always)] pub fn ptr(&self) -> &Weak> { &self.ptr } + #[inline(always)] pub fn ptr_mut(&mut self) -> &mut Weak> { &mut self.ptr } + } -impl Clone for WeakRwLock { - fn clone(&self) -> Self { - Self { ptr: self.ptr.clone() } - } -} + impl PartialEq for ArcUnsafe { + fn eq(&self, other: &Self) -> bool { self.ptr_eq(other) } + } + impl Eq for ArcUnsafe { } + + impl Clone for WeakUnsafe { + fn clone(&self) -> Self { + Self { ptr: self.ptr.clone() } + } + } + + impl PartialEq for WeakUnsafe { + fn eq(&self, other: &Self) -> bool { self.ptr.ptr_eq(&other.ptr) } + } + impl Eq for WeakUnsafe { } + + // IDENTITY HASHING + impl std::hash::Hash for ArcUnsafe { + fn hash(&self, state: &mut H) { + let address = Arc::as_ptr(&self.ptr); + address.hash(state); + } + } + + impl std::hash::Hash for WeakUnsafe { + fn hash(&self, state: &mut H) { + let address = Weak::as_ptr(&self.ptr); + address.hash(state); + } + } + + impl weak_table::traits::WeakElement for WeakUnsafe { + type Strong = ArcUnsafe; + fn new(view: &ArcUnsafe) -> Self { view.downgrade() } + fn view(&self) -> Option> { self.upgrade() } + fn clone(view: &ArcUnsafe) -> ArcUnsafe { view.clone() } + } + + impl std::ops::Deref for ArcUnsafe { + type Target = std::cell::UnsafeCell; + fn deref(&self) -> &Self::Target { &self.ptr } + } + + // --- IMPLEMENTATIONS: OrdArcUnsafe (Ordered) --- + impl OrdArcUnsafe { + pub fn downgrade(&self) -> OrdWeakUnsafe { + OrdWeakUnsafe:: { + ord: self.ord, + ptr: Arc::downgrade(&self.ptr), + } + } + } + + impl OrdWeakUnsafe { + pub fn upgrade_force(&self) -> OrdArcUnsafe { + OrdArcUnsafe:: { + ord: self.ord, + ptr: self.ptr.upgrade().unwrap(), + } + } + pub fn upgrade(&self) -> Option> { + self.ptr.upgrade().map(|x| OrdArcUnsafe:: { ord: self.ord, ptr: x }) + } + pub fn ptr_eq(&self, other: &Self) -> bool { + Weak::ptr_eq(&self.ptr, &other.ptr) + } + } + + impl Clone for OrdArcUnsafe { + fn clone(&self) -> Self { + Self::new_ptr(Arc::clone(self.ptr()), self.ord) + } + } + + impl OrdArcUnsafe { + pub fn new_ptr(ptr: Arc>, ord: U) -> Self { + Self { ord, ptr } + } + pub fn new_value(obj: T, ord: U) -> Self { + Self::new_ptr(Arc::new(UnsafeCell::new(obj)), ord) + } + #[inline(always)] + pub fn ptr(&self) -> &Arc> { + &self.ptr + } + #[inline(always)] + pub fn ptr_mut(&mut self) -> &mut Arc> { + &mut self.ptr + } + + // Helper methods to match UnsafePtr-like usage + #[inline(always)] + pub fn read_recursive(&self) -> &T { + unsafe { &*self.ptr.get() } + } + #[inline(always)] + pub fn write(&self) -> &mut T { + unsafe { &mut *self.ptr.get() } + } + + #[inline(always)] + pub fn get_ord(&self) -> U { + self.ord + } + } + + impl OrdWeakUnsafe { + #[inline(always)] + pub fn ptr(&self) -> &Weak> { + &self.ptr + } + } + + // Ordering and Hashing for OrdArcUnsafe + impl PartialEq for OrdArcUnsafe { + fn eq(&self, other: &Self) -> bool { self.ord.eq(&other.ord) } + } + impl Eq for OrdArcUnsafe {} + + impl Ord for OrdArcUnsafe { + fn cmp(&self, other: &Self) -> std::cmp::Ordering { self.ord.cmp(&other.ord) } + } + + impl PartialOrd for OrdArcUnsafe { + fn partial_cmp(&self, other: &Self) -> Option { Some(self.cmp(other)) } + } + + impl std::hash::Hash for OrdArcUnsafe { + fn hash(&self, state: &mut H) { self.ord.hash(state); } + } + + // Ordering and Hashing for OrdWeakUnsafe + impl PartialEq for OrdWeakUnsafe { + fn eq(&self, other: &Self) -> bool { self.ord.eq(&other.ord) } + } + impl Eq for OrdWeakUnsafe {} + + impl Ord for OrdWeakUnsafe { + fn cmp(&self, other: &Self) -> std::cmp::Ordering { self.ord.cmp(&other.ord) } + } + + impl PartialOrd for OrdWeakUnsafe { + fn partial_cmp(&self, other: &Self) -> Option { Some(self.cmp(other)) } + } + + impl std::hash::Hash for OrdWeakUnsafe { + fn hash(&self, state: &mut H) { self.ord.hash(state); } + } + + impl Clone for OrdWeakUnsafe { + fn clone(&self) -> Self { + Self { ord: self.ord, ptr: self.ptr.clone() } + } + } + + // WeakTable Integration + impl weak_table::traits::WeakElement for OrdWeakUnsafe { + type Strong = OrdArcUnsafe; + fn new(view: &Self::Strong) -> Self { view.downgrade() } + fn view(&self) -> Option { self.upgrade() } + fn clone(view: &Self::Strong) -> Self::Strong { view.clone() } + } + + impl std::ops::Deref for OrdArcUnsafe { + type Target = std::cell::UnsafeCell; + fn deref(&self) -> &Self::Target { &self.ptr } + } -impl PartialEq for WeakRwLock { - fn eq(&self, other: &Self) -> bool { - self.ptr.ptr_eq(&other.ptr) + // --- TYPE ALIASES --- + // Defaults to Ordered Unsafe Pointer + pub type ArcManualSafeLock = OrdArcUnsafe; + pub type WeakManualSafeLock = OrdWeakUnsafe; } } -impl Eq for WeakRwLock {} +// ========================================================================================= +// SAFE IMPLEMENTATION (DEFAULT) +// Uses parking_lot::RwLock for standard thread safety. +// Includes Standard (ArcRwLock) and Ordered (OrdArcRwLock) variants. +// ========================================================================================= +cfg_if::cfg_if! { + if #[cfg(not(feature="unsafe_pointer"))] { + use crate::parking_lot::lock_api::{RwLockReadGuard, RwLockWriteGuard}; + use crate::parking_lot::{RawRwLock, RwLock}; -impl std::ops::Deref for ArcRwLock { - type Target = RwLock; - fn deref(&self) -> &Self::Target { - &self.ptr + // --- TRAIT DEFINITION --- + pub trait RwLockPtr { + fn new_ptr(ptr: Arc>) -> Self; + fn new_value(obj: ObjType) -> Self; + fn ptr(&self) -> &Arc>; + fn ptr_mut(&mut self) -> &mut Arc>; + + #[inline(always)] + fn read_recursive(&self) -> RwLockReadGuard { + self.ptr().read_recursive() + } + + #[inline(always)] + fn write(&self) -> RwLockWriteGuard { + self.ptr().write() + } + + fn ptr_eq(&self, other: &Self) -> bool { + Arc::ptr_eq(self.ptr(), other.ptr()) + } + } + + // --- STRUCT DEFINITIONS --- + + // 1. Standard ArcRwLock + pub struct ArcRwLock { + ptr: Arc>, + } + + pub struct WeakRwLock { + ptr: Weak>, + } + + // 2. Ordered ArcRwLock (For deterministic sorting) + pub struct OrdArcRwLock { + pub ord: U, + pub ptr: Arc>, + } + + pub struct OrdWeakRwLock { + pub ord: U, + pub ptr: Weak>, + } + + // --- IMPLEMENTATIONS: ArcRwLock (Standard) --- + impl ArcRwLock { + pub fn downgrade(&self) -> WeakRwLock { + WeakRwLock:: { + ptr: Arc::downgrade(&self.ptr), + } + } + } + + impl WeakRwLock { + pub fn upgrade_force(&self) -> ArcRwLock { + ArcRwLock:: { + ptr: self.ptr.upgrade().unwrap(), + } + } + pub fn upgrade(&self) -> Option> { + self.ptr.upgrade().map(|x| ArcRwLock:: { ptr: x }) + } + pub fn ptr_eq(&self, other: &Self) -> bool { + Weak::ptr_eq(&self.ptr, &other.ptr) + } + } + + impl Clone for ArcRwLock { + fn clone(&self) -> Self { + Self::new_ptr(Arc::clone(self.ptr())) + } + } + + impl RwLockPtr for ArcRwLock { + fn new_ptr(ptr: Arc>) -> Self { Self { ptr } } + fn new_value(obj: T) -> Self { Self::new_ptr(Arc::new(RwLock::new(obj))) } + #[inline(always)] fn ptr(&self) -> &Arc> { &self.ptr } + #[inline(always)] fn ptr_mut(&mut self) -> &mut Arc> { &mut self.ptr } + } + + impl WeakRwLock { + #[inline(always)] pub fn ptr(&self) -> &Weak> { &self.ptr } + #[inline(always)] fn ptr_mut(&mut self) -> &mut Weak> { &mut self.ptr } + } + + impl PartialEq for ArcRwLock { + fn eq(&self, other: &Self) -> bool { self.ptr_eq(other) } + } + impl Eq for ArcRwLock {} + + impl Clone for WeakRwLock { + fn clone(&self) -> Self { Self { ptr: self.ptr.clone() } } + } + + impl PartialEq for WeakRwLock { + fn eq(&self, other: &Self) -> bool { self.ptr.ptr_eq(&other.ptr) } + } + impl Eq for WeakRwLock {} + + impl std::hash::Hash for ArcRwLock { + fn hash(&self, state: &mut H) { + let address = Arc::as_ptr(&self.ptr); + address.hash(state); + } + } + + impl std::hash::Hash for WeakRwLock { + fn hash(&self, state: &mut H) { + let address = Weak::as_ptr(&self.ptr); + address.hash(state); + } + } + + impl weak_table::traits::WeakElement for WeakRwLock { + type Strong = ArcRwLock; + fn new(view: &ArcRwLock) -> Self { view.downgrade() } + fn view(&self) -> Option> { self.upgrade() } + fn clone(view: &ArcRwLock) -> ArcRwLock { view.clone() } + } + + impl std::ops::Deref for ArcRwLock { + type Target = RwLock; + fn deref(&self) -> &Self::Target { &self.ptr } + } + + // --- IMPLEMENTATIONS: OrdArcRwLock (Ordered) --- + impl OrdArcRwLock { + pub fn downgrade(&self) -> OrdWeakRwLock { + OrdWeakRwLock:: { + ord: self.ord, + ptr: Arc::downgrade(&self.ptr), + } + } + } + + impl OrdWeakRwLock { + pub fn upgrade_force(&self) -> OrdArcRwLock { + OrdArcRwLock:: { + ord: self.ord, + ptr: self.ptr.upgrade().unwrap(), + } + } + pub fn upgrade(&self) -> Option> { + self.ptr.upgrade().map(|x| OrdArcRwLock:: { ord: self.ord, ptr: x }) + } + pub fn ptr_eq(&self, other: &Self) -> bool { + Weak::ptr_eq(&self.ptr, &other.ptr) + } + } + + impl Clone for OrdArcRwLock { + fn clone(&self) -> Self { + Self::new_ptr(Arc::clone(self.ptr()), self.ord) + } + } + + impl OrdArcRwLock { + pub fn new_ptr(ptr: Arc>, ord: U) -> Self { + Self { ord, ptr } + } + pub fn new_value(obj: T, ord: U) -> Self { + Self::new_ptr(Arc::new(RwLock::new(obj)), ord) + } + #[inline(always)] + pub fn ptr(&self) -> &Arc> { + &self.ptr + } + #[inline(always)] + pub fn ptr_mut(&mut self) -> &mut Arc> { + &mut self.ptr + } + #[inline(always)] + pub fn read_recursive(&self) -> RwLockReadGuard { + self.ptr.read_recursive() + } + #[inline(always)] + pub fn write(&self) -> RwLockWriteGuard { + self.ptr.write() + } + #[inline(always)] + pub fn get_ord(&self) -> U { + self.ord + } + } + + impl OrdWeakRwLock { + #[inline(always)] + pub fn ptr(&self) -> &Weak> { + &self.ptr + } + #[inline(always)] + fn ptr_mut(&mut self) -> &mut Weak> { + &mut self.ptr + } + } + + impl PartialEq for OrdArcRwLock { + fn eq(&self, other: &Self) -> bool { + self.ord.eq(&other.ord) + } + } + impl Eq for OrdArcRwLock {} + + impl Ord for OrdArcRwLock { + fn cmp(&self, other: &Self) -> std::cmp::Ordering { + self.ord.cmp(&other.ord) + } + } + + impl Ord for OrdWeakRwLock { + fn cmp(&self, other: &Self) -> std::cmp::Ordering { + self.ord.cmp(&other.ord) + } + } + impl Eq for OrdWeakRwLock {} + + impl PartialEq for OrdWeakRwLock { + fn eq(&self, other: &Self) -> bool { + self.ord.eq(&other.ord) + } + } + + impl PartialOrd for OrdArcRwLock { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } + } + + impl PartialOrd for OrdWeakRwLock { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } + } + + impl Clone for OrdWeakRwLock { + fn clone(&self) -> Self { + Self { + ord: self.ord, + ptr: self.ptr.clone(), + } + } + } + + impl std::ops::Deref for OrdArcRwLock { + type Target = RwLock; + fn deref(&self) -> &Self::Target { + &self.ptr + } + } + + impl std::hash::Hash for OrdArcRwLock { + fn hash(&self, state: &mut H) { + let address = Arc::as_ptr(&self.ptr); + (address, self.ord).hash(state); + } + } + + impl std::hash::Hash for OrdWeakRwLock { + fn hash(&self, state: &mut H) { + let address = Weak::as_ptr(&self.ptr); + (address, self.ord).hash(state); + } + } + + impl weak_table::traits::WeakElement for OrdWeakRwLock { + type Strong = OrdArcRwLock; + fn new(view: &Self::Strong) -> Self { view.downgrade() } + fn view(&self) -> Option { self.upgrade() } + fn clone(view: &Self::Strong) -> Self::Strong { view.clone() } + } + + // --- TYPE ALIASES --- + // Defaults to Ordered Safe Lock with a default tuple key + // pub type ArcManualSafeLock = ArcRwLock; + // pub type WeakManualSafeLock = WeakRwLock; + pub type ArcManualSafeLock = OrdArcRwLock; + pub type WeakManualSafeLock = OrdWeakRwLock; } } +// ========================================================================================= +// TESTS +// ========================================================================================= #[cfg(test)] mod tests { use super::*; @@ -122,8 +603,17 @@ mod tests { idx: usize, } - type TesterPtr = ArcRwLock; - type TesterWeak = WeakRwLock; + // Dynamic Type Alias for Testing + // NOTE: Both now point to their respective Ordered variants for consistency. + cfg_if::cfg_if! { + if #[cfg(feature="unsafe_pointer")] { + type TesterPtr = ArcManualSafeLock; + type TesterWeak = WeakManualSafeLock; + } else { + type TesterPtr = ArcManualSafeLock; + type TesterWeak = WeakManualSafeLock; + } + } impl std::fmt::Debug for TesterPtr { fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result { @@ -141,11 +631,22 @@ mod tests { #[test] fn pointers_test_1() { // cargo test pointers_test_1 -- --nocapture - let ptr = TesterPtr::new_value(Tester { idx: 0 }); + let ptr = TesterPtr::new_value(Tester { idx: 0 }, (0, 0)); let weak = ptr.downgrade(); - ptr.write().idx = 1; + + // Testing Write + { + ptr.write().idx = 1; + } + + // Testing Weak Upgrade and Read assert_eq!(weak.upgrade_force().read_recursive().idx, 1); - weak.upgrade_force().write().idx = 2; + + // Testing Write via Weak + { + weak.upgrade_force().write().idx = 2; + } + assert_eq!(ptr.read_recursive().idx, 2); } -} +} \ No newline at end of file diff --git a/src/primal_module.rs b/src/primal_module.rs index 9a4ffbc8..d5aa2de5 100644 --- a/src/primal_module.rs +++ b/src/primal_module.rs @@ -9,7 +9,7 @@ use std::sync::Arc; use crate::dual_module::*; use crate::itertools::Itertools; use crate::ordered_float::OrderedFloat; -use crate::primal_module_serial::ClusterAffinity; +use crate::primal_module_serial::{ClusterAffinity, PrimalClusterPtr, PrimalClusterWeak}; use crate::relaxer_optimizer::OptimizerResult; use crate::util::*; use crate::visualize::*; @@ -28,7 +28,7 @@ pub trait PrimalModuleImpl { fn clear(&mut self); /// load a new decoding problem given dual interface: note that all nodes MUST be defect node - fn load(&mut self, interface_ptr: &DualModuleInterfacePtr, dual_module: &mut D); + fn load(&mut self, interface_ptr: &DualModuleInterfacePtr, dual_module: &mut D, partition_id: usize); /// analyze the reason why dual module cannot further grow, update primal data structure (alternating tree, temporary matches, etc) /// and then tell dual module what to do to resolve these conflicts; @@ -69,8 +69,8 @@ pub trait PrimalModuleImpl { syndrome_pattern: Arc, dual_module: &mut impl DualModuleImpl, ) { - interface.load(syndrome_pattern, dual_module); - self.load(interface, dual_module); + interface.load(syndrome_pattern, dual_module, 0); // TODO: double check the partition id here + self.load(interface, dual_module, 0); // TODO: double check the partition id here self.solve_step_callback_interface_loaded(interface, dual_module, |_, _, _, _| {}) } @@ -161,7 +161,7 @@ pub trait PrimalModuleImpl { if moved_out_set.contains(to_flip) { moved_out_set.remove(to_flip); } else { - moved_out_set.insert(*to_flip); + moved_out_set.insert(to_flip.clone()); } } Arc::new(SyndromePattern::new_vertices(moved_out_set.into_iter().collect())) @@ -179,8 +179,8 @@ pub trait PrimalModuleImpl { // then call the solver to if let Some(visualizer) = visualizer { let callback = Self::visualizer_callback(visualizer); - interface.load(syndrome_pattern, dual_module); - self.load(interface, dual_module); + interface.load(syndrome_pattern, dual_module, 0); // TODO: double check the partition id here + self.load(interface, dual_module, 0); // TODO: double check the partition id here self.solve_step_callback_interface_loaded(interface, dual_module, callback); visualizer .snapshot_combined("solved".to_string(), vec![interface, dual_module, self]) @@ -230,10 +230,10 @@ pub trait PrimalModuleImpl { let cluster_affs = self.get_sorted_clusters_aff(); for cluster_affinity in cluster_affs.into_iter().sorted() { - let cluster_index = cluster_affinity.cluster_index; + let cluster_ptr = cluster_affinity.cluster_ptr; let mut dual_node_deltas = FastIterMap::new(); let (mut resolved, optimizer_result) = - self.resolve_cluster_tune(cluster_index, interface, dual_module, &mut dual_node_deltas); + self.resolve_cluster_tune(cluster_ptr, interface, dual_module, &mut dual_node_deltas); let mut obstacles = dual_module.get_obstacles_tune(optimizer_result, dual_node_deltas); @@ -322,9 +322,10 @@ pub trait PrimalModuleImpl { dual_module: &mut impl DualModuleImpl, ) -> (OutputSubgraph, WeightRange) { let output_subgraph = self.subgraph(interface, dual_module); + let internal_subgraph = OutputSubgraph::get_internal_subgraph(&output_subgraph); let weight_range = WeightRange::new( interface.sum_dual_variables() + dual_module.get_negative_weight_sum(), - dual_module.get_subgraph_weight(&output_subgraph.subgraph) + dual_module.get_negative_weight_sum(), + dual_module.get_subgraph_weight(internal_subgraph) + dual_module.get_negative_weight_sum(), ); (output_subgraph, weight_range) } @@ -341,14 +342,14 @@ pub trait PrimalModuleImpl { } /// in "tune" mode, return the list of clusters that need to be resolved - fn pending_clusters(&mut self) -> Vec { + fn pending_clusters(&mut self) -> Vec { panic!("not implemented `pending_clusters`"); } /// check if a cluster has been solved, if not then resolve it fn resolve_cluster( &mut self, - _cluster_index: NodeIndex, + _cluster_ptr: PrimalClusterPtr, _interface_ptr: &DualModuleInterfacePtr, _dual_module: &mut impl DualModuleImpl, ) -> bool { @@ -358,11 +359,11 @@ pub trait PrimalModuleImpl { /// `resolve_cluster` but in tuning mode, optimizer result denotes what the optimizer has accomplished fn resolve_cluster_tune( &mut self, - _cluster_index: NodeIndex, + _cluster_ptr: PrimalClusterPtr, _interface_ptr: &DualModuleInterfacePtr, _dual_module: &mut impl DualModuleImpl, // _dual_node_deltas: &mut FastIterMap, - _dual_node_deltas: &mut FastIterMap, + _dual_node_deltas: &mut FastIterMap, ) -> (bool, OptimizerResult) { panic!("not implemented `resolve_cluster_tune`"); } diff --git a/src/primal_module_serial.rs b/src/primal_module_serial.rs index 13bd4a4c..5975f1b4 100644 --- a/src/primal_module_serial.rs +++ b/src/primal_module_serial.rs @@ -14,6 +14,7 @@ use crate::primal_module::*; use crate::relaxer_optimizer::*; use crate::util::*; use crate::visualize::*; +use crate::dual_module_pq::{EdgePtr, VertexPtr, EdgeWeak}; use std::collections::VecDeque; use std::fmt::Debug; @@ -38,7 +39,7 @@ pub struct PrimalModuleSerial { pub plugins: Arc, /// how many plugins are actually executed for every cluster pub plugin_count: Arc>, - pub plugin_pending_clusters: Vec, + pub plugin_pending_clusters: Vec, /// configuration pub config: PrimalModuleSerialConfig, /// the time spent on resolving the obstacles @@ -50,22 +51,28 @@ pub struct PrimalModuleSerial { pub cluster_weights_initialized: bool, } -#[derive(Eq, Debug, Clone, Default)] +#[derive(Eq, Clone)] pub struct ClusterAffinity { - pub cluster_index: NodeIndex, + pub cluster_ptr: PrimalClusterPtr, pub affinity: Affinity, } -impl std::hash::Hash for ClusterAffinity { - fn hash(&self, state: &mut H) { - self.cluster_index.hash(state); - self.affinity.hash(state); +impl std::fmt::Debug for ClusterAffinity { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // Assuming your generic pointer has a way to get the index (like .1 or .index) + // If PrimalClusterPtr is OrdArcRwLock, the index is usually in the second tuple element. + let cluster_index = self.cluster_ptr.read_recursive().cluster_index; // Or however you access the index in your lock wrapper + + f.debug_struct("ClusterAffinity") + .field("cluster_ptr", &cluster_index) // Only print the ID, not the whole object + // .field("other_field", &self.other_field) + .finish() } } impl PartialEq for ClusterAffinity { fn eq(&self, other: &Self) -> bool { - self.affinity == other.affinity && self.cluster_index == other.cluster_index + self.affinity == other.affinity && self.cluster_ptr.eq(&other.cluster_ptr) } } @@ -76,9 +83,9 @@ impl Ord for ClusterAffinity { match other.affinity.cmp(&self.affinity) { std::cmp::Ordering::Equal => { // If affinities are equal, compare cluster_index in ascending order - self.cluster_index.cmp(&other.cluster_index) + self.cluster_ptr.read_recursive().cluster_index.cmp(&other.cluster_ptr.read_recursive().cluster_index) } - other => other.reverse(), + other => other, } } } @@ -89,6 +96,12 @@ impl PartialOrd for ClusterAffinity { } } +impl std::hash::Hash for ClusterAffinity { + fn hash(&self, state: &mut H) { + (self.cluster_ptr.clone(), self.affinity.clone()).hash(state); + } +} + pub enum Unionable { Can, DoesNotNeed, @@ -131,8 +144,8 @@ pub struct PrimalModuleSerialNode { pub cluster_weak: PrimalClusterWeak, } -pub type PrimalModuleSerialNodePtr = ArcRwLock; -pub type PrimalModuleSerialNodeWeak = WeakRwLock; +pub type PrimalModuleSerialNodePtr = ArcManualSafeLock; +pub type PrimalModuleSerialNodeWeak = WeakManualSafeLock; pub struct PrimalCluster { /// the index in the cluster @@ -140,13 +153,13 @@ pub struct PrimalCluster { /// the nodes that belongs to this cluster pub nodes: Vec, /// all the edges ever exists in any hair - pub edges: FastIterSet, + pub edges: FastIterSet, /// all the vertices ever touched by any tight edge - pub vertices: FastIterSet, + pub vertices: FastIterSet, /// the parity matrix to determine whether it's a valid cluster and also find new ways to increase the dual pub matrix: EchelonMatrix, /// the parity subgraph result, only valid when it's solved - pub subgraph: Option, + pub subgraph: Option, /// plugin manager helps to execute the plugin and find an executable relaxer pub plugin_manager: PluginManager, /// optimizing the direction of relaxers @@ -154,10 +167,12 @@ pub struct PrimalCluster { /// HIHGS solution stored for incrmental lp #[cfg(feature = "incr_lp")] //note: really depends where we want the error to manifest pub incr_solution: Option>>, + /// the partition id of the cluster + pub partition_id: usize, } -pub type PrimalClusterPtr = ArcRwLock; -pub type PrimalClusterWeak = WeakRwLock; +pub type PrimalClusterPtr = ArcManualSafeLock; +pub type PrimalClusterWeak = WeakManualSafeLock; impl PrimalModuleImpl for PrimalModuleSerial { fn new_empty(_initializer: &Arc) -> Self { @@ -189,7 +204,7 @@ impl PrimalModuleImpl for PrimalModuleSerial { } #[allow(clippy::unnecessary_cast)] - fn load(&mut self, interface_ptr: &DualModuleInterfacePtr, _dual_module: &mut D) { + fn load(&mut self, interface_ptr: &DualModuleInterfacePtr, _dual_module: &mut D, partition_id: usize) { let interface = interface_ptr.read_recursive(); for index in 0..interface.nodes.len() as NodeIndex { let dual_node_ptr = &interface.nodes[index as usize]; @@ -217,19 +232,23 @@ impl PrimalModuleImpl for PrimalModuleSerial { nodes: vec![], edges: node.invalid_subgraph.hair.clone(), vertices: node.invalid_subgraph.vertices.clone(), - matrix: node.invalid_subgraph.generate_matrix(&interface.decoding_graph), + matrix: node.invalid_subgraph.generate_matrix(), subgraph: None, plugin_manager: PluginManager::new(self.plugins.clone(), self.plugin_count.clone()), relaxer_optimizer: RelaxerOptimizer::new(), #[cfg(all(feature = "incr_lp", feature = "highs"))] incr_solution: None, - }); + partition_id, + }, (partition_id, self.clusters.len() as NodeIndex)); // create the primal node of this defect node and insert into cluster let primal_node_ptr = PrimalModuleSerialNodePtr::new_value(PrimalModuleSerialNode { dual_node_ptr: dual_node_ptr.clone(), cluster_weak: primal_cluster_ptr.downgrade(), - }); + }, + (partition_id, node.index as usize)); + drop(node); primal_cluster_ptr.write().nodes.push(primal_node_ptr.clone()); + dual_node_ptr.write().primal_module_serial_node = Some(primal_node_ptr.clone().downgrade()); // add to self self.nodes.push(primal_node_ptr); self.clusters.push(primal_cluster_ptr); @@ -289,10 +308,15 @@ impl PrimalModuleImpl for PrimalModuleSerial { cluster.cluster_index, cluster.vertices, cluster.edges ) }) - .iter(), + , ); } - OutputSubgraph::new(subgraph, _dual_module.get_negative_edges()) + let subgraph_index: Vec = subgraph + .clone() + .into_iter() + .map(|edge_weak| edge_weak.upgrade_force().read_recursive().edge_index) + .collect(); + OutputSubgraph::new(subgraph_index, _dual_module.get_negative_edges(), subgraph) } /// check if there are more plugins to be applied @@ -304,7 +328,7 @@ impl PrimalModuleImpl for PrimalModuleSerial { return if *self.plugin_count.read_recursive() < self.plugins.len() { // increment the plugin count *self.plugin_count.write() += 1; - self.plugin_pending_clusters = (0..self.clusters.len()).collect(); + self.plugin_pending_clusters = self.clusters.iter().map(|c| c.downgrade()).collect::>(); true } else { false @@ -312,7 +336,7 @@ impl PrimalModuleImpl for PrimalModuleSerial { } /// get the pending clusters - fn pending_clusters(&mut self) -> Vec { + fn pending_clusters(&mut self) -> Vec { self.plugin_pending_clusters.clone() } @@ -322,11 +346,10 @@ impl PrimalModuleImpl for PrimalModuleSerial { #[allow(clippy::unnecessary_cast)] fn resolve_cluster( &mut self, - cluster_index: NodeIndex, + cluster_ptr: PrimalClusterPtr, interface_ptr: &DualModuleInterfacePtr, dual_module: &mut impl DualModuleImpl, ) -> bool { - let cluster_ptr = self.clusters[cluster_index as usize].clone(); let mut cluster = cluster_ptr.write(); if cluster.nodes.is_empty() { return true; // no longer a cluster, no need to handle @@ -338,10 +361,10 @@ impl PrimalModuleImpl for PrimalModuleSerial { } // update the matrix with new tight edges let cluster = &mut *cluster; - for &edge_index in cluster.edges.iter() { + for edge_ptr in cluster.edges.iter() { cluster .matrix - .update_edge_tightness(edge_index, dual_module.is_edge_tight(edge_index)); + .update_edge_tightness(edge_ptr.downgrade(), dual_module.is_edge_tight(edge_ptr.clone())); } // find an executable relaxer from the plugin manager @@ -356,21 +379,22 @@ impl PrimalModuleImpl for PrimalModuleSerial { let cluster_mut = &mut *cluster; // must first get mutable reference let plugin_manager = &mut cluster_mut.plugin_manager; let matrix = &mut cluster_mut.matrix; - plugin_manager.find_relaxer(decoding_graph, matrix, &positive_dual_variables) + plugin_manager.find_relaxer(dual_module, matrix, &positive_dual_variables) }; // if a relaxer is found, execute it and return if let Some(relaxer) = relaxer { for (invalid_subgraph, grow_rate) in relaxer.get_direction().iter() { - let (existing, dual_node_ptr) = interface_ptr.find_or_create_node(invalid_subgraph, dual_module); + let (existing, dual_node_ptr) = interface_ptr.find_or_create_node(invalid_subgraph, dual_module, cluster.partition_id); if !existing { // create the corresponding primal node and add it to cluster let primal_node_ptr = PrimalModuleSerialNodePtr::new_value(PrimalModuleSerialNode { dual_node_ptr: dual_node_ptr.clone(), cluster_weak: cluster_ptr.downgrade(), - }); + }, (cluster.partition_id, dual_node_ptr.read_recursive().index as usize)); cluster.nodes.push(primal_node_ptr.clone()); - self.nodes.push(primal_node_ptr); + self.nodes.push(primal_node_ptr.clone()); + dual_node_ptr.write().primal_module_serial_node = Some(primal_node_ptr.downgrade()); } dual_module.set_grow_rate(&dual_node_ptr, grow_rate.clone()); @@ -383,7 +407,7 @@ impl PrimalModuleImpl for PrimalModuleSerial { // subgraph with minimum weight from all plugins as the starting point to do local minimum // find a local minimum (hopefully a global minimum) - let weight_of = |edge_index: EdgeIndex| dual_module.get_edge_weight(edge_index); + let weight_of = |edge_weak: EdgeWeak| dual_module.get_edge_weight(edge_weak); cluster.subgraph = Some(cluster.matrix.get_solution_local_minimum(weight_of).expect("satisfiable")); true } @@ -392,17 +416,16 @@ impl PrimalModuleImpl for PrimalModuleSerial { #[allow(clippy::unnecessary_cast)] fn resolve_cluster_tune( &mut self, - cluster_index: NodeIndex, + cluster_ptr: PrimalClusterPtr, interface_ptr: &DualModuleInterfacePtr, dual_module: &mut impl DualModuleImpl, // dual_node_deltas: &mut FastIterMap, - dual_node_deltas: &mut FastIterMap, + dual_node_deltas: &mut FastIterMap, ) -> (bool, OptimizerResult) { let mut optimizer_result = OptimizerResult::default(); #[cfg(feature = "incr_lp")] - let mut cluster_ptr = self.clusters[cluster_index as usize].clone(); + let mut cluster_ptr = cluster_ptr.clone(); #[cfg(not(feature = "incr_lp"))] - let cluster_ptr = self.clusters[cluster_index as usize].clone(); let mut cluster_temp = cluster_ptr.write(); if cluster_temp.nodes.is_empty() { return (true, optimizer_result); // no longer a cluster, no need to handle @@ -416,10 +439,10 @@ impl PrimalModuleImpl for PrimalModuleSerial { #[cfg(not(feature = "incr_lp"))] let cluster = &mut *cluster_temp; - for &edge_index in cluster.edges.iter() { + for edge_ptr in cluster.edges.iter() { cluster .matrix - .update_edge_tightness(edge_index, dual_module.is_edge_tight_tune(edge_index)); + .update_edge_tightness(edge_ptr.downgrade(), dual_module.is_edge_tight_tune(edge_ptr.clone())); } // find an executable relaxer from the plugin manager @@ -430,11 +453,10 @@ impl PrimalModuleImpl for PrimalModuleSerial { .map(|p| p.read_recursive().dual_node_ptr.clone()) .filter(|dual_node_ptr| !dual_node_ptr.read_recursive().dual_variable_at_last_updated_time.is_zero()) .collect(); - let decoding_graph = &interface_ptr.read_recursive().decoding_graph; let cluster_mut = &mut *cluster; // must first get mutable reference let plugin_manager = &mut cluster_mut.plugin_manager; let matrix = &mut cluster_mut.matrix; - plugin_manager.find_relaxer(decoding_graph, matrix, &positive_dual_variables) + plugin_manager.find_relaxer(dual_module, matrix, &positive_dual_variables) }; // Yue added 2025.1.31: also check for local minimum during the algorithm; otherwise when we increase @@ -442,7 +464,7 @@ impl PrimalModuleImpl for PrimalModuleSerial { // more complicated dual solution does not necessarily mean better logical error rate. Rather, if we // keep looking for smaller weighted solutions in the middle, the result is hopefully better. if !self.config.only_solve_primal_once { - let weight_of = |edge_index: EdgeIndex| dual_module.get_edge_weight(edge_index); + let weight_of = |edge_weak: EdgeWeak| edge_weak.upgrade_force().read_recursive().weight.clone(); if let Some(subgraph) = cluster.matrix.get_solution_local_minimum(weight_of) { if let Some(original_subgraph) = &cluster.subgraph { let original_weight = dual_module.get_subgraph_weight(original_subgraph); @@ -475,7 +497,7 @@ impl PrimalModuleImpl for PrimalModuleSerial { ) }) .collect(); - let edge_slacks: FastIterMap = dual_variables + let edge_slacks: FastIterMap = dual_variables .keys() .flat_map(|invalid_subgraph: &Arc| invalid_subgraph.hair.iter().cloned()) .chain( @@ -485,7 +507,7 @@ impl PrimalModuleImpl for PrimalModuleSerial { .flat_map(|invalid_subgraph| invalid_subgraph.hair.iter().cloned()), ) .unique() - .map(|edge_index| (edge_index, dual_module.get_edge_slack_tune(edge_index))) + .map(|edge_ptr| (edge_ptr.clone(), dual_module.get_edge_slack_tune(edge_ptr.clone()))) .collect(); let (new_relaxer, early_returned) = cluster.relaxer_optimizer.optimize(relaxer, edge_slacks, dual_variables); @@ -627,23 +649,24 @@ impl PrimalModuleImpl for PrimalModuleSerial { for (invalid_subgraph, grow_rate) in relaxer.get_direction().iter() { if let Some((existing, dual_node_ptr)) = - interface_ptr.find_or_create_node_tune(invalid_subgraph, dual_module) + interface_ptr.find_or_create_node_tune(invalid_subgraph, dual_module, cluster.partition_id) { if !existing { // create the corresponding primal node and add it to cluster let primal_node_ptr = PrimalModuleSerialNodePtr::new_value(PrimalModuleSerialNode { dual_node_ptr: dual_node_ptr.clone(), cluster_weak: cluster_ptr.downgrade(), - }); + }, (cluster.partition_id, dual_node_ptr.read_recursive().index as usize)); cluster.nodes.push(primal_node_ptr.clone()); - self.nodes.push(primal_node_ptr); + self.nodes.push(primal_node_ptr.clone()); + dual_node_ptr.write().primal_module_serial_node = Some(primal_node_ptr.downgrade()); } // Document the desired deltas let index = dual_node_ptr.read_recursive().index; dual_node_deltas.insert( OrderedDualNodePtr::new(index, dual_node_ptr), - (grow_rate.clone(), cluster_index), + (grow_rate.clone(), cluster_ptr.clone()), ); } } @@ -654,7 +677,7 @@ impl PrimalModuleImpl for PrimalModuleSerial { // find a local minimum (hopefully a global minimum) if self.config.only_solve_primal_once { - let weight_of = |edge_index: EdgeIndex| dual_module.get_edge_weight(edge_index); + let weight_of = |edge_weak: EdgeWeak| dual_module.get_edge_weight(edge_weak); cluster.subgraph = Some(cluster.matrix.get_solution_local_minimum(weight_of).expect("satisfiable")); } @@ -666,12 +689,12 @@ impl PrimalModuleImpl for PrimalModuleSerial { let pending_clusters = self.pending_clusters(); let mut sorted_clusters_aff = FastIterSet::default(); - for cluster_index in pending_clusters.iter() { - let cluster_ptr = self.clusters[*cluster_index].clone(); - let affinity = dual_module.calculate_cluster_affinity(cluster_ptr); + for cluster_weak in pending_clusters.iter() { + let cluster_ptr = cluster_weak.upgrade_force(); + let affinity = dual_module.calculate_cluster_affinity(cluster_ptr.clone()); if let Some(affinity) = affinity { sorted_clusters_aff.insert(ClusterAffinity { - cluster_index: *cluster_index, + cluster_ptr: cluster_ptr.clone(), affinity, }); } @@ -721,21 +744,24 @@ impl PrimalModuleSerial { &self, dual_node_ptr_1: &DualNodePtr, dual_node_ptr_2: &DualNodePtr, - decoding_graph: &DecodingHyperGraph, _dual_module: &mut impl DualModuleImpl, // note: remove if not for cluster-based ) { // cluster_1 will become the union of cluster_1 and cluster_2 // and cluster_2 will be outdated - let node_index_1 = dual_node_ptr_1.read_recursive().index; - let node_index_2 = dual_node_ptr_2.read_recursive().index; - if node_index_1 == node_index_2 { + if dual_node_ptr_1.eq(dual_node_ptr_2) { return; // already the same node } - let primal_node_1 = self.nodes[node_index_1 as usize].read_recursive(); - let primal_node_2 = self.nodes[node_index_2 as usize].read_recursive(); - if primal_node_1.cluster_weak.ptr_eq(&primal_node_2.cluster_weak) { + let primal_node_1_weak = dual_node_ptr_1.read_recursive().primal_module_serial_node.clone().unwrap(); + let primal_node_2_weak = dual_node_ptr_2.read_recursive().primal_module_serial_node.clone().unwrap(); + let primal_node_1_ptr = primal_node_1_weak.upgrade_force(); + let primal_node_2_ptr = primal_node_2_weak.upgrade_force(); + let primal_node_1 = primal_node_1_ptr.read_recursive(); + let primal_node_2 = primal_node_2_ptr.read_recursive(); + + if primal_node_1.cluster_weak.eq(&primal_node_2.cluster_weak) { return; // already in the same cluster } + let cluster_ptr_1 = primal_node_1.cluster_weak.upgrade_force(); let cluster_ptr_2 = primal_node_2.cluster_weak.upgrade_force(); drop(primal_node_1); @@ -844,12 +870,13 @@ impl PrimalModuleSerial { // cluster_1.subgraph = None; // mark as no subgraph } - for &vertex_index in cluster_2.vertices.iter() { - if !cluster_1.vertices.contains(&vertex_index) { - cluster_1.vertices.insert(vertex_index); - let incident_edges = decoding_graph.get_vertex_neighbors(vertex_index); - let parity = decoding_graph.is_vertex_defect(vertex_index); - cluster_1.matrix.add_constraint(vertex_index, incident_edges, parity); + for vertex_ptr in cluster_2.vertices.iter() { + if !cluster_1.vertices.contains(&vertex_ptr) { + cluster_1.vertices.insert(vertex_ptr.clone()); + let vertex = vertex_ptr.read_recursive(); + let incident_edges = &vertex.edges; + let parity = vertex.is_defect; + cluster_1.matrix.add_constraint(vertex_ptr.downgrade().clone(), incident_edges, parity); } } cluster_1.relaxer_optimizer.append(&mut cluster_2.relaxer_optimizer); @@ -864,14 +891,13 @@ impl PrimalModuleSerial { dual_module: &mut impl DualModuleImpl, ) -> bool { debug_assert!(!dual_report.is_unbounded() && dual_report.get_valid_growth().is_none()); - let mut active_clusters = FastIterSet::::new(); + let mut active_clusters = FastIterSet::::new(); let interface = interface_ptr.read_recursive(); - let decoding_graph = &interface.decoding_graph; while let Some(obstacle) = dual_report.pop() { match obstacle { - Obstacle::Conflict { edge_index } => { + Obstacle::Conflict { edge_ptr } => { // union all the dual nodes in the edge index and create new dual node by adding this edge to `internal_edges` - let dual_nodes = dual_module.get_edge_nodes(edge_index); + let dual_nodes = dual_module.get_edge_nodes(edge_ptr.clone()); debug_assert!( !dual_nodes.is_empty(), "should not conflict if no dual nodes are contributing" @@ -880,34 +906,30 @@ impl PrimalModuleSerial { // first union all the dual nodes for dual_node_ptr in dual_nodes.iter().skip(1) { // self.union(dual_node_ptr_0, dual_node_ptr, &interface.decoding_graph); - self.union(dual_node_ptr_0, dual_node_ptr, &interface.decoding_graph, dual_module); + self.union(dual_node_ptr_0, dual_node_ptr, dual_module); } - let cluster_ptr = self.nodes[dual_node_ptr_0.read_recursive().index as usize] - .read_recursive() - .cluster_weak - .upgrade_force(); + let primal_node_weak = dual_node_ptr_0.read_recursive().primal_module_serial_node.clone().unwrap(); + let cluster_ptr = primal_node_weak.upgrade_force().read_recursive().cluster_weak.upgrade_force(); let mut cluster = cluster_ptr.write(); // then add new constraints because these edges may touch new vertices - let incident_vertices = decoding_graph.get_edge_neighbors(edge_index); - for &vertex_index in incident_vertices.iter() { - if !cluster.vertices.contains(&vertex_index) { - cluster.vertices.insert(vertex_index); - let incident_edges = decoding_graph.get_vertex_neighbors(vertex_index); - let parity = decoding_graph.is_vertex_defect(vertex_index); - cluster.matrix.add_constraint(vertex_index, incident_edges, parity); + let incident_vertices = &edge_ptr.read_recursive().vertices; + for vertex_weak in incident_vertices.iter() { + let vertex_ptr = vertex_weak.upgrade_force(); + if !cluster.vertices.contains(&vertex_ptr) { + cluster.vertices.insert(vertex_ptr.clone()); + let incident_edges = &vertex_ptr.read_recursive().edges; + let parity = vertex_ptr.read_recursive().is_defect; + cluster.matrix.add_constraint(vertex_ptr.downgrade().clone(), incident_edges, parity); } } - cluster.edges.insert(edge_index); + cluster.edges.insert(edge_ptr.clone()); // add to active cluster so that it's processed later - active_clusters.insert(cluster.cluster_index); + active_clusters.insert(cluster_ptr.clone()); } Obstacle::ShrinkToZero { dual_node_ptr } => { - let cluster_ptr = self.nodes[dual_node_ptr.index as usize] - .read_recursive() - .cluster_weak - .upgrade_force(); - let cluster_index = cluster_ptr.read_recursive().cluster_index; - active_clusters.insert(cluster_index); + let primal_node_weak = dual_node_ptr.ptr.read_recursive().primal_module_serial_node.clone().unwrap(); + let cluster_ptr = primal_node_weak.upgrade_force().read_recursive().cluster_weak.upgrade_force(); + active_clusters.insert(cluster_ptr.clone()); } } } @@ -916,8 +938,8 @@ impl PrimalModuleSerial { *self.plugin_count.write() = 0; // force only the first plugin } let mut all_solved = true; - for &cluster_index in active_clusters.iter() { - let solved = self.resolve_cluster(cluster_index, interface_ptr, dual_module); + for cluster_ptr in active_clusters.iter() { + let solved = self.resolve_cluster(cluster_ptr.clone(), interface_ptr, dual_module); all_solved &= solved; } if !all_solved { @@ -936,14 +958,13 @@ impl PrimalModuleSerial { dual_module: &mut impl DualModuleImpl, ) -> bool { debug_assert!(!dual_report.is_unbounded() && dual_report.get_valid_growth().is_none()); - let mut active_clusters = FastIterSet::::new(); + let mut active_clusters = FastIterSet::::new(); let interface = interface_ptr.read_recursive(); - let decoding_graph = &interface.decoding_graph; while let Some(obstacle) = dual_report.pop() { match obstacle { - Obstacle::Conflict { edge_index } => { + Obstacle::Conflict { edge_ptr } => { // union all the dual nodes in the edge index and create new dual node by adding this edge to `internal_edges` - let dual_nodes = dual_module.get_edge_nodes(edge_index); + let dual_nodes = dual_module.get_edge_nodes(edge_ptr.clone()); debug_assert!( !dual_nodes.is_empty(), "should not conflict if no dual nodes are contributing" @@ -952,34 +973,31 @@ impl PrimalModuleSerial { // first union all the dual nodes for dual_node_ptr in dual_nodes.iter().skip(1) { // self.union(dual_node_ptr_0, dual_node_ptr, &interface.decoding_graph); - self.union(dual_node_ptr_0, dual_node_ptr, &interface.decoding_graph, dual_module); + self.union(dual_node_ptr_0, dual_node_ptr, dual_module); } - let cluster_ptr = self.nodes[dual_node_ptr_0.read_recursive().index as usize] - .read_recursive() - .cluster_weak - .upgrade_force(); + let primal_node_weak = dual_node_ptr_0.read_recursive().primal_module_serial_node.clone().unwrap(); + let cluster_ptr = primal_node_weak.upgrade_force().read_recursive().cluster_weak.upgrade_force(); let mut cluster = cluster_ptr.write(); // then add new constraints because these edges may touch new vertices - let incident_vertices = decoding_graph.get_edge_neighbors(edge_index); - for &vertex_index in incident_vertices.iter() { - if !cluster.vertices.contains(&vertex_index) { - cluster.vertices.insert(vertex_index); - let incident_edges = decoding_graph.get_vertex_neighbors(vertex_index); - let parity = decoding_graph.is_vertex_defect(vertex_index); - cluster.matrix.add_constraint(vertex_index, incident_edges, parity); + let incident_vertices = &edge_ptr.read_recursive().vertices; + for vertex_weak in incident_vertices.iter() { + let vertex_ptr = vertex_weak.upgrade_force(); + if !cluster.vertices.contains(&vertex_ptr) { + let vertex = vertex_ptr.read_recursive(); + cluster.vertices.insert(vertex_ptr.clone()); + let incident_edges = &vertex.edges; + let parity = vertex.is_defect; + cluster.matrix.add_constraint(vertex_weak.clone(), incident_edges, parity); } } - cluster.edges.insert(edge_index); + cluster.edges.insert(edge_ptr.clone()); // add to active cluster so that it's processed later - active_clusters.insert(cluster.cluster_index); + active_clusters.insert(cluster_ptr.clone()); } Obstacle::ShrinkToZero { dual_node_ptr } => { - let cluster_ptr = self.nodes[dual_node_ptr.index as usize] - .read_recursive() - .cluster_weak - .upgrade_force(); - let cluster_index = cluster_ptr.read_recursive().cluster_index; - active_clusters.insert(cluster_index); + let primal_node_weak = dual_node_ptr.ptr.read_recursive().primal_module_serial_node.clone().unwrap(); + let cluster_ptr = primal_node_weak.upgrade_force().read_recursive().cluster_weak.upgrade_force(); + active_clusters.insert(cluster_ptr.clone()); } } } @@ -988,8 +1006,8 @@ impl PrimalModuleSerial { *self.plugin_count.write() = 0; // force only the first plugin } let mut all_solved = true; - for &cluster_index in active_clusters.iter() { - let solved = self.resolve_cluster(cluster_index, interface_ptr, dual_module); + for cluster_ptr in active_clusters.iter() { + let solved = self.resolve_cluster(cluster_ptr.clone(), interface_ptr, dual_module); all_solved &= solved; } if !all_solved { @@ -1010,8 +1028,8 @@ impl PrimalModuleSerial { } // check that all clusters have passed the plugins loop { - while let Some(cluster_index) = self.plugin_pending_clusters.pop() { - let solved = self.resolve_cluster(cluster_index, interface_ptr, dual_module); + while let Some(cluster_weak) = self.plugin_pending_clusters.pop() { + let solved = self.resolve_cluster(cluster_weak.upgrade_force(), interface_ptr, dual_module); if !solved { return false; // let the dual module to handle one } @@ -1019,7 +1037,7 @@ impl PrimalModuleSerial { if *self.plugin_count.read_recursive() < self.plugins.len() { // increment the plugin count *self.plugin_count.write() += 1; - self.plugin_pending_clusters = (0..self.clusters.len()).collect(); + self.plugin_pending_clusters = self.clusters.iter().map(|c| c.downgrade()).collect::>(); } else { break; // nothing more to check } @@ -1035,15 +1053,14 @@ impl PrimalModuleSerial { interface_ptr: &DualModuleInterfacePtr, dual_module: &mut impl DualModuleImpl, ) -> (FastIterSet, bool) { - let mut active_clusters = FastIterSet::::new(); + let mut active_clusters = FastIterSet::::new(); let interface = interface_ptr.read_recursive(); - let decoding_graph = &interface.decoding_graph; for obstacle in dual_report.into_iter() { match obstacle { - Obstacle::Conflict { edge_index } => { + Obstacle::Conflict { edge_ptr } => { // union all the dual nodes in the edge index and create new dual node by adding this edge to `internal_edges` - let dual_nodes = dual_module.get_edge_nodes(edge_index); + let dual_nodes = dual_module.get_edge_nodes(edge_ptr.clone()); debug_assert!( !dual_nodes.is_empty(), "should not conflict if no dual nodes are contributing" @@ -1053,34 +1070,31 @@ impl PrimalModuleSerial { // first union all the dual nodes for dual_node_ptr in dual_nodes.iter().skip(1) { // self.union(dual_node_ptr_0, dual_node_ptr, &interface.decoding_graph); - self.union(dual_node_ptr_0, dual_node_ptr, &interface.decoding_graph, dual_module); + self.union(dual_node_ptr_0, dual_node_ptr, dual_module); } - let cluster_ptr = self.nodes[dual_node_ptr_0.read_recursive().index as usize] - .read_recursive() - .cluster_weak - .upgrade_force(); + // TODO: Double check this + let primal_node_weak = dual_node_ptr_0.read_recursive().primal_module_serial_node.clone().unwrap(); + let cluster_ptr = primal_node_weak.upgrade_force().read_recursive().cluster_weak.upgrade_force(); let mut cluster = cluster_ptr.write(); // then add new constraints because these edges may touch new vertices - let incident_vertices = decoding_graph.get_edge_neighbors(edge_index); - for &vertex_index in incident_vertices.iter() { - if !cluster.vertices.contains(&vertex_index) { - cluster.vertices.insert(vertex_index); - let incident_edges = decoding_graph.get_vertex_neighbors(vertex_index); - let parity = decoding_graph.is_vertex_defect(vertex_index); - cluster.matrix.add_constraint(vertex_index, incident_edges, parity); + let incident_vertices = &edge_ptr.read_recursive().vertices; + for vertex_weak in incident_vertices.iter() { + let vertex_ptr = vertex_weak.upgrade_force(); + if !cluster.vertices.contains(&vertex_ptr) { + cluster.vertices.insert(vertex_ptr.clone()); + let incident_edges = &vertex_ptr.read_recursive().edges; + let parity = vertex_ptr.read_recursive().is_defect; + cluster.matrix.add_constraint(vertex_weak.clone(), incident_edges, parity); } } - cluster.edges.insert(edge_index); + cluster.edges.insert(edge_ptr.clone()); // add to active cluster so that it's processed later - active_clusters.insert(cluster.cluster_index); + active_clusters.insert(cluster_ptr.clone()); } Obstacle::ShrinkToZero { dual_node_ptr } => { - let cluster_ptr = self.nodes[dual_node_ptr.index as usize] - .read_recursive() - .cluster_weak - .upgrade_force(); - let cluster_index = cluster_ptr.read_recursive().cluster_index; - active_clusters.insert(cluster_index); + let primal_node_weak = dual_node_ptr.ptr.read_recursive().primal_module_serial_node.clone().unwrap(); + let cluster_ptr = primal_node_weak.upgrade_force().read_recursive().cluster_weak.upgrade_force(); + active_clusters.insert(cluster_ptr.clone()); } } } @@ -1092,9 +1106,9 @@ impl PrimalModuleSerial { let mut all_solved = true; let mut dual_node_deltas = FastIterMap::new(); let mut optimizer_result = OptimizerResult::default(); - for &cluster_index in active_clusters.iter() { + for cluster_ptr in active_clusters.iter() { let (solved, other) = - self.resolve_cluster_tune(cluster_index, interface_ptr, dual_module, &mut dual_node_deltas); + self.resolve_cluster_tune(cluster_ptr.clone(), interface_ptr, dual_module, &mut dual_node_deltas); if !solved { // todo: investigate more return (dual_module.get_obstacles_tune(other, dual_node_deltas), false); @@ -1190,7 +1204,7 @@ pub mod tests { // primal_module.config = serde_json::from_value(json!({"timeout":1})).unwrap(); // try to work on a simple syndrome let decoding_graph = DecodingHyperGraph::new_defects(model_graph, defect_vertices.clone()); - let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone()); + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); primal_module.solve_visualizer( &interface_ptr, decoding_graph.syndrome_pattern.clone(), @@ -1251,7 +1265,7 @@ pub mod tests { defect_vertices, final_dual, plugins, - DualModulePQ::new_empty(&model_graph.initializer), + DualModulePQ::new_empty(&model_graph.initializer, 0), model_graph, Some(visualizer), ) diff --git a/src/primal_module_union_find.rs b/src/primal_module_union_find.rs index 78cf67ee..afc12d83 100644 --- a/src/primal_module_union_find.rs +++ b/src/primal_module_union_find.rs @@ -12,6 +12,8 @@ use crate::invalid_subgraph::*; use crate::num_traits::Zero; use crate::pointers::*; use crate::primal_module::*; +use crate::dual_module::{DualNodePtr}; +use crate::dual_module_pq::EdgePtr; use crate::union_find::*; use crate::util::*; use crate::visualize::*; @@ -23,17 +25,21 @@ use std::sync::Arc; pub struct PrimalModuleUnionFind { /// union find data structure union_find: UnionFind, + /// Maps your calculated Global ID -> Internal UnionFind Index + node_map: FastIterMap, } type UnionFind = UnionFindGeneric; - +const PARTITION_STRIDE: usize = 1_000_000; // Large enough to never overlap /// define your own union-find node data structure like this #[derive(Debug, Clone)] pub struct PrimalModuleUnionFindNode { /// all the internal edges - pub internal_edges: FastIterSet, + pub internal_edges: FastIterSet, /// the corresponding node index with these internal edges pub node_index: NodeIndex, + // /// the dual node pointer + // pub node_ptr: DualNodePtr, } /// example trait implementation @@ -63,10 +69,18 @@ impl UnionNodeTrait for PrimalModuleUnionFindNode { } } +impl PrimalModuleUnionFindNode { + fn get_global_id(node: &DualNodePtr) -> usize { + let (partition_id, local_index) = node.get_ord(); // Assuming you store (part, idx) + (partition_id * PARTITION_STRIDE) + local_index + } +} + impl PrimalModuleImpl for PrimalModuleUnionFind { fn new_empty(_initializer: &Arc) -> Self { Self { union_find: UnionFind::new(0), + node_map: FastIterMap::new(), } } @@ -75,10 +89,10 @@ impl PrimalModuleImpl for PrimalModuleUnionFind { } #[allow(clippy::unnecessary_cast)] - fn load(&mut self, interface_ptr: &DualModuleInterfacePtr, _dual_module: &mut D) { + fn load(&mut self, interface_ptr: &DualModuleInterfacePtr, _dual_module: &mut D, partition_id: usize) { let interface = interface_ptr.read_recursive(); - for index in 0..interface.nodes.len() as NodeIndex { - let node_ptr = &interface.nodes[index as usize]; + for index in 0..interface.nodes.len() { + let node_ptr = &interface.nodes[index]; let node = node_ptr.read_recursive(); debug_assert!( node.invalid_subgraph.edges.is_empty(), @@ -97,10 +111,12 @@ impl PrimalModuleImpl for PrimalModuleUnionFind { self.union_find.size(), "must load defect nodes in order, did you forget to call solver.clear()?" ); - self.union_find.insert(PrimalModuleUnionFindNode { + let internal_id = self.union_find.insert(PrimalModuleUnionFindNode { internal_edges: FastIterSet::new(), node_index: node.index, }); + let global_id = PrimalModuleUnionFindNode::get_global_id(&node_ptr); + self.node_map.insert(global_id, node.index); } } @@ -112,28 +128,32 @@ impl PrimalModuleImpl for PrimalModuleUnionFind { dual_module: &mut impl DualModuleImpl, ) -> bool { debug_assert!(!dual_report.is_unbounded() && dual_report.get_valid_growth().is_none()); - let mut active_clusters = FastIterSet::::new(); + let mut active_clusters = FastIterSet::::new(); // set of internal cluster indices instead of global indices while let Some(obstacle) = dual_report.pop() { match obstacle { - Obstacle::Conflict { edge_index } => { + Obstacle::Conflict { edge_ptr } => { // union all the dual nodes in the edge index and create new dual node by adding this edge to `internal_edges` - let dual_nodes = dual_module.get_edge_nodes(edge_index); + let dual_nodes = dual_module.get_edge_nodes(edge_ptr.clone()); debug_assert!( !dual_nodes.is_empty(), "should not conflict if no dual nodes are contributing" ); - let cluster_index = dual_nodes[0].read_recursive().index; + let cluster_global_index = PrimalModuleUnionFindNode::get_global_id(&dual_nodes[0]); + let cluster_internal_index = self.node_map.get(&cluster_global_index).expect("all nodes should be in the map"); for dual_node_ptr in dual_nodes.iter() { dual_module.set_grow_rate(dual_node_ptr, Rational::zero()); - let node_index = dual_node_ptr.read_recursive().index; - active_clusters.remove(&(self.union_find.find(node_index as usize) as NodeIndex)); - self.union_find.union(cluster_index as usize, node_index as usize); + let node_global_index = PrimalModuleUnionFindNode::get_global_id(dual_node_ptr); + let node_internal_index = self.node_map.get(&node_global_index).expect("all nodes should be in the map"); + active_clusters.remove(&(self.union_find.find(*cluster_internal_index as usize) as NodeIndex)); + active_clusters.remove(&(self.union_find.find(*node_internal_index) as NodeIndex)); + self.union_find.union(*cluster_internal_index as usize, *node_internal_index as usize); } + let new_root_index = self.union_find.find(*cluster_internal_index); self.union_find - .get_mut(cluster_index as usize) + .get_mut(*cluster_internal_index as usize) .internal_edges - .insert(edge_index); - active_clusters.insert(self.union_find.find(cluster_index as usize) as NodeIndex); + .insert(edge_ptr); + active_clusters.insert(new_root_index as NodeIndex); } _ => { unreachable!() @@ -142,23 +162,27 @@ impl PrimalModuleImpl for PrimalModuleUnionFind { } for &cluster_index in active_clusters.iter() { if interface_ptr - .read_recursive() - .decoding_graph .is_valid_cluster_auto_vertices(&self.union_find.get(cluster_index as usize).internal_edges) { // do nothing } else { - let new_cluster_node_index = self.union_find.size() as NodeIndex; + let new_cluster_node_index = self.union_find.size() as NodeIndex; // it is an internal index self.union_find.insert(PrimalModuleUnionFindNode { internal_edges: FastIterSet::new(), node_index: new_cluster_node_index, }); + self.union_find.union(cluster_index as usize, new_cluster_node_index as usize); + // Get the final merged edges from the new root (which might be new_cluster_node_index or cluster_index depending on UF weight) + let final_root = self.union_find.find(new_cluster_node_index as usize); let invalid_subgraph = InvalidSubgraph::new_ptr( - self.union_find.get(cluster_index as usize).internal_edges.clone(), - &interface_ptr.read_recursive().decoding_graph, + self.union_find.get(final_root).internal_edges.clone(), dual_module ); - interface_ptr.create_node(invalid_subgraph, dual_module); + let created_dual_node_ptr = interface_ptr.create_node(invalid_subgraph, dual_module, interface_ptr.get_ord().0); + // we should add this new node to the node map + let global_id = PrimalModuleUnionFindNode::get_global_id(&created_dual_node_ptr); + let internal_id = created_dual_node_ptr.get_ord().0; + self.node_map.insert(global_id, internal_id); } } false @@ -176,13 +200,16 @@ impl PrimalModuleImpl for PrimalModuleUnionFind { if !valid_clusters.contains(&root_index) { valid_clusters.insert(root_index); let cluster_subgraph = interface_ptr - .read_recursive() - .decoding_graph .find_valid_subgraph_auto_vertices(&self.union_find.get(root_index).internal_edges) .expect("must be valid cluster"); - subgraph.extend(cluster_subgraph.iter()); + subgraph.extend(cluster_subgraph.into_iter()); } } + let subgraph_index = subgraph + .clone() + .into_iter() + .map(|e| e.upgrade_force().read_recursive().edge_index) + .collect::>(); // let mut subgraph_set = subgraph.into_iter().collect::>(); // for to_flip in _dual_module.get_negative_edges().iter() { @@ -195,7 +222,7 @@ impl PrimalModuleImpl for PrimalModuleUnionFind { // OutputSubgraph::new(subgraph_set.into_iter().collect(), Default::default()) // note: note implmented to handle negative weights yet - OutputSubgraph::new(subgraph, _dual_module.get_negative_edges()) + OutputSubgraph::new(subgraph_index, _dual_module.get_negative_edges(), subgraph) } } @@ -233,7 +260,7 @@ pub mod tests { let mut primal_module = PrimalModuleUnionFind::new_empty(&model_graph.initializer); // try to work on a simple syndrome code.set_defect_vertices(&defect_vertices); - let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone()); + let interface_ptr = DualModuleInterfacePtr::new(model_graph.clone(), 0); primal_module.solve_visualizer( &interface_ptr, Arc::new(code.get_syndrome()), @@ -301,7 +328,7 @@ pub mod tests { code, defect_vertices, final_dual, - DualModulePQ::new_empty(&model_graph.initializer), + DualModulePQ::new_empty(&model_graph.initializer, 0), model_graph, Some(visualizer), ) diff --git a/src/relaxer.rs b/src/relaxer.rs index 1435f4ef..e61ee163 100644 --- a/src/relaxer.rs +++ b/src/relaxer.rs @@ -6,6 +6,7 @@ use std::cmp::Ordering; // use std::collections::hash_map::DefaultHasher; use std::hash::{Hash, Hasher}; use std::sync::Arc; +use crate::dual_module_pq::{EdgePtr}; #[derive(Clone, Eq, Derivative, Default)] #[derivative(Debug)] @@ -17,9 +18,9 @@ pub struct Relaxer { direction: FastIterMap, Rational>, /// the edges that will be untightened after growing along `direction`; /// basically all the edges that have negative `overall_growing_rate` - untighten_edges: FastIterMap, + untighten_edges: FastIterMap, /// the edges that will grow - growing_edges: FastIterMap, + growing_edges: FastIterMap, } impl Hash for Relaxer { @@ -70,21 +71,21 @@ impl Relaxer { pub fn new_raw(direction: FastIterMap, Rational>) -> Self { let mut edges = FastIterMap::new(); for (invalid_subgraph, speed) in direction.iter() { - for &edge_index in invalid_subgraph.hair.iter() { - if let Some(mut edge) = edges.get_mut(&edge_index) { + for edge_ptr in invalid_subgraph.hair.iter() { + if let Some(mut edge) = edges.get_mut(edge_ptr) { *edge += speed; continue; } - edges.insert(edge_index, speed.clone()); + edges.insert(edge_ptr.clone(), speed.clone()); } } let mut untighten_edges = FastIterMap::new(); let mut growing_edges = FastIterMap::new(); - for (edge_index, speed) in edges { + for (edge_ptr, speed) in edges { if speed.is_negative() { - untighten_edges.insert(edge_index, speed); + untighten_edges.insert(edge_ptr, speed); } else if speed.is_positive() { - growing_edges.insert(edge_index, speed); + growing_edges.insert(edge_ptr, speed); } } let mut relaxer = Self { @@ -128,11 +129,11 @@ impl Relaxer { &self.direction } - pub fn get_growing_edges(&self) -> &FastIterMap { + pub fn get_growing_edges(&self) -> &FastIterMap { &self.growing_edges } - pub fn get_untighten_edges(&self) -> &FastIterMap { + pub fn get_untighten_edges(&self) -> &FastIterMap { &self.untighten_edges } } @@ -142,6 +143,8 @@ mod tests { use super::*; use crate::decoding_hypergraph::tests::*; use crate::invalid_subgraph::tests::*; + use crate::dual_module_pq::DualModulePQ; + use crate::dual_module::{DualModuleInterfacePtr, DualModuleImpl}; use num_traits::One; #[test] @@ -149,13 +152,18 @@ mod tests { // cargo test relaxer_good -- --nocapture let visualize_filename = "relaxer_good.json".to_string(); let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); - let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete( + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); // initialize vertex and edge pointers + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); + interface_ptr.load(decoding_graph.syndrome_pattern.clone(), &mut dual_module, 0); // this is needed to load the defect vertices + + let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete_from_indices( vec![7].into_iter().collect(), FastIterSet::new(), - decoding_graph.as_ref(), + &mut dual_module )); use num_traits::One; - let relaxer = Relaxer::new([(invalid_subgraph, Rational::one())].into()); + let relaxer = Relaxer::new([(invalid_subgraph, Rational::one())].into_iter().collect()); println!("relaxer: {relaxer:?}"); assert!(relaxer.untighten_edges.is_empty()); } @@ -166,24 +174,36 @@ mod tests { // cargo test relaxer_bad -- --nocapture let visualize_filename = "relaxer_bad.json".to_string(); let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); - let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete( + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); // initialize vertex and edge pointers + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); + interface_ptr.load(decoding_graph.syndrome_pattern.clone(), &mut dual_module, 0); // this is needed to load the defect vertices + + let invalid_subgraph = Arc::new(InvalidSubgraph::new_complete_from_indices( vec![7].into_iter().collect(), FastIterSet::new(), - decoding_graph.as_ref(), + &mut dual_module )); - let relaxer: Relaxer = Relaxer::new([(invalid_subgraph, Rational::zero())].into()); + let relaxer: Relaxer = Relaxer::new([(invalid_subgraph, Rational::zero())].into_iter().collect()); println!("relaxer: {relaxer:?}"); // should not print because it panics } #[test] fn relaxer_hash() { // cargo test relaxer_hash -- --nocapture + let visualize_filename = "relaxer_hash.json".to_string(); + let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); // initialize vertex and edge pointers + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); + interface_ptr.load(decoding_graph.syndrome_pattern.clone(), &mut dual_module, 0); // this is needed to load the defect vertices let vertices: FastIterSet = [1, 2, 3].into(); let edges: FastIterSet = [4, 5].into(); let hair: FastIterSet = [6, 7, 8].into(); - let invalid_subgraph = InvalidSubgraph::new_raw(vertices.clone(), edges.clone(), hair.clone()); - let relaxer_1 = Relaxer::new([(Arc::new(invalid_subgraph.clone()), Rational::one())].into()); - let relaxer_2 = Relaxer::new([(Arc::new(invalid_subgraph), Rational::one())].into()); + let invalid_subgraph = + InvalidSubgraph::new_raw_from_indices(vertices.clone(), edges.clone(), hair.clone(), &mut dual_module); + let relaxer_1 = Relaxer::new([(Arc::new(invalid_subgraph.clone()), Rational::one())].into_iter().collect()); + let relaxer_2 = Relaxer::new([(Arc::new(invalid_subgraph), Rational::one())].into_iter().collect()); assert_eq!(relaxer_1, relaxer_2); // they should have the same hash value assert_eq!( diff --git a/src/relaxer_forest.rs b/src/relaxer_forest.rs index 48587a50..34e01083 100644 --- a/src/relaxer_forest.rs +++ b/src/relaxer_forest.rs @@ -8,6 +8,7 @@ use crate::num_traits::Zero; use crate::relaxer::*; use crate::util::*; use num_traits::Signed; +use crate::dual_module_pq::{EdgePtr, EdgeWeak}; use std::sync::Arc; @@ -17,13 +18,13 @@ pub type RelaxerVec = Vec; pub struct RelaxerForest { /// keep track of the remaining tight edges for quick validation: /// these edges cannot grow unless untightened by some relaxers - tight_edges: FastIterSet, + tight_edges: FastIterSet, /// keep track of the subgraphs that are allowed to shrink: /// these should be all positive dual variables, all others are yS = 0 shrinkable_subgraphs: FastIterSet>, /// each untightened edge corresponds to a relaxer with speed: /// to untighten the edge for a unit length, how much should a relaxer be executed - edge_untightener: FastIterMap, Rational)>, + edge_untightener: FastIterMap, Rational)>, /// expanded relaxer results, as part of the dynamic programming: /// the expanded relaxer is a valid relaxer only growing of initial un-tight edges, /// not any edges untightened by other relaxers @@ -36,11 +37,11 @@ pub const FOREST_ERR_MSG_UNSHRINKABLE: &str = "invalid relaxer: try to shrink a impl RelaxerForest { pub fn new(tight_edges: IterEdge, shrinkable_subgraphs: IterSubgraph) -> Self where - IterEdge: Iterator, + IterEdge: Iterator, IterSubgraph: Iterator>, { Self { - tight_edges: FastIterSet::from_iter(tight_edges), + tight_edges: FastIterSet::from_iter(tight_edges.map(|e| e.upgrade_force())), shrinkable_subgraphs: FastIterSet::from_iter(shrinkable_subgraphs), edge_untightener: FastIterMap::new(), expanded_relaxers: FastIterMap::new(), @@ -53,8 +54,9 @@ impl RelaxerForest { // non-negative overall speed and effectiveness check relaxer.sanity_check()?; // a relaxer cannot grow any tight edge - for (edge_index, _) in relaxer.get_growing_edges().iter() { - if self.tight_edges.contains(edge_index) && !self.edge_untightener.contains_key(edge_index) { + for (edge_ptr, _) in relaxer.get_growing_edges().iter() { + if self.tight_edges.contains(edge_ptr) && !self.edge_untightener.contains_key(edge_ptr) { + let edge_index = edge_ptr.read_recursive().edge_index; return Err(format!("{FOREST_ERR_MSG_GROW_TIGHT_EDGE}: {edge_index}")); } } @@ -72,10 +74,10 @@ impl RelaxerForest { // validate only at debug mode to improve speed debug_assert_eq!(self.validate(&relaxer), Ok(())); // add this relaxer to the forest - for (edge_index, speed) in relaxer.get_untighten_edges().iter() { + for (edge_ptr, speed) in relaxer.get_untighten_edges().iter() { debug_assert!(speed.is_negative()); - if !self.edge_untightener.contains_key(edge_index) { - self.edge_untightener.insert(*edge_index, (relaxer.clone(), -speed.recip())); + if !self.edge_untightener.contains_key(edge_ptr) { + self.edge_untightener.insert(edge_ptr.clone(), (relaxer.clone(), -speed.recip())); } } } @@ -84,17 +86,18 @@ impl RelaxerForest { if self.expanded_relaxers.contains_key(relaxer) { return; } - let mut untightened_edges: FastIterMap = FastIterMap::new(); + let mut untightened_edges: FastIterMap = FastIterMap::new(); let mut directions: FastIterMap, Rational> = relaxer.get_direction().clone(); - for (edge_index, speed) in relaxer.get_growing_edges().iter() { + for (edge_ptr, speed) in relaxer.get_growing_edges().iter() { debug_assert!(speed.is_positive()); - if self.tight_edges.contains(edge_index) { + if self.tight_edges.contains(edge_ptr) { + let edge_index = edge_ptr.read_recursive().edge_index; debug_assert!( - self.edge_untightener.contains_key(edge_index), + self.edge_untightener.contains_key(edge_ptr), "edge {} is tight but no untightener presents, thus new relaxer cannot grow on it", edge_index ); - let require_speed = if let Some(mut existing_speed) = untightened_edges.get_mut(edge_index) { + let require_speed = if let Some(mut existing_speed) = untightened_edges.get_mut(edge_ptr) { if *existing_speed >= *speed { *existing_speed -= speed; Rational::zero() @@ -108,9 +111,9 @@ impl RelaxerForest { }; if require_speed.is_positive() { // we need to invoke another relaxer to untighten this edge - let edge_relaxer = self.edge_untightener.get(edge_index).unwrap().0.clone(); + let edge_relaxer = self.edge_untightener.get(edge_ptr).unwrap().0.clone(); self.compute_expanded(&edge_relaxer); - let (edge_relaxer, speed_ratio) = self.edge_untightener.get(edge_index).unwrap(); + let (edge_relaxer, speed_ratio) = self.edge_untightener.get(edge_ptr).unwrap(); debug_assert!(speed_ratio.is_positive()); let expanded_edge_relaxer = self.expanded_relaxers.get(edge_relaxer).unwrap(); for (subgraph, original_speed) in expanded_edge_relaxer.get_direction().iter() { @@ -121,17 +124,17 @@ impl RelaxerForest { } directions.insert(subgraph.clone(), new_speed); } - for (edge_index, original_speed) in expanded_edge_relaxer.get_untighten_edges().iter() { + for (edge_ptr, original_speed) in expanded_edge_relaxer.get_untighten_edges().iter() { debug_assert!(original_speed.is_negative()); let new_speed = -original_speed * speed_ratio * require_speed.clone(); - if let Some(mut speed) = untightened_edges.get_mut(edge_index) { + if let Some(mut speed) = untightened_edges.get_mut(edge_ptr) { *speed += new_speed; continue; } - untightened_edges.insert(*edge_index, new_speed); + untightened_edges.insert(edge_ptr.clone(), new_speed); } - debug_assert_eq!(untightened_edges.get(edge_index), Some(&require_speed)); - *untightened_edges.get_mut(edge_index).unwrap() -= require_speed; + debug_assert_eq!(untightened_edges.get(edge_ptr), Some(&require_speed)); + *untightened_edges.get_mut(edge_ptr).unwrap() -= require_speed; } } } @@ -156,30 +159,46 @@ impl RelaxerForest { pub mod tests { use super::*; use num_traits::{FromPrimitive, One}; + use crate::decoding_hypergraph::tests::color_code_5_decoding_graph; + use crate::dual_module::DualModuleImpl; + use crate::dual_module::DualModuleInterfacePtr; + use crate::dual_module_pq::DualModulePQ; #[test] fn relaxer_forest_example() { // cargo test relaxer_forest_example -- --nocapture + let visualize_filename = "relaxer_forest_example.json".to_string(); + let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); + interface_ptr.load(decoding_graph.syndrome_pattern.clone(), &mut dual_module, 0); // this is needed to load the defect vertices + let tight_edges = [0, 1, 2, 3, 4, 5, 6]; let shrinkable_subgraphs = [ - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [1, 2, 3].into())), - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [4, 5].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [1, 2, 3].into(), &mut dual_module)), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [4, 5].into(), &mut dual_module)), ]; - let mut relaxer_forest = RelaxerForest::new(tight_edges.into_iter(), shrinkable_subgraphs.iter().cloned()); - let invalid_subgraph_1 = Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [7, 8, 9].into())); + let tight_edges_weak = dual_module + .get_edge_ptr_vec(&tight_edges) + .into_iter() + .map(|e| e.downgrade()) + .collect::>(); + let mut relaxer_forest = RelaxerForest::new(tight_edges_weak.into_iter(), shrinkable_subgraphs.iter().cloned()); + let invalid_subgraph_1 = Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [7, 8, 9].into(), &mut dual_module)); let relaxer_1 = Arc::new(Relaxer::new_raw( [ (invalid_subgraph_1.clone(), Rational::one()), (shrinkable_subgraphs[0].clone(), -Rational::one()), ] - .into(), + .into_iter().collect(), )); let expanded_1 = relaxer_forest.expand(&relaxer_1); assert_eq!(expanded_1, *relaxer_1); relaxer_forest.add(relaxer_1); // now add a relaxer that is relying on relaxer_1 - let invalid_subgraph_2 = Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [1, 2, 7].into())); - let relaxer_2 = Arc::new(Relaxer::new_raw([(invalid_subgraph_2.clone(), Rational::one())].into())); + let invalid_subgraph_2 = Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [1, 2, 7].into(), &mut dual_module)); + let relaxer_2 = Arc::new(Relaxer::new_raw([(invalid_subgraph_2.clone(), Rational::one())].into_iter().collect())); let expanded_2 = relaxer_forest.expand(&relaxer_2); let expected_relaxer = Relaxer::new( [ @@ -187,7 +206,7 @@ pub mod tests { (shrinkable_subgraphs[0].clone(), -Rational::one()), (invalid_subgraph_2, Rational::one()), ] - .into(), + .into_iter().collect(), ); // println!("{expanded_2:#?}"); // println!("{expected_relaxer:#?}"); @@ -197,29 +216,41 @@ pub mod tests { #[test] fn relaxer_forest_require_multiple() { // cargo test relaxer_forest_require_multiple -- --nocapture + let visualize_filename = "relaxer_forest_require_multiple.json".to_string(); + let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); + interface_ptr.load(decoding_graph.syndrome_pattern.clone(), &mut dual_module, 0); // this is needed to load the defect vertices + let tight_edges = [0, 1, 2, 3, 4, 5, 6]; let shrinkable_subgraphs = [ - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [1, 2].into())), - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [3].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [1, 2].into(), &mut dual_module)), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [3].into(), &mut dual_module)), ]; - let mut relaxer_forest = RelaxerForest::new(tight_edges.into_iter(), shrinkable_subgraphs.iter().cloned()); - let invalid_subgraph_1 = Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [7, 8, 9].into())); + let tight_edges_weak = dual_module + .get_edge_ptr_vec(&tight_edges) + .into_iter() + .map(|e| e.downgrade()) + .collect::>(); + let mut relaxer_forest = RelaxerForest::new(tight_edges_weak.into_iter(), shrinkable_subgraphs.iter().cloned()); + let invalid_subgraph_1 = Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [7, 8, 9].into(), &mut dual_module)); let relaxer_1 = Arc::new(Relaxer::new_raw( [ (invalid_subgraph_1.clone(), Rational::one()), (shrinkable_subgraphs[0].clone(), -Rational::one()), ] - .into(), + .into_iter().collect(), )); relaxer_forest.add(relaxer_1); - let invalid_subgraph_2 = Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [1, 2, 7].into())); - let invalid_subgraph_3 = Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [2].into())); + let invalid_subgraph_2 = Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [1, 2, 7].into(), &mut dual_module)); + let invalid_subgraph_3 = Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [2].into(), &mut dual_module)); let relaxer_2 = Arc::new(Relaxer::new_raw( [ (invalid_subgraph_2.clone(), Rational::one()), (invalid_subgraph_3.clone(), Rational::one()), ] - .into(), + .into_iter().collect(), )); let expanded_2 = relaxer_forest.expand(&relaxer_2); assert_eq!( @@ -231,7 +262,7 @@ pub mod tests { (invalid_subgraph_1, Rational::from_usize(2).unwrap()), (shrinkable_subgraphs[0].clone(), -Rational::from_usize(2).unwrap()), ] - .into() + .into_iter().collect() ) ); // println!("{expanded_2:#?}"); @@ -240,28 +271,40 @@ pub mod tests { #[test] fn relaxer_forest_relaxing_same_edge() { // cargo test relaxer_forest_relaxing_same_edge -- --nocapture + let visualize_filename = "relaxer_forest_relaxing_same_edge.json".to_string(); + let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); // initialize vertex and edge pointers + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); + interface_ptr.load(decoding_graph.syndrome_pattern.clone(), &mut dual_module, 0); // this is needed to load the defect vertices + let tight_edges = [0, 1, 2, 3, 4, 5, 6]; let shrinkable_subgraphs = [ - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [1, 2].into())), - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [2, 3].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [1, 2].into(), &mut dual_module)), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [2, 3].into(), &mut dual_module)), ]; - let mut relaxer_forest = RelaxerForest::new(tight_edges.into_iter(), shrinkable_subgraphs.iter().cloned()); - let invalid_subgraph_1 = Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [7, 8, 9].into())); + let tight_edges_weak = dual_module + .get_edge_ptr_vec(&tight_edges) + .into_iter() + .map(|e| e.downgrade()) + .collect::>(); + let mut relaxer_forest = RelaxerForest::new(tight_edges_weak.into_iter(), shrinkable_subgraphs.iter().cloned()); + let invalid_subgraph_1 = Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [7, 8, 9].into(), &mut dual_module)); let relaxer_1 = Arc::new(Relaxer::new_raw( [ (invalid_subgraph_1.clone(), Rational::one()), (shrinkable_subgraphs[0].clone(), -Rational::one()), ] - .into(), + .into_iter().collect(), )); relaxer_forest.add(relaxer_1); - let invalid_subgraph_2 = Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [10, 11].into())); + let invalid_subgraph_2 = Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [10, 11].into(), &mut dual_module)); let relaxer_2 = Arc::new(Relaxer::new_raw( [ (invalid_subgraph_2.clone(), Rational::one()), (shrinkable_subgraphs[1].clone(), -Rational::one()), ] - .into(), + .into_iter().collect(), )); relaxer_forest.add(relaxer_2); } @@ -269,20 +312,32 @@ pub mod tests { #[test] fn relaxer_forest_validate() { // cargo test relaxer_forest_validate -- --nocapture + let visualize_filename = "relaxer_forest_validate.json".to_string(); + let (decoding_graph, ..) = color_code_5_decoding_graph(vec![7, 1], visualize_filename); + let initializer = decoding_graph.model_graph.initializer.clone(); + let mut dual_module = DualModulePQ::new_empty(&initializer, 0); // initialize vertex and edge pointers + let interface_ptr = DualModuleInterfacePtr::new(decoding_graph.model_graph.clone(), 0); + interface_ptr.load(decoding_graph.syndrome_pattern.clone(), &mut dual_module, 0); // this is needed to load the defect vertices + let tight_edges = [0, 1, 2, 3, 4, 5, 6]; let shrinkable_subgraphs = [ - Arc::new(InvalidSubgraph::new_raw([1].into(), [].into(), [1, 2].into())), - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([1].into(), [].into(), [1, 2].into(), &mut dual_module)), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [].into(), &mut dual_module)), ]; - let relaxer_forest = RelaxerForest::new(tight_edges.into_iter(), shrinkable_subgraphs.iter().cloned()); + let tight_edges_weak = dual_module + .get_edge_ptr_vec(&tight_edges) + .into_iter() + .map(|e| e.downgrade()) + .collect::>(); + let relaxer_forest = RelaxerForest::new(tight_edges_weak.into_iter(), shrinkable_subgraphs.iter().cloned()); println!("relaxer_forest: {:?}", relaxer_forest.shrinkable_subgraphs); // invalid relaxer is forbidden let invalid_relaxer = Relaxer::new_raw( [( - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [].into(), &mut dual_module)), -Rational::one(), )] - .into(), + .into_iter().collect(), ); let error_message = relaxer_forest.validate(&invalid_relaxer).expect_err("should panic"); assert_eq!( @@ -292,10 +347,10 @@ pub mod tests { // relaxer that increases a tight edge is forbidden let relaxer = Relaxer::new_raw( [( - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [1].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [1].into(), &mut dual_module)), Rational::one(), )] - .into(), + .into_iter().collect(), ); let error_message = relaxer_forest.validate(&relaxer).expect_err("should panic"); assert_eq!( @@ -306,15 +361,15 @@ pub mod tests { let relaxer = Relaxer::new_raw( [ ( - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [9].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [9].into(), &mut dual_module)), Rational::one(), ), ( - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [2, 3].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [2, 3].into(), &mut dual_module)), -Rational::one(), ), ] - .into(), + .into_iter().collect(), ); let error_message = relaxer_forest.validate(&relaxer).expect_err("should panic"); assert_eq!( @@ -324,10 +379,10 @@ pub mod tests { // otherwise a relaxer is ok let relaxer = Relaxer::new_raw( [( - Arc::new(InvalidSubgraph::new_raw([].into(), [].into(), [9].into())), + Arc::new(InvalidSubgraph::new_raw_from_indices([].into(), [].into(), [9].into(), &mut dual_module)), Rational::one(), )] - .into(), + .into_iter().collect(), ); relaxer_forest.validate(&relaxer).unwrap(); } diff --git a/src/relaxer_optimizer.rs b/src/relaxer_optimizer.rs index 9e37a1c8..b588c44f 100644 --- a/src/relaxer_optimizer.rs +++ b/src/relaxer_optimizer.rs @@ -9,6 +9,7 @@ use crate::invalid_subgraph::*; use crate::relaxer::*; use crate::util::*; +use crate::dual_module_pq::EdgePtr; use std::sync::Arc; @@ -219,7 +220,7 @@ impl RelaxerOptimizer { pub fn optimize( &mut self, relaxer: Relaxer, - edge_slacks: FastIterMap, + edge_slacks: FastIterMap, mut dual_variables: FastIterMap, Rational>, ) -> (Relaxer, bool) { use highs::{HighsModelStatus, RowProblem, Sense}; @@ -240,8 +241,8 @@ impl RelaxerOptimizer { let mut x_vars = vec![]; let mut y_vars = vec![]; let mut invalid_subgraphs = Vec::with_capacity(dual_variables.len()); - let mut edge_contributor: FastIterMap> = - edge_slacks.keys().map(|&edge_index| (edge_index, vec![])).collect(); + let mut edge_contributor: FastIterMap> = + edge_slacks.keys().map(|edge_ptr| (edge_ptr.clone(), vec![])).collect(); for (var_index, (invalid_subgraph, dual_variable)) in dual_variables.iter().enumerate() { // constraint of the dual variable >= 0 @@ -257,14 +258,14 @@ impl RelaxerOptimizer { ); invalid_subgraphs.push(invalid_subgraph.clone()); - for &edge_index in invalid_subgraph.hair.iter() { - edge_contributor.get_mut(&edge_index).unwrap().push(var_index); + for edge_ptr in invalid_subgraph.hair.iter() { + edge_contributor.get_mut(edge_ptr).unwrap().push(var_index); } } - for (&edge_index, ref slack) in edge_slacks.iter() { + for (edge_ptr, ref slack) in edge_slacks.iter() { let mut row_entries = vec![]; - for &var_index in edge_contributor[&edge_index].iter() { + for &var_index in edge_contributor[edge_ptr].iter() { row_entries.push((x_vars[var_index], 1.0)); row_entries.push((y_vars[var_index], -1.0)); } diff --git a/src/util.rs b/src/util.rs index 1b5d5e76..23953dd6 100644 --- a/src/util.rs +++ b/src/util.rs @@ -21,6 +21,7 @@ use std::fs::File; use std::io::prelude::*; use std::io::{BufReader, BufWriter}; use std::time::Instant; +use crate::dual_module_pq::{EdgeWeak, EdgePtr}; cfg_if::cfg_if! { if #[cfg(feature="f64_weight")] { @@ -70,6 +71,7 @@ pub type Weight = Rational; pub type EdgeIndex = usize; pub type VertexIndex = usize; pub type HeraldIndex = usize; +pub type HeraldPtr = EdgePtr; pub type KnownSafeRefCell = std::cell::RefCell; pub type NodeIndex = VertexIndex; @@ -378,19 +380,32 @@ impl SolverInitializer { #[allow(clippy::unnecessary_cast)] pub fn get_subgraph_total_weight(&self, subgraph: &OutputSubgraph) -> Weight { - let mut weight = Weight::zero(); - for &edge_index in subgraph.iter() { - weight += self.weighted_edges[edge_index as usize].weight.clone(); + let internal_subgraph = OutputSubgraph::get_internal_subgraph(&subgraph); + if internal_subgraph.is_empty() { + let mut weight = Weight::zero(); + for &edge_index in subgraph.iter() { + weight += self.weighted_edges[edge_index as usize].weight.clone(); + } + return weight; + } else { + let mut weight = Weight::zero(); + for edge_weak in internal_subgraph.iter() { + weight += edge_weak.upgrade_force().read_recursive().weight.clone(); + } + return weight; } - weight } #[allow(clippy::unnecessary_cast)] pub fn get_subgraph_syndrome(&self, subgraph: &OutputSubgraph) -> FastIterSet { + let internal_subgraph = OutputSubgraph::get_internal_subgraph(&subgraph); let mut defect_vertices = FastIterSet::new(); - for &edge_index in subgraph.iter() { - let HyperEdge { vertices, .. } = &self.weighted_edges[edge_index as usize]; - for &vertex_index in vertices.iter() { + for edge_weak in internal_subgraph.iter() { + let edge_ptr = edge_weak.upgrade_force(); + let edge = edge_ptr.read_recursive(); + let vertices = &edge.vertices; + let unique_vertices = vertices.into_iter().map(|v| v.upgrade_force().read_recursive().vertex_index).collect::>(); + for &vertex_index in unique_vertices.iter() { if defect_vertices.contains(&vertex_index) { defect_vertices.remove(&vertex_index); // println!("duplicate defect vertex: {}", vertex_index); @@ -647,17 +662,20 @@ impl F64Rng for DeterministicRng { /// the result of MWPF algorithm: a parity subgraph (defined by some edges that, /// if are selected, will generate the parity result in the syndrome) pub type Subgraph = Vec; +pub type InternalSubgraph = Vec; pub struct OutputSubgraph { pub subgraph: Subgraph, pub flip_edge_indices: hashbrown::HashSet, + internal_subgraph: InternalSubgraph, // for internal use only, not exposed to users } impl OutputSubgraph { - pub fn new(subgraph: Subgraph, flip_edge_indices: hashbrown::HashSet) -> Self { + pub fn new(subgraph: Subgraph, flip_edge_indices: hashbrown::HashSet, internal_subgraph: InternalSubgraph) -> Self { Self { subgraph, flip_edge_indices, + internal_subgraph, } } @@ -677,14 +695,19 @@ impl OutputSubgraph { flip_edge_indices: &mut self.flip_edge_indices, } } -} -impl From for OutputSubgraph { - fn from(value: Subgraph) -> Self { - Self::new(value, hashbrown::HashSet::new()) + pub fn get_internal_subgraph(&self) -> &InternalSubgraph { + &self.internal_subgraph } } +// TODO: double check if we should comment this out +// impl From for OutputSubgraph { +// fn from(value: Subgraph) -> Self { +// Self::new(value, hashbrown::HashSet::new()) +// } +// } + // consuming iterators // Implementing `IntoIterator` for `&OutputSubgraph` (for `iter`) impl<'a> IntoIterator for &'a OutputSubgraph { @@ -1373,7 +1396,7 @@ pub mod tests { flip_edge_indices.insert(2); flip_edge_indices.insert(5); - let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices); + let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices, vec![]); // Expected behavior: `2` is skipped, and `5` is added at the end. let result: Vec<_> = output_subgraph.iter().cloned().collect(); @@ -1385,7 +1408,7 @@ pub mod tests { let subgraph = vec![1, 2, 3]; let flip_edge_indices = HashSet::new(); - let output_subgraph = OutputSubgraph::new(subgraph.clone(), flip_edge_indices); + let output_subgraph = OutputSubgraph::new(subgraph.clone(), flip_edge_indices, vec![]); // With empty `flip_edge_indices`, should just return all elements in `subgraph`. let result: Vec<_> = output_subgraph.iter().cloned().collect(); @@ -1401,7 +1424,7 @@ pub mod tests { flip_edge_indices.insert(3); flip_edge_indices.insert(4); - let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices); + let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices, vec![]); // Expected behavior: all elements in `subgraph` are skipped, and `4` is added at the end. let result: Vec<_> = output_subgraph.iter().cloned().collect(); @@ -1415,7 +1438,7 @@ pub mod tests { flip_edge_indices.insert(2); flip_edge_indices.insert(5); - let mut output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices); + let mut output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices, vec![]); // Modify elements during mutable iteration for elem in output_subgraph.iter_mut() { @@ -1432,7 +1455,7 @@ pub mod tests { let subgraph = vec![10, 20, 30]; let flip_edge_indices = HashSet::new(); // Empty flip edge indices - let mut output_subgraph = OutputSubgraph::new(subgraph.clone(), flip_edge_indices); + let mut output_subgraph = OutputSubgraph::new(subgraph.clone(), flip_edge_indices, vec![]); // Expected to iterate through all without any modifications to flip_edge_indices for elem in output_subgraph.iter_mut() { @@ -1451,7 +1474,7 @@ pub mod tests { flip_edge_indices.insert(2); flip_edge_indices.insert(5); - let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices); + let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices, vec![]); // Consuming iterator, so `output_subgraph` cannot be used afterward let result: Vec<_> = output_subgraph.into_iter().collect(); @@ -1469,7 +1492,7 @@ pub mod tests { flip_edge_indices.insert(3); flip_edge_indices.insert(4); - let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices); + let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices, vec![]); // Consuming iterator, expected to yield only `4` at the end since all `subgraph` elements are flipped. let result: Vec<_> = output_subgraph.into_iter().collect(); @@ -1483,7 +1506,7 @@ pub mod tests { flip_edge_indices.insert(1); flip_edge_indices.insert(2); - let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices); + let output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices, vec![]); // With empty `subgraph`, should only yield elements in `flip_edge_indices` let mut result: Vec<_> = output_subgraph.iter().cloned().collect(); @@ -1498,7 +1521,7 @@ pub mod tests { flip_edge_indices.insert(2); flip_edge_indices.insert(5); - let mut output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices); + let mut output_subgraph = OutputSubgraph::new(subgraph, flip_edge_indices, vec![]); // Expected behavior: `2` is skipped, and `5` is added at the end. let result: Vec<_> = output_subgraph.iter_mut().map(|x| *x).collect();