Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1,132 changes: 419 additions & 713 deletions data/expected_outputs/smc-dials.json

Large diffs are not rendered by default.

417 changes: 417 additions & 0 deletions data/expected_outputs/smc-furniture-small.json

Large diffs are not rendered by default.

1,613 changes: 841 additions & 772 deletions data/expected_outputs/smc-furniture.json

Large diffs are not rendered by default.

708 changes: 369 additions & 339 deletions data/expected_outputs/smc-nuts-bolts.json

Large diffs are not rendered by default.

1,549 changes: 587 additions & 962 deletions data/expected_outputs/smc-wheels.json

Large diffs are not rendered by default.

10 changes: 10 additions & 0 deletions data/regression/furniture-small.json
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
[
"(T (repeat (repeat (T (C (C (T (T c (M 2 0 0 0)) (M 1.125 0 0 0)) (T (T c (M 2 0 0 0)) (M 1.625 0 0 0))) (T r (M 0.84375 0 0 0))) (M 1 0 0 0)) 1 (M 1 0 0 0)) 1 (M 1 0 0 0)) (M 1 0 -6.9375 0))",
"(T (repeat (repeat (T (C (T r (M 1.125 0 0 0)) (T c (M 0.84375 0 0 0))) (M 1 0 0 0)) 1 (M 1 0 0 0)) 1 (M 1 0 0 0)) (M 1 0 -7.1875 0))",
"(T (repeat (repeat (T (C (T (T (repeat (T l (M 1 0 -0.5 1.20711)) 8 (M 1 0.785398 0 0)) (M 0.75 0 0 0)) (M 1.125 0 0 0)) (T c (M 0.84375 0 0 0))) (M 1 0 0 0)) 1 (M 1 0 0 0)) 1 (M 1 0 0 0)) (M 1 0 -7.1875 0))",
"(T (repeat (repeat (T (C (T (T (repeat (T l (M 1 0 -0.5 1.20711)) 8 (M 1 0.785398 0 0)) (M 0.75 0 0 0)) (M 1.125 0 0 0)) (T c (M 0.84375 0 0 0))) (M 1 0 0 0)) 2 (M 1 0 14.375 0)) 1 (M 1 0 0 0)) (M 1 0 -7.1875 0))",
"(C (C (T (T (repeat (repeat (T (C (T (T (r_s 13.5 3) (M 1 0 0 0)) (M 1 0 0 0)) (T (repeat (repeat (T (T c (M 0.84375 0 0 0)) (M 1 0 0 0)) 2 (M 1 0 5.625 0)) 1 (M 1 0 0 0)) (M 1 0 -2.8125 0))) (M 1 0 0 1.5)) 1 (M 1 0 0 0)) 2 (M 1 0 0 3.75)) (M 1 0 0 0)) (M 1 0 0 0.75)) (T (T (r_s 15 8.25) (M 1 0 0 4.125)) (M 1 0 0 0))) (T (repeat (repeat (T (T (T (T l (M 1 0 -0.5 0)) (M 1 1.5708 0 0)) (M 3 0 0 0)) (M 1 0 0 -1.5)) 3 (M 1 0 7.5 0)) 1 (M 1 0 0 0)) (M 1 0 -7.5 0)))",
"(C (C (T (T (repeat (repeat (T (C (T (T (r_s 13.5 3) (M 1 0 0 0)) (M 1 0 0 0)) (T (repeat (repeat (T (T c (M 0.84375 0 0 0)) (M 1 0 0 0)) 2 (M 1 0 5.625 0)) 1 (M 1 0 0 0)) (M 1 0 -2.8125 0))) (M 1 0 0 1.5)) 1 (M 1 0 0 0)) 2 (M 1 0 0 3.75)) (M 1 0 0 0)) (M 1 0 0 0.75)) (T (T (r_s 15 8.25) (M 1 0 0 4.125)) (M 1 0 0 0))) (T (repeat (repeat (T (T (T (T l (M 1 0 -0.5 0)) (M 1 1.5708 0 0)) (M 6 0 0 0)) (M 1 0 0 -3)) 2 (M 1 0 15 0)) 1 (M 1 0 0 0)) (M 1 0 -7.5 0)))",
"(C (T (repeat (repeat (T (C (T (T (r_s 9 4) (M 1 0 0 0)) (M 1 0 0 0)) (T (repeat (repeat (T (C (T (T (repeat (T l (M 1 0 -0.5 1.20711)) 8 (M 1 0.785398 0 0)) (M 0.75 0 0 0)) (M 1.125 0 0 0)) (T c (M 0.84375 0 0 0))) (M 1 0 0 0)) 2 (M 1 0 2.875 0)) 1 (M 1 0 0 0)) (M 1 0 -1.4375 0))) (M 1 0 0 0)) 1 (M 1 0 0 0)) 4 (M 1 0 0 5)) (M 1 0 0 -7.5)) (T (T (r_s 11 21) (M 1 0 0 0)) (M 1 0 0 0)))",
"(C (C (T (T (repeat (repeat (T (C (T (T (r_s 13.5 3) (M 1 0 0 0)) (M 1 0 0 0)) (T (repeat (repeat (T (T r (M 0.84375 0 0 0)) (M 1 0 0 0)) 2 (M 1 0 5.625 0)) 1 (M 1 0 0 0)) (M 1 0 -2.8125 0))) (M 1 0 0 1.5)) 1 (M 1 0 0 0)) 2 (M 1 0 0 3.75)) (M 1 0 0 0)) (M 1 0 0 0.75)) (T (T (r_s 15 8.25) (M 1 0 0 4.125)) (M 1 0 0 0))) (T (repeat (repeat (T (T (T (T l (M 1 0 -0.5 0)) (M 1 1.5708 0 0)) (M 6 0 0 0)) (M 1 0 0 -3)) 4 (M 1 0 5 0)) 1 (M 1 0 0 0)) (M 1 0 -7.5 0)))"
]
45 changes: 45 additions & 0 deletions src/pattern_args.rs
Original file line number Diff line number Diff line change
Expand Up @@ -284,6 +284,51 @@ impl PatternArgs {
}
}

pub fn sort_args(&mut self, shared: &SharedData) {
let mut min_zip_for_ivar: Vec<Option<&Vec<ZNode>>> = vec![None; self.variables.len()];
for labeled in self.arg_choices.iter() {
let ivar = labeled.ivar;
let zip = &shared.zip_of_zid[labeled.zid];
if min_zip_for_ivar[ivar].is_none_or(|current| compare_zips(zip, current) == std::cmp::Ordering::Less) {
min_zip_for_ivar[ivar] = Some(zip);
}
}
let mut ivar_order: Vec<usize> = (0..self.variables.len()).collect();
ivar_order.sort_by(|&i, &j| compare_zips(min_zip_for_ivar[i].unwrap(),
min_zip_for_ivar[j].unwrap()));
// ivar_order.reverse();
let mut new_variables = vec![(0u32, VariableType::Unvalidated); self.variables.len()];
for (new_ivar, &old_ivar) in ivar_order.iter().enumerate() {
new_variables[new_ivar] = self.variables[old_ivar];
}
self.arg_choices.iter_mut().for_each(|labeled| {
let old_ivar = labeled.ivar;
let new_ivar = ivar_order.iter().position(|&x| x == old_ivar).unwrap();
labeled.ivar = new_ivar;
});
self.variables = new_variables;
}

}

fn compare_zips(zip1: &[ZNode], zip2: &[ZNode]) -> std::cmp::Ordering {
// compare two zips lexicographically
for (node1, node2) in zip1.iter().zip(zip2.iter()) {
if node1 == node2 {
continue;
}
assert!(*node1 != ZNode::Body && *node2 != ZNode::Body, "these zippers should be compatibly children of one argument");
if *node1 == ZNode::Arg {
assert!(*node2 == ZNode::Func);
return std::cmp::Ordering::Less; // args come first
}
assert!(*node1 == ZNode::Func && *node2 == ZNode::Arg, "these zippers should be compatibly children of one argument");
return std::cmp::Ordering::Greater;
}
if zip1.len() == zip2.len() {
return std::cmp::Ordering::Equal;
}
panic!("Zippers need to be the same length at this point, otherwise one variable is a child of another.");
}

pub struct LocationsForReusableArgs<'a> {
Expand Down
8 changes: 7 additions & 1 deletion src/smc.rs
Original file line number Diff line number Diff line change
Expand Up @@ -352,7 +352,7 @@ pub fn compression_step_smc(
return vec![];
};

let finished_pattern = FinishedPattern::new(best.pattern.clone(), &shared);
let finished_pattern = FinishedPattern::new(sorted_arguments(best.pattern, &shared), &shared);
let result = CompressionStepResult::new(
finished_pattern,
inv_name,
Expand All @@ -363,3 +363,9 @@ pub fn compression_step_smc(
);
vec![result]
}

fn sorted_arguments(pattern: Pattern, shared: &SharedData) -> Pattern {
let mut pattern = pattern;
pattern.pattern_args.sort_args(shared);
pattern
}
8 changes: 7 additions & 1 deletion tests/integration_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -134,14 +134,20 @@ const SMC_ARGS: &str = " --smc --smc-particles 1000 --smc-extra-steps 40";

#[test]
fn smc_regression_tests() {
let args = "-i10".to_owned() + SMC_ARGS;
let args = "-i10 --rewrite-check".to_owned() + SMC_ARGS;
compare_out_jsons_testing("data/cogsci/nuts-bolts.json", "data/expected_outputs/smc-nuts-bolts.json", &args, InputFormat::ProgramsList);
compare_out_jsons_testing("data/cogsci/wheels.json", "data/expected_outputs/smc-wheels.json", &args, InputFormat::ProgramsList);
compare_out_jsons_testing("data/cogsci/furniture.json", "data/expected_outputs/smc-furniture.json", &args, InputFormat::ProgramsList);
compare_out_jsons_testing("data/cogsci/dials.json", "data/expected_outputs/smc-dials.json", &args, InputFormat::ProgramsList);
// compare_out_jsons("data/cogsci/city.json", "data/expected_outputs/smc-city.json", &args, InputFormat::ProgramsList);
}

#[test]
fn smc_regression_tests_small() {
let args = "-i10 --rewrite-check".to_owned() + SMC_ARGS;
compare_out_jsons_testing("data/regression/furniture-small.json", "data/expected_outputs/smc-furniture-small.json", &args, InputFormat::ProgramsList);
}

fn python_args() -> String {
DFA_ARGS.to_owned() + " --symvar-prefix &"
}
Expand Down