remove BuildIf, BuildNot, BuildThen TermIterState variants

This commit is contained in:
Mark Thom
2023-01-30 23:26:50 -07:00
committed by Mark
parent cb59c3003a
commit b205abe949
4 changed files with 112 additions and 84 deletions

View File

@@ -139,27 +139,22 @@ impl VariableFixtures {
}; };
} }
pub(crate) fn mark_temp_var( pub(crate) fn mark_temp_var(&mut self, var_info: &VarInfo) {
&mut self,
generated_var_index: usize,
lvl: Level,
classify_info: &ClassifyInfo,
term_loc: GenContext,
) {
let chunk_num = term_loc.chunk_num(); let chunk_num = term_loc.chunk_num();
let var = Var::from(var_info.var_ptr);
let mut status = self.temp_vars.swap_remove(&generated_var_index).unwrap_or_else(|| { let mut status = self.temp_vars.swap_remove(&var).unwrap_or_else(|| {
TempVarStatus { TempVarStatus {
chunk_num, chunk_num,
temp_var_data: TempVarData::new(classify_info.arity), temp_var_data: TempVarData::new(var_info.classify_info.arity),
} }
}); });
if let Level::Shallow = lvl { if let Level::Shallow = var_info.lvl {
self.record_temp_info(&mut status, classify_info.arg_c, term_loc); self.record_temp_info(&mut status, var_info.classify_info.arg_c, term_loc);
} }
self.temp_vars.insert(Var::Generated(generated_var_index), status); self.temp_vars.insert(var, status);
} }
} }

View File

