#include "DefinitionIntroduction.hpp"
#include "Kernel/Clause.hpp"
#include "Kernel/TermIterators.hpp"
#include "Kernel/InferenceStore.hpp"
#include "Lib/Metaiterators.hpp"
struct IncompleteFunction {
unsigned functor, arity, remaining;
bool typeCon;
};
static Term *lgg(Term *left, Term *right) {
ASS_EQ(left->functor(), right->functor())
if(left == right)
return left;
std::vector<TermList> args;
std::vector<IncompleteFunction> skeleton;
DHMap<std::pair<TermList, TermList>, unsigned> substitution;
unsigned fresh = 0;
SubtermIterator left_subterms(left);
SubtermIterator right_subterms(right);
while(left_subterms.hasNext()) {
ALWAYS(right_subterms.hasNext());
TermList left = left_subterms.next();
TermList right = right_subterms.next();
if(left.isTerm() && right.isTerm() && left.term()->functor() == right.term()->functor()) {
Term *leftt = left.term();
unsigned functor = leftt->functor();
unsigned arity = leftt->arity();
unsigned remaining = arity;
bool typeCon = leftt->isSort();
skeleton.push_back({functor, arity, remaining, typeCon});
}
else {
unsigned mapped;
if(!substitution.find({left, right}, mapped))
substitution.insert({left, right}, mapped = fresh++);
args.emplace_back(mapped, false);
left_subterms.right();
right_subterms.right();
if(!skeleton.empty())
skeleton.back().remaining--;
}
while(!skeleton.empty() && !skeleton.back().remaining) {
IncompleteFunction record = skeleton.back();
skeleton.pop_back();
size_t before = args.size() - record.arity;
Term *term = record.typeCon
? AtomicSort::create(record.functor, record.arity, args.data() + before)
: Term::create(record.functor, record.arity, args.data() + before);
args.resize(before);
args.emplace_back(term);
if(!skeleton.empty())
skeleton.back().remaining--;
else
break;
}
}
ASS_EQ(args.size(), left->arity());
return Term::create(left->functor(), left->arity(), args.data());
}
namespace Inferences
{
void DefinitionIntroduction::introduceDefinitionFor(Term *t) {
if(auto [_, inserted] = _defined.insert(t); !inserted)
return;
DHMap<unsigned, TermList> domain_sorts;
TermList range_sort = SortHelper::getResultSort(t);
SortHelper::collectVariableSorts(t, domain_sorts);
std::vector<TermList> domain_sort_vector;
std::vector<TermList> variables;
unsigned term_arity = 0, sort_arity = 0;
Renaming sort_rename;
for(auto [x, sort] : iterTraits(domain_sorts.items()))
if(sort == AtomicSort::superSort()) {
sort_arity++;
variables.emplace_back(x, false);
sort_rename.getOrBind(x);
}
for(auto [x, sort] : iterTraits(domain_sorts.items()))
if(sort != AtomicSort::superSort()) {
term_arity++;
variables.emplace_back(x, false);
domain_sort_vector.push_back(sort_rename.apply(sort));
}
unsigned functor = env.signature->addFreshFunction(term_arity + sort_arity, "sF");
OperatorType *type = OperatorType::getFunctionType(
term_arity,
domain_sort_vector.data(),
sort_rename.apply(range_sort),
sort_arity
);
env.signature->getFunction(functor)->setType(type);
Term *def = Term::create(functor, sort_arity + term_arity, variables.data());
Literal *eq = Literal::createEquality(true, TermList(def), TermList(t), range_sort);
auto definition = Clause::fromLiterals({eq}, NonspecificInference0(UnitInputType::AXIOM, InferenceRule::FUNCTION_DEFINITION));
InferenceStore::instance()->recordIntroducedSymbol(definition, SymbolType::FUNC, functor);
_definitions.push_back(definition);
}
void DefinitionIntroduction::process(Term *t) {
std::vector<Entry> &entries = _entries[t->functor()];
for(Entry &entry : entries) {
Term *gen = lgg(entry.term, t);
if(gen->allArgumentsAreVariables() && gen->getDistinctVars() == gen->arity())
continue;
entry.term = gen;
entry.weight += t->weight();
if(entry.weight > env.options->functionDefinitionIntroduction()) {
introduceDefinitionFor(entry.term);
std::swap(entry, entries.back());
entries.pop_back();
}
return;
}
entries.push_back({t, t->weight()});
}
void DefinitionIntroduction::process(Clause *cl) {
if(cl->inference().rule() == InferenceRule::FUNCTION_DEFINITION)
return;
while(_entries.size() < env.signature->functions())
_entries.emplace_back();
for(Literal *l : cl->iterLits())
for(Term *t : iterTraits(NonVariableNonTypeIterator(l)))
if(!t->allArgumentsAreVariables() || t->getDistinctVars() < t->arity())
process(t);
}
}