#include <algorithm>
#include "Lib/BinaryHeap.hpp"
#include "Lib/DArray.hpp"
#include "Lib/DHMap.hpp"
#include "Lib/Hash.hpp"
#include "Lib/TriangularArray.hpp"
#include "Clause.hpp"
#include "Matcher.hpp"
#include "Term.hpp"
#include "TermIterators.hpp"
#include "MLMatcher.hpp"
namespace {
using namespace std;
using namespace Lib;
using namespace Kernel;
typedef DHMap<unsigned,unsigned, IdentityHash, DefaultHash> UUMap;
struct ArrayStoringBinder
{
ArrayStoringBinder(TermList* arr, UUMap& v2pos)
: _arr(arr), _v2pos(v2pos) {}
bool bind(unsigned var, TermList term)
{
_arr[_v2pos.get(var)]=term;
return true;
}
void specVar(unsigned var, TermList term)
{ ASSERTION_VIOLATION; }
private:
TermList* _arr;
UUMap& _v2pos;
};
bool createLiteralBindings(Literal* baseLit, LiteralList const* alts, Clause* instCl, Literal* resolvedLit,
unsigned*& boundVarData, TermList**& altBindingPtrs, TermList*& altBindingData)
{
static UUMap variablePositions;
static BinaryHeap<unsigned,Int> varNums;
variablePositions.reset();
varNums.reset();
VariableIterator bvit(baseLit);
while(bvit.hasNext()) {
unsigned var=bvit.next().var();
varNums.insert(var);
}
unsigned nextPos=0;
while(!varNums.isEmpty()) {
unsigned var=varNums.pop();
while(!varNums.isEmpty() && varNums.top()==var) {
varNums.pop();
}
ALWAYS(variablePositions.insert(var, nextPos));
*(boundVarData++) = var;
nextPos++;
}
unsigned numVars=nextPos;
LiteralList::Iterator ait(alts);
while(ait.hasNext()) {
Literal* alit=ait.next();
if(alit==resolvedLit) {
continue;
}
if(alit->isEquality()) {
if(MatchingUtils::matchArgs(baseLit,alit)) {
ArrayStoringBinder binder(altBindingData, variablePositions);
MatchingUtils::matchArgs(baseLit,alit,binder);
*altBindingPtrs=altBindingData;
altBindingPtrs++;
altBindingData+=numVars;
if(resolvedLit) {
(altBindingData++)->setContent(0);
} else {
(altBindingData++)->setContent(instCl->getLiteralPosition(alit));
}
}
if(MatchingUtils::matchReversedArgs(baseLit, alit)) {
ArrayStoringBinder binder(altBindingData, variablePositions);
MatchingUtils::matchReversedArgs(baseLit, alit, binder);
*altBindingPtrs=altBindingData;
altBindingPtrs++;
altBindingData+=numVars;
if(resolvedLit) {
(altBindingData++)->setContent(0);
} else {
(altBindingData++)->setContent(instCl->getLiteralPosition(alit));
}
}
} else {
if(numVars) {
ArrayStoringBinder binder(altBindingData, variablePositions);
ALWAYS(MatchingUtils::matchArgs(baseLit,alit,binder));
}
*altBindingPtrs=altBindingData;
altBindingPtrs++;
altBindingData+=numVars;
if(resolvedLit) {
(altBindingData++)->setContent(0);
} else {
(altBindingData++)->setContent((uint64_t)instCl->getLiteralPosition(alit));
}
}
}
if(resolvedLit && resolvedLit->complementaryHeader()==baseLit->header()) {
if(!baseLit->arity() || MatchingUtils::matchArgs(baseLit,resolvedLit)) {
if(numVars) {
ArrayStoringBinder binder(altBindingData, variablePositions);
MatchingUtils::matchArgs(baseLit,resolvedLit,binder);
}
*altBindingPtrs=altBindingData;
altBindingPtrs++;
altBindingData+=numVars;
(altBindingData++)->setContent(1);
}
if(baseLit->isEquality() && MatchingUtils::matchReversedArgs(baseLit, resolvedLit)) {
ArrayStoringBinder binder(altBindingData, variablePositions);
MatchingUtils::matchReversedArgs(baseLit, resolvedLit, binder);
*altBindingPtrs=altBindingData;
altBindingPtrs++;
altBindingData+=numVars;
(altBindingData++)->setContent(1);
}
}
return true;
}
struct MatchingData {
unsigned len;
unsigned* varCnts;
unsigned** boundVarNums;
TermList*** altBindings;
TriangularArray<unsigned>* remaining;
unsigned* nextAlts;
TriangularArray<pair<int,int>* >* intersections;
Literal** bases;
LiteralList const* const* alts;
Clause* instance;
Literal* resolvedLit;
unsigned* boundVarNumStorage;
TermList** altBindingPtrStorage;
TermList* altBindingStorage;
pair<int,int>* intersectionStorage;
enum InitResult {
OK,
MUST_BACKTRACK,
NO_ALTERNATIVE
};
unsigned getRemainingInCurrent(unsigned bi) const
{
return remaining->get(bi,bi);
}
unsigned getAltRecordIndex(unsigned bi, unsigned alti) const
{
return static_cast<unsigned>(altBindings[bi][alti][varCnts[bi]].content());
}
bool compatible(unsigned b1Index, TermList* i1Bindings,
unsigned b2Index, unsigned i2AltIndex, pair<int,int>* iinfo) const
{
TermList* i2Bindings=altBindings[b2Index][i2AltIndex];
while(iinfo->first!=-1) {
if(i1Bindings[iinfo->first]!=i2Bindings[iinfo->second]) {
return false;
}
iinfo++;
}
return true;
}
bool bindAlt(unsigned bIndex, unsigned altIndex)
{
TermList* curBindings=altBindings[bIndex][altIndex];
for(unsigned i=bIndex+1; i<len; i++) {
if(!isInitialized(i)) {
break;
}
pair<int,int>* iinfo=getIntersectInfo(bIndex, i);
unsigned remAlts=remaining->get(i,bIndex);
if(iinfo->first!=-1) {
for(unsigned ai=0;ai<remAlts;ai++) {
if(!compatible(bIndex,curBindings,i,ai,iinfo)) {
remAlts--;
std::swap(altBindings[i][ai], altBindings[i][remAlts]);
ai--;
}
}
}
if(remAlts==0) {
return false;
}
remaining->set(i,bIndex+1,remAlts);
}
return true;
}
pair<int,int>* getIntersectInfo(unsigned b1, unsigned b2)
{
ASS_L(b1, b2);
pair<int,int>* res=intersections->get(b2,b1);
if( res ) {
return res;
}
intersections->set(b2,b1, intersectionStorage);
res=intersectionStorage;
unsigned b1vcnt=varCnts[b1];
unsigned b2vcnt=varCnts[b2];
unsigned* b1vn=boundVarNums[b1];
unsigned* b1vnStop=boundVarNums[b1]+b1vcnt;
unsigned* b2vn=boundVarNums[b2];
unsigned* b2vnStop=boundVarNums[b2]+b2vcnt;
int b1VarIndex=0;
int b2VarIndex=0;
while(true) {
while(b1vn!=b1vnStop && *b1vn<*b2vn) { b1vn++; b1VarIndex++; }
if(b1vn==b1vnStop) { break; }
while(b2vn!=b2vnStop && *b1vn>*b2vn) { b2vn++; b2VarIndex++; }
if(b2vn==b2vnStop) { break; }
if(*b1vn==*b2vn) {
intersectionStorage->first=b1VarIndex;
intersectionStorage->second=b2VarIndex;
intersectionStorage++;
b1vn++; b1VarIndex++;
b2vn++; b2VarIndex++;
if(b1vn==b1vnStop || b2vn==b2vnStop) { break; }
}
}
intersectionStorage->first=-1;
intersectionStorage++;
return res;
}
bool isInitialized(unsigned bIndex) const {
return boundVarNums[bIndex];
}
InitResult ensureInit(unsigned bIndex)
{
if(!isInitialized(bIndex)) {
boundVarNums[bIndex]=boundVarNumStorage;
altBindings[bIndex]=altBindingPtrStorage;
ALWAYS(createLiteralBindings(bases[bIndex], alts[bIndex], instance, resolvedLit,
boundVarNumStorage, altBindingPtrStorage, altBindingStorage));
varCnts[bIndex]=boundVarNumStorage-boundVarNums[bIndex];
unsigned altCnt=altBindingPtrStorage-altBindings[bIndex];
if(altCnt==0) {
return NO_ALTERNATIVE;
}
remaining->set(bIndex, 0, altCnt);
unsigned remAlts=0;
for(unsigned pbi=0;pbi<bIndex;pbi++) { pair<int,int>* iinfo=getIntersectInfo(pbi, bIndex);
remAlts=remaining->get(bIndex, pbi);
if(iinfo->first!=-1) {
TermList* pbBindings=altBindings[pbi][nextAlts[pbi]-1];
for(unsigned ai=0;ai<remAlts;ai++) {
if(!compatible(pbi, pbBindings, bIndex, ai, iinfo)) {
remAlts--;
std::swap(altBindings[bIndex][ai], altBindings[bIndex][remAlts]);
ai--;
}
}
}
remaining->set(bIndex,pbi+1,remAlts);
}
if(bIndex>0 && remAlts==0) {
return MUST_BACKTRACK;
}
}
return OK;
}
};
}
namespace Kernel
{
using namespace Lib;
class MLMatcher::Impl final
{
public:
USE_ALLOCATOR(MLMatcher::Impl);
Impl();
~Impl() = default;
void init(Literal** baseLits, unsigned baseLen, Clause* instance, LiteralList const* const* alts, Literal* resolvedLit, bool multiset);
bool nextMatch();
void getMatchedAltsBitmap(std::vector<bool>& outMatchedBitmap) const;
void getBindings(std::unordered_map<unsigned, TermList>& outBindings) const;
MLMatchStats getStats() const { return stats; }
Impl(Impl const&) = delete;
Impl(Impl&&) = delete;
Impl& operator=(Impl const&) = delete;
Impl& operator=(Impl&&) = delete;
private:
void initMatchingData(Literal** baseLits0, unsigned baseLen, Clause* instance, LiteralList const* const* alts, Literal* resolvedLit);
private:
DArray<Literal*> s_baseLits;
DArray<LiteralList const*> s_altsArr;
DArray<unsigned> s_varCnts;
DArray<unsigned*> s_boundVarNums;
DArray<TermList**> s_altPtrs;
TriangularArray<unsigned> s_remaining;
TriangularArray<pair<int,int>* > s_intersections;
DArray<unsigned> s_nextAlts;
DArray<unsigned> s_boundVarNumData;
DArray<TermList*> s_altBindingPtrs;
DArray<TermList> s_altBindingsData;
DArray<pair<int,int> > s_intersectionData;
MatchingData s_matchingData;
DArray<unsigned> s_matchRecord;
unsigned s_currBLit;
bool s_multiset;
MLMatchStats stats;
};
MLMatcher::Impl::Impl()
: s_baseLits(32)
, s_altsArr(32)
, s_varCnts(32)
, s_boundVarNums(32)
, s_altPtrs(32)
, s_remaining(32)
, s_intersections(32)
, s_nextAlts(32)
, s_boundVarNumData(64)
, s_altBindingPtrs(128)
, s_altBindingsData(256)
, s_intersectionData(128)
, s_matchRecord(32)
{ }
void MLMatcher::Impl::initMatchingData(Literal** baseLits0, unsigned baseLen, Clause* instance, LiteralList const* const* alts, Literal* resolvedLit)
{
s_baseLits.initFromArray(baseLen,baseLits0);
s_altsArr.initFromArray(baseLen,alts);
s_varCnts.ensure(baseLen);
s_boundVarNums.init(baseLen,0);
s_altPtrs.ensure(baseLen);
s_remaining.setSide(baseLen);
s_nextAlts.ensure(baseLen);
s_intersections.setSide(baseLen);
s_intersections.zeroAll();
unsigned zeroAlts=0;
unsigned singleAlts=0;
size_t baseLitVars=0;
size_t altCnt=0;
size_t altBindingsCnt=0;
unsigned mostDistVarsLit=0;
unsigned mostDistVarsCnt=s_baseLits[0]->getDistinctVars();
auto swapLits = [this] (int i, int j) {
std::swap(s_baseLits[i], s_baseLits[j]);
std::swap(s_altsArr[i], s_altsArr[j]);
};
for(unsigned i=0;i<baseLen;i++) {
unsigned distVars=s_baseLits[i]->getDistinctVars();
baseLitVars+=distVars;
unsigned currAltCnt=0;
LiteralList::Iterator ait(s_altsArr[i]);
while(ait.hasNext()) {
currAltCnt++;
if(ait.next()->isEquality()) {
currAltCnt++;
}
}
altCnt+=currAltCnt+2; altBindingsCnt+=(distVars+1)*(currAltCnt+2);
ASS_LE(zeroAlts, singleAlts);
ASS_LE(singleAlts, i);
if(currAltCnt==0) {
if(zeroAlts!=i) {
if(singleAlts!=zeroAlts) {
swapLits(singleAlts, zeroAlts);
}
swapLits(i, zeroAlts);
if(mostDistVarsLit==singleAlts) {
mostDistVarsLit=i;
}
}
zeroAlts++;
singleAlts++;
} else if(currAltCnt==1 && !(resolvedLit && resolvedLit->couldBeInstanceOf(s_baseLits[i], true)) ) {
if(singleAlts!=i) {
swapLits(i, singleAlts);
if(mostDistVarsLit==singleAlts) {
mostDistVarsLit=i;
}
}
singleAlts++;
} else if(i>0 && mostDistVarsCnt<distVars) {
mostDistVarsLit=i;
mostDistVarsCnt=distVars;
}
}
if(mostDistVarsLit>singleAlts) {
swapLits(mostDistVarsLit, singleAlts);
}
s_boundVarNumData.ensure(baseLitVars);
s_altBindingPtrs.ensure(altCnt);
s_altBindingsData.ensure(altBindingsCnt);
s_intersectionData.ensure((baseLitVars+baseLen)*baseLen);
s_matchingData.len=baseLen;
s_matchingData.varCnts=s_varCnts.array();
s_matchingData.boundVarNums=s_boundVarNums.array();
s_matchingData.altBindings=s_altPtrs.array();
s_matchingData.remaining=&s_remaining;
s_matchingData.nextAlts=s_nextAlts.array();
s_matchingData.intersections=&s_intersections;
s_matchingData.bases=s_baseLits.array();
s_matchingData.alts=s_altsArr.array();
s_matchingData.instance=instance;
s_matchingData.resolvedLit=resolvedLit;
s_matchingData.boundVarNumStorage=s_boundVarNumData.array();
s_matchingData.altBindingPtrStorage=s_altBindingPtrs.array();
s_matchingData.altBindingStorage=s_altBindingsData.array();
s_matchingData.intersectionStorage=s_intersectionData.array();
}
void MLMatcher::Impl::init(Literal** baseLits, unsigned baseLen, Clause* instance, LiteralList const* const* alts, Literal* resolvedLit, bool multiset)
{
if (resolvedLit) {
ASS(!multiset);
}
initMatchingData(baseLits, baseLen, instance, alts, resolvedLit);
unsigned matchRecordLen = resolvedLit ? 2 : instance->length();
s_matchRecord.init(matchRecordLen, 0xFFFFFFFF);
ASS_EQ(s_matchRecord.size(), matchRecordLen);
s_matchingData.nextAlts[0] = 0;
s_currBLit = 0;
s_multiset = multiset;
stats = MLMatchStats{};
}
bool MLMatcher::Impl::nextMatch()
{
MatchingData* const md = &s_matchingData;
while (true) {
MatchingData::InitResult ires = md->ensureInit(s_currBLit);
if (ires != MatchingData::OK) {
if (ires == MatchingData::MUST_BACKTRACK) {
s_currBLit--;
continue;
} else {
ASS_EQ(ires, MatchingData::NO_ALTERNATIVE);
return false;
}
}
unsigned maxAlt = md->getRemainingInCurrent(s_currBLit);
while (md->nextAlts[s_currBLit] < maxAlt &&
(
( s_multiset && s_matchRecord[md->getAltRecordIndex(s_currBLit, md->nextAlts[s_currBLit])] < s_currBLit )
|| !md->bindAlt(s_currBLit, md->nextAlts[s_currBLit])
)
) {
md->nextAlts[s_currBLit]++;
}
if (md->nextAlts[s_currBLit] < maxAlt) {
if (md->nextAlts[s_currBLit] < maxAlt - 1) {
stats.numDecisions += 1;
}
unsigned matchRecordIndex=md->getAltRecordIndex(s_currBLit, md->nextAlts[s_currBLit]);
for (unsigned i = 0; i < s_matchRecord.size(); i++) {
if (s_matchRecord[i] == s_currBLit) {
s_matchRecord[i]=0xFFFFFFFF;
}
}
ASS(!s_multiset || s_matchRecord[matchRecordIndex]>s_currBLit); if (s_matchRecord[matchRecordIndex]>s_currBLit) {
s_matchRecord[matchRecordIndex]=s_currBLit;
}
md->nextAlts[s_currBLit]++;
s_currBLit++;
if(s_currBLit == md->len) {
if(md->resolvedLit && s_matchRecord[1] >= md->len) {
s_currBLit--;
continue;
}
s_currBLit--; stats.result = true;
return true;
}
md->nextAlts[s_currBLit]=0;
} else {
ASS_GE(md->nextAlts[s_currBLit], maxAlt);
if(s_currBLit==0) { return false; }
s_currBLit--;
}
}
ASSERTION_VIOLATION; }
void MLMatcher::Impl::getMatchedAltsBitmap(std::vector<bool>& outMatchedBitmap) const
{
MatchingData const* const md = &s_matchingData;
ASS(!md->resolvedLit);
outMatchedBitmap.clear();
outMatchedBitmap.resize(md->instance->length(), false);
for (unsigned bi = 0; bi < md->len; ++bi) {
unsigned alti = md->nextAlts[bi] - 1;
unsigned i = md->getAltRecordIndex(bi, alti);
outMatchedBitmap[i] = true;
}
}
void MLMatcher::Impl::getBindings(std::unordered_map<unsigned, TermList>& outBindings) const
{
MatchingData const* const md = &s_matchingData;
ASS(!md->resolvedLit);
ASS(outBindings.empty());
for (unsigned bi = 0; bi < md->len; ++bi) {
unsigned alti = md->nextAlts[bi] - 1;
for (unsigned vi = 0; vi < md->varCnts[bi]; ++vi) {
unsigned var = md->boundVarNums[bi][vi];
TermList trm = md->altBindings[bi][alti][vi];
DEBUG_CODE(auto res =) outBindings.insert({var, trm});
#if VDEBUG
auto it = res.first;
bool inserted = res.second;
if (!inserted) {
ASS_EQ(it->second, trm);
}
#endif
}
}
}
MLMatcher::MLMatcher()
{
m_impl = std::make_unique<MLMatcher::Impl>();
}
void MLMatcher::init(Literal** baseLits, unsigned baseLen, Clause* instance, LiteralList const* const* alts, Literal* resolvedLit, bool multiset)
{
ASS(m_impl);
m_impl->init(baseLits, baseLen, instance, alts, resolvedLit, multiset);
}
MLMatcher::~MLMatcher() = default;
bool MLMatcher::nextMatch()
{
ASS(m_impl);
return m_impl->nextMatch();
}
void MLMatcher::getMatchedAltsBitmap(std::vector<bool>& outMatchedBitmap) const
{
ASS(m_impl);
m_impl->getMatchedAltsBitmap(outMatchedBitmap);
}
void MLMatcher::getBindings(std::unordered_map<unsigned, TermList>& outBindings) const
{
ASS(m_impl);
m_impl->getBindings(outBindings);
}
MLMatchStats MLMatcher::getStats() const
{
ASS(m_impl);
return m_impl->getStats();
}
static MLMatcher matcher;
bool MLMatcher::canBeMatched(Literal** baseLits, unsigned baseLen, Clause* instance, LiteralList const* const* alts, Literal* resolvedLit, bool multiset)
{
matcher.init(baseLits, baseLen, instance, alts, resolvedLit, multiset);
return matcher.nextMatch();
}
MLMatchStats MLMatcher::getStaticStats()
{
return matcher.getStats();
}
}