@@ -106,10 +106,10 @@ impl ChunkType {
pub enum QueryTerm { pub enum QueryTerm {
// register, clause type, subterms, clause call policy. // register, clause type, subterms, clause call policy.
Clause(Cell<RegType>, ClauseType, Vec<Term>, CallPolicy), Clause(Cell<RegType>, ClauseType, Vec<Term>, CallPolicy),
Cut, Fail,
Not(Vec<QueryTerm>), GlobalCut,
IfThen(Vec<QueryTerm>, Vec<QueryTerm>), GetCutPoint(usize),
LocalCut(Cell<VarReg>), // for IfThen. LocalCut(usize),
Branch(Vec<Vec<QueryTerm>>), Branch(Vec<Vec<QueryTerm>>),
ChunkTypeBoundary(ChunkType), ChunkTypeBoundary(ChunkType),
} }

View File

@@ -12,8 +12,9 @@ use std::iter::*;
use std::vec::Vec; use std::vec::Vec;
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)] #[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
pub(crate) struct VarPtr { pub(crate) enum VarPtr {
ptr: std::ptr::NonNull<Var>, ToVar(std::ptr::NonNull<Var>),
InSitu(usize),
} }
impl From<&Var> for VarPtr { impl From<&Var> for VarPtr {
@@ -26,17 +27,27 @@ impl From<&Var> for VarPtr {
} }
impl From<VarPtr> for Var { impl From<VarPtr> for Var {
#[inline] #[inline(always)]
fn from(value: VarPtr) -> Var { fn from(value: VarPtr) -> Var {
unsafe { match value {
(*value.ptr.as_ptr()).clone() VarPtr::ToPtr(ptr) => unsafe {
(*ptr.ptr.as_ptr()).clone()
},
VarPtr::InSitu(var_num) => {
Var::Generated(var_num)
}
} }
} }
} }
impl VarPtr { impl VarPtr {
pub(crate) fn set(&mut self, value: Var) { pub(crate) fn set(&mut self, value: Var) {
unsafe { *self.ptr.as_mut() = value; } match self {
VarPtr::ToVar(ref mut ptr) =>
unsafe { *ptr.as_mut() = value },
VarPtr::InSitu(_) => {
}
}
} }
} }

View File

@@ -150,9 +150,9 @@ enum TraversalState {
// add the last disjunct to a QueryTerm::Branch, continuing from // add the last disjunct to a QueryTerm::Branch, continuing from
// where it leaves off. // where it leaves off.
BuildFinalDisjunct(usize), BuildFinalDisjunct(usize),
BuildIf(usize, Term), // build the P term of P -> Q Fail,
BuildThen(usize, Vec<QueryTerm>), // build the Q term of P -> Q GetCutPoint(usize),
BuildNot(usize), // build the P term of \+ P LocalCut(usize),
ResetCallPolicy(CallPolicy), ResetCallPolicy(CallPolicy),
Term(Term), Term(Term),
AddBranchNum(BranchNumber), // set current_branch_number, add it to the root set AddBranchNum(BranchNumber), // set current_branch_number, add it to the root set
@@ -186,6 +186,7 @@ pub struct VariableClassifier {
current_branch_num: BranchNumber, current_branch_num: BranchNumber,
current_chunk_num: usize, current_chunk_num: usize,
branch_map: BranchMap, branch_map: BranchMap,
var_num: usize,
root_set: RootSet, root_set: RootSet,
} }
@@ -196,12 +197,23 @@ pub enum VarClassification {
Perm, Perm,
} }
#[derive(Clone, Debug)]
pub struct VarRecord { pub struct VarRecord {
pub classification: VarClassification, pub classification: VarClassification,
pub chunk_occurrences: Vec<usize>, pub chunk_occurrences: Vec<usize>,
pub num_occurrences: usize, pub num_occurrences: usize,
} }
impl Default for VarRecord {
fn default() -> Self {
VarRecord {
classification: VarClassification::Void,
chunk_occurrences: vec![],
num_occurrences: 0,
}
}
}
pub struct VarData { pub struct VarData {
pub records: Vec<VarRecord>, pub records: Vec<VarRecord>,
pub fixtures: VariableFixtures, pub fixtures: VariableFixtures,
@@ -269,7 +281,7 @@ fn insert_set_last_chunk_type(
while let Some(traversal_st) = iter.next() { while let Some(traversal_st) = iter.next() {
match traversal_st { match traversal_st {
TraversalState::Term(term) | TraversalState::BuildIf(_, term) => { TraversalState::Term(term) => {
will_break = false; will_break = false;
match term_in_other_chunk(&term) { match term_in_other_chunk(&term) {
@@ -288,7 +300,7 @@ fn insert_set_last_chunk_type(
} }
} }
_ => { _ => {
unreachable!(); state_stack.push(traversal_st);
} }
} }
} }
@@ -305,12 +317,13 @@ impl VariableClassifier {
current_chunk_num: 0, current_chunk_num: 0,
branch_map: BranchMap(BranchMapInt::new()), branch_map: BranchMap(BranchMapInt::new()),
root_set: RootSet::new(), root_set: RootSet::new(),
var_num: 0,
} }
} }
pub fn classify_fact(mut self, term: Term) -> Result<ClassifyFactResult, CompilationError> { pub fn classify_fact(mut self, term: Term) -> Result<ClassifyFactResult, CompilationError> {
self.classify_head_variables(&term)?; self.classify_head_variables(&term)?;
Ok((term, self.branch_map.separate_and_classify_variables())) Ok((term, self.branch_map.separate_and_classify_variables(self.var_num)))
} }
pub fn classify_rule<'a, LS: LoadState<'a>>( pub fn classify_rule<'a, LS: LoadState<'a>>(
@@ -322,7 +335,7 @@ impl VariableClassifier {
self.classify_head_variables(&head)?; self.classify_head_variables(&head)?;
let query_terms = self.classify_body_variables(loader, body)?; let query_terms = self.classify_body_variables(loader, body)?;
Ok((head, query_terms, self.branch_map.separate_and_classify_variables())) Ok((head, query_terms, self.branch_map.separate_and_classify_variables(self.var_num)))
} }
fn merge_branches(&mut self) { fn merge_branches(&mut self) {
@@ -396,6 +409,20 @@ impl VariableClassifier {
chunk_info.vars.push(var_info); chunk_info.vars.push(var_info);
} }
fn probe_in_situ_var(&mut self, chunk_type: ChunkType, var_num: usize) {
let classify_info = ClassifyInfo { arg_c: 0, arity: 0 };
let var_info = VarInfo {
var_ptr: VarPtr::InSitu(var_num),
classify_info,
lvl: Level::Shallow,
};
let term_loc = chunk_type.to_gen_context(self.current_chunk_num);
self.probe_body_var(Var::Generated(var_num), term_loc, var_info);
}
fn classify_head_variables(&mut self, term: &Term) -> Result<(), CompilationError> { fn classify_head_variables(&mut self, term: &Term) -> Result<(), CompilationError> {
match term { match term {
Term::Clause(..) | Term::Literal(_, Literal::Atom(_)) => { Term::Clause(..) | Term::Literal(_, Literal::Atom(_)) => {
@@ -403,10 +430,7 @@ impl VariableClassifier {
_ => return Err(CompilationError::InvalidRuleHead), _ => return Err(CompilationError::InvalidRuleHead),
} }
let mut classify_info = ClassifyInfo { let mut classify_info = ClassifyInfo { arg_c: 0, arity: term.arity() };
arg_c: 0,
arity: term.arity(),
};
// false argument to breadth_first_iter because the root is not iterable. // false argument to breadth_first_iter because the root is not iterable.
for term_ref in breadth_first_iter(term, false) { for term_ref in breadth_first_iter(term, false) {
@@ -491,19 +515,20 @@ impl VariableClassifier {
TraversalState::BuildFinalDisjunct(preceding_len) => { TraversalState::BuildFinalDisjunct(preceding_len) => {
flatten_into_disjunct(&mut build_stack, preceding_len); flatten_into_disjunct(&mut build_stack, preceding_len);
} }
TraversalState::BuildIf(preceding_len, then_term) => { TraversalState::GetCutPoint(var_num) => {
let iter = build_stack.drain(preceding_len ..); let term_loc = chunk_type.to_gen_context(self.current_chunk_num);
state_stack.push(TraversalState::BuildThen(preceding_len, iter.collect())); self.probe_in_situ_var(term_loc, var_num);
state_stack.push(TraversalState::Term(then_term)); build_stack.push(QueryTerm::GetCutPoint(var_num));
} }
TraversalState::BuildThen(preceding_len, if_terms) => { TraversalState::LocalCut(var_num) => {
let iter = build_stack.drain(preceding_len ..); let term_loc = chunk_type.to_gen_context(self.current_chunk_num);
build_stack.push(QueryTerm::IfThen(if_terms, iter.collect()));
self.probe_in_situ_var(term_loc, var_num);
build_stack.push(QueryTerm::LocalCut(var_num));
} }
TraversalState::BuildNot(preceding_len) => { TraversalState::Fail => {
let iter = build_stack.drain(preceding_len ..); build_stack.push(QueryTerm::Fail);
build_stack.push(QueryTerm::Not(iter.collect()));
} }
TraversalState::Term(term) => { TraversalState::Term(term) => {
match term { match term {
@@ -567,17 +592,14 @@ impl VariableClassifier {
let then_term = terms.pop().unwrap(); let then_term = terms.pop().unwrap();
let if_term = terms.pop().unwrap(); let if_term = terms.pop().unwrap();
let build_stack_len = build_stack.len(); let iter = vec![TraversalState::Term(then_term),
TraversalState::LocalCut(self.var_num),
// TODO: insert GetCutPoint between TraversalState::Term(if_term),
// the two traversal states and detect TraversalState::GetCutPoint(self.var_num)]
// that as a chunk boundary in
// insert_set_last_chunk_type ??
let iter = vec![TraversalState::BuildIf(build_stack_len, then_term),
TraversalState::Term(if_term)]
.into_iter(); .into_iter();
self.var_num += 1;
if ChunkType::Mid != chunk_type { if ChunkType::Mid != chunk_type {
if insert_set_last_chunk_type(&mut state_stack, iter) { if insert_set_last_chunk_type(&mut state_stack, iter) {
if chunk_type.is_last() { if chunk_type.is_last() {
@@ -587,10 +609,12 @@ impl VariableClassifier {
} }
} }
Term::Clause(_, atom!("\\+"), terms) if terms.len() == 1 => { Term::Clause(_, atom!("\\+"), terms) if terms.len() == 1 => {
let build_stack_len = build_stack.len(); state_stack.push(TraversalState::Fail);
state_stack.push(TraversalState::LocalCut(self.var_num));
state_stack.push(TraversalState::BuildNot(build_stack_len));
state_stack.push(TraversalState::Term(terms[0])); state_stack.push(TraversalState::Term(terms[0]));
state_stack.push(TraversalState::GetCutPoint(self.var_num));
self.var_num += 1;
} }
Term::Clause(_, atom!(":"), mut terms) if terms.len() == 2 => { Term::Clause(_, atom!(":"), mut terms) if terms.len() == 2 => {
let term_loc = chunk_type.to_gen_context(self.current_chunk_num); let term_loc = chunk_type.to_gen_context(self.current_chunk_num);
@@ -692,7 +716,7 @@ impl VariableClassifier {
); );
} }
Term::Literal(_, Literal::Atom(atom!("!")) | Literal::Char('!')) => { Term::Literal(_, Literal::Atom(atom!("!")) | Literal::Char('!')) => {
build_stack.push(QueryTerm::Cut); build_stack.push(QueryTerm::GlobalCut);
} }
Term::Literal(cell, Literal::Atom(name)) => { Term::Literal(cell, Literal::Atom(name)) => {
if !ClauseType::is_inbuilt(name, 0) { if !ClauseType::is_inbuilt(name, 0) {
@@ -722,19 +746,25 @@ impl VariableClassifier {
} }
impl BranchMap { impl BranchMap {
pub fn separate_and_classify_variables(&mut self) -> VarData { pub fn separate_and_classify_variables(&mut self, mut var_num: usize) -> VarData {
let mut var_num = 0usize;
let mut var_data = VarData { let mut var_data = VarData {
records: vec![], records: vec![VarRecord::default(); self.len()],
fixtures: VariableFixtures::new(), fixtures: VariableFixtures::new(),
}; };
for branches in self.values_mut() { for (var, branches) in self.iter_mut() {
for branch in branches.iter_mut() { for branch in branches.iter_mut() {
let mut num_occurrences = 0; let mut num_occurrences = 0;
let mut chunk_occurrences = vec![];
let classification = if branch.chunks.len() > 1 { let idx = if let Var::Generated(var_num) = var {
*var_num
} else {
var_num += 1;
var_num - 1
};
var_data.records[idx].classification =
if branch.chunks.len() > 1 {
VarClassification::Perm VarClassification::Perm
} else { } else {
branch.chunks branch.chunks
@@ -747,18 +777,15 @@ impl BranchMap {
.unwrap_or(VarClassification::Void) .unwrap_or(VarClassification::Void)
}; };
var_data.records[idx].chunk_occurrences.reserve(branch.chunks.len());
for chunk in branch.chunks.iter_mut() { for chunk in branch.chunks.iter_mut() {
num_occurrences += chunk.vars.len(); var_data.records[idx].num_occurrences += chunk.vars.len();
if let VarClassification::Temp = classification { if let VarClassification::Temp = classification {
for var_info in chunk.vars.iter_mut() { for var_info in chunk.vars.iter_mut() {
var_info.var_ptr.set(Var::Generated(var_num)); var_info.var_ptr.set(Var::Generated(var_num));
var_data.fixtures.mark_temp_var( var_data.fixtures.mark_temp_var(&var_info);
var_num,
var_info.lvl,
&var_info.classify_info,
chunk.term_loc,
);
} }
} else { } else {
for var_info in chunk.vars.iter_mut() { for var_info in chunk.vars.iter_mut() {
@@ -766,13 +793,8 @@ impl BranchMap {
} }
} }
chunk_occurrences.push(chunk.chunk_num); var_data.records[idx].chunk_occurrences.push(chunk.chunk_num);
} }
let record = VarRecord { classification, chunk_occurrences, num_occurrences };
var_data.records.push(record);
var_num += 1;
} }
} }