variable classification al a carte

This commit is contained in:
Mark Thom
2022-11-13 10:13:37 -07:00
committed by Mark
parent 170818759d
commit a66d666bed
8 changed files with 198 additions and 188 deletions

View File

@@ -9,6 +9,7 @@ paper "Compiling Large Disjunctions" to Scryer Prolog.
*/
use crate::atom_table::*;
use crate::fixtures::VariableFixtures;
use crate::forms::*;
use crate::instructions::*;
use crate::iterators::*;
@@ -34,7 +35,7 @@ struct BranchNumber {
impl Default for BranchNumber {
fn default() -> Self {
Self {
branch_num: Rational::from(1 << 10),
branch_num: Rational::from(1 << 63),
delta: Rational::from(1),
}
}
@@ -86,16 +87,19 @@ impl BranchNumber {
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct VarInfo {
var_ptr: VarPtr,
classify_info: ClassifyInfo,
lvl: Level,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct ChunkInfo {
chunk_num: usize,
vars: Vec<VarPtr>, // pointer to incidence
}
impl ChunkInfo {
fn new(chunk_num: usize) -> Self {
ChunkInfo { chunk_num, vars: vec![] }
}
term_loc: GenContext,
// pointer to incidence, term occurrence arity.
vars: Vec<VarInfo>,
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
@@ -140,6 +144,23 @@ enum ChunkType {
Last,
}
impl ChunkType {
#[inline(always)]
fn to_gen_context(self, chunk_num: usize) -> GenContext {
match self {
ChunkType::Head => GenContext::Head,
ChunkType::Mid => GenContext::Mid(chunk_num),
ChunkType::Last => GenContext::Last(chunk_num),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct ClassifyInfo {
arg_c: usize,
arity: usize,
}
enum TraversalState {
// construct a QueryTerm::Branch with number of disjuncts, reset
// the chunk type to that of the chunk preceding the disjunct.
@@ -199,8 +220,15 @@ pub struct VarRecord {
pub num_occurrences: usize,
}
pub type ClassifyFactResult = (Term, Vec<VarRecord>);
pub type ClassifyRuleResult = (Term, Vec<QueryTerm>, Vec<VarRecord>);
// TODO: already exists a VarData! although it may no longer exist??
// Also, the name is too similar to VarInfo. Think of better names!
pub struct VarData {
pub records: Vec<VarRecord>,
pub fixtures: VariableFixtures,
}
pub type ClassifyFactResult = (Term, VarData);
pub type ClassifyRuleResult = (Term, Vec<QueryTerm>, VarData);
fn merge_branch_seq<Iter: Iterator<Item = BranchInfo>>(branches: Iter) -> BranchInfo {
let mut branch_info = BranchInfo::new(BranchNumber::default());
@@ -238,19 +266,20 @@ fn flatten_into_disjunct(build_stack: &mut Vec<QueryTerm>, preceding_len: usize)
fn term_in_other_chunk(term: &Term) -> Option<bool> {
match term {
Term::Clause(_, name, terms) => Some(!ClauseType::is_inbuilt(name, terms.len())),
Term::Clause(_, name, terms) => Some(!ClauseType::is_inbuilt(*name, terms.len())),
Term::Literal(_, Literal::Atom(atom!("!"))) |
Term::Literal(_, Literal::Char('!')) => Some(false),
Term::Literal(_, Literal::Atom(name)) => Some(!ClauseType::is_inbuilt(name, 0)),
Term::Literal(_, Literal::Atom(name)) => Some(!ClauseType::is_inbuilt(*name, 0)),
Term::Var(..) => Some(true),
_ => None,
}
}
// returns true if the insertion of SetLastChunkType was the final push.
// expects that iter iterates over a conjunct of Terms in reverse order.
fn insert_set_last_chunk_type(
state_stack: &mut Vec<TraversalState>,
iter: impl Iterator<Item = TraversalState>,
mut iter: impl Iterator<Item = TraversalState>,
) -> bool {
let beg = state_stack.len();
let mut idx = beg;
@@ -269,6 +298,7 @@ fn insert_set_last_chunk_type(
if will_break {
state_stack.push(TraversalState::SetLastChunkType);
state_stack.push(traversal_st);
break;
} else {
state_stack.push(traversal_st);
@@ -356,15 +386,22 @@ impl VariableClassifier {
}
fn probe_body_term(&mut self, term: &Term, term_loc: GenContext) {
let mut classify_info = ClassifyInfo { arg_c: 0, arity: term.arity() };
// second arg is true to iterate the root, which may be a variable
for term_ref in breadth_first_iter(term, true) {
if let TermRef::Var(_, _, var_name) = term_ref {
self.probe_body_var(Var::from(var_name), term_loc);
if let TermRef::Var(lvl, _, var_name) = term_ref {
let var_info = VarInfo { var_ptr: VarPtr::from(&var_name), lvl, classify_info };
self.probe_body_var(var_name, term_loc, var_info);
}
if let Level::Shallow = term_ref.level() {
classify_info.arg_c += 1;
}
}
}
fn probe_body_var(&mut self, var_name: Var, chunk_type: ChunkType) {
fn probe_body_var(&mut self, var_name: Var, term_loc: GenContext, var_info: VarInfo) {
let branch_info_v = self.branch_map.entry(var_name)
.or_insert_with(|| vec![]);
@@ -387,11 +424,15 @@ impl VariableClassifier {
};
if needs_new_chunk {
branch_info.chunks.push(ChunkInfo::new(self.current_chunk_num));
branch_info.chunks.push(ChunkInfo {
chunk_num: self.current_chunk_num,
term_loc,
vars: vec![],
});
}
let chunk_info = branch_info.chunks.last_mut().unwrap();
chunk_info.vars.push(VarPtr::from(&var_name));
chunk_info.vars.push(var_info);
}
fn classify_head_variables(&mut self, term: &Term) -> Result<(), CompilationError> {
@@ -401,9 +442,14 @@ impl VariableClassifier {
_ => return Err(CompilationError::InvalidRuleHead),
}
let mut classify_info = ClassifyInfo {
arg_c: 0,
arity: term.arity(),
};
// false argument to breadth_first_iter because the root is not iterable.
for term_ref in breadth_first_iter(term, false) {
if let TermRef::Var(_, _, var_name) = term_ref {
if let TermRef::Var(lvl, _, var_name) = term_ref {
// the body of the if let here is an inlined
// "probe_head_var". note the difference between it
// and "probe_body_var".
@@ -420,11 +466,21 @@ impl VariableClassifier {
let needs_new_chunk = branch_info.chunks.is_empty();
if needs_new_chunk {
branch_info.chunks.push(ChunkInfo::new(self.current_chunk_num));
branch_info.chunks.push(ChunkInfo {
chunk_num: self.current_chunk_num,
term_loc: GenContext::Head,
vars: vec![]
});
}
let chunk_info = branch_info.chunks.last_mut().unwrap();
chunk_info.vars.push(VarPtr::from(&var_name));
let var_info = VarInfo { var_ptr: VarPtr::from(&var_name), classify_info, lvl };
chunk_info.vars.push(var_info);
}
if let Level::Shallow = term_ref.level() {
classify_info.arg_c += 1;
}
}
@@ -584,6 +640,8 @@ impl VariableClassifier {
}
}
Term::Clause(_, atom!(":"), mut terms) if terms.len() == 2 => {
let term_loc = chunk_type.to_gen_context(self.current_chunk_num);
let predicate_name = terms.pop().unwrap();
let module_name = terms.pop().unwrap();
@@ -592,7 +650,7 @@ impl VariableClassifier {
Term::Literal(_, Literal::Atom(module_name)),
Term::Literal(_, Literal::Atom(predicate_name)),
) => {
if !ClauseType::is_inbuilt(name, 0) {
if !ClauseType::is_inbuilt(predicate_name, 0) {
state_stack.push(TraversalState::IncrChunkNum);
}
@@ -649,6 +707,8 @@ impl VariableClassifier {
}
}
Term::Clause(cell, atom!("$call_with_inference_counting"), terms) if terms.len() == 1 => {
let term_loc = chunk_type.to_gen_context(self.current_chunk_num);
for term in terms.iter() {
self.probe_body_term(term, term_loc);
}
@@ -663,6 +723,8 @@ impl VariableClassifier {
state_stack.push(TraversalState::IncrChunkNum);
}
let term_loc = chunk_type.to_gen_context(self.current_chunk_num);
for term in terms.iter() {
self.probe_body_term(term, term_loc);
}
@@ -694,6 +756,7 @@ impl VariableClassifier {
),
);
}
_ => {
return Err(CompilationError::InadmissibleQueryTerm);
}
@@ -707,25 +770,18 @@ impl VariableClassifier {
}
impl BranchMap {
pub fn separate_and_classify_variables(&mut self) -> Vec<VarRecord> {
let mut var_num = 0usize;
let mut records = vec![];
pub fn separate_and_classify_variables(&mut self) -> VarData {
let mut var_num = 0usize;
let mut var_data = VarData {
records: vec![],
fixtures: VariableFixtures::new(),
};
for branches in self.values_mut() {
for branch in branches.iter_mut() {
let mut num_occurrences = 0;
let mut chunk_occurrences = vec![];
for chunk in branch.chunks.iter_mut() {
num_occurrences += chunk.vars.len();
for var in chunk.vars.iter_mut() {
var.set(Var::Generated(var_num));
}
chunk_occurrences.push(chunk.chunk_num);
}
let classification = if branch.chunks.len() > 1 {
VarClassification::Perm
} else {
@@ -739,13 +795,37 @@ impl BranchMap {
.unwrap_or(VarClassification::Void)
};
records.push(VarRecord { classification, chunk_occurrences, num_occurrences });
for chunk in branch.chunks.iter_mut() {
num_occurrences += chunk.vars.len();
if let VarClassification::Temp = classification {
for var_info in chunk.vars.iter_mut() {
var_info.var_ptr.set(Var::Generated(var_num));
var_data.fixtures.mark_temp_var(
var_num,
var_info.lvl,
&var_info.classify_info,
chunk.term_loc,
);
}
} else {
for var_info in chunk.vars.iter_mut() {
var_info.var_ptr.set(Var::Generated(var_num));
}
}
chunk_occurrences.push(chunk.chunk_num);
}
let record = VarRecord { classification, chunk_occurrences, num_occurrences };
var_data.records.push(record);
var_num += 1;
}
}
debug_assert_eq!(records.len(), var_num);
debug_assert_eq!(var_data.records.len(), var_num);
records
var_data
}
}