From 38e71db7e89b71406ec73a1f199373b688159658 Mon Sep 17 00:00:00 2001 From: Fedor Isakov Date: Sat, 28 Sep 2019 12:57:40 +0200 Subject: [PATCH] Revive and refactor ReteNetwork-based rule matching algorithm. ReteNetwork is faster on large samples. RuleMatchingProbe is mutable and not persistent, no need to update value. --- .../mps/logic/reactor/core/Dispatcher.kt | 18 +- .../mps/logic/reactor/core/RuleMatcher.kt | 14 +- .../logic/reactor/core/RuleMatchingProbe.kt | 8 +- .../core/internal/ReteRuleMatcherImpl.kt | 385 +++++++++--------- .../reactor/core/internal/RuleMatcherImpl.kt | 79 ++-- 5 files changed, 278 insertions(+), 226 deletions(-) diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/Dispatcher.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/Dispatcher.kt index f4436aff..15451c52 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/Dispatcher.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/Dispatcher.kt @@ -16,8 +16,6 @@ package jetbrains.mps.logic.reactor.core -import jetbrains.mps.logic.reactor.util.Profiler - typealias DispatchingFrontState = Map @@ -71,7 +69,9 @@ class Dispatcher (val ruleIndex: RuleIndex) { private constructor(pred: DispatchingFront, matching: Iterable) { this.ruletag2probe = pred.ruletag2probe matching.forEach { probe -> - ruletag2probe[probe.rule().uniqueTag()] = probe + if (RULE_MATCHER_PROBE_PERSISTENT) { + ruletag2probe[probe.rule().uniqueTag()] = probe + } allMatches.addAll(probe.matches()) } } @@ -126,8 +126,10 @@ class Dispatcher (val ruleIndex: RuleIndex) { */ internal fun consume(consumedMatch: RuleMatchEx): DispatchingFront { ruletag2probe[consumedMatch.rule().uniqueTag()]?.let { - ruletag2probe[consumedMatch.rule().uniqueTag()] = - it.consume(consumedMatch) + val probe = it.consume(consumedMatch) + if (RULE_MATCHER_PROBE_PERSISTENT) { + ruletag2probe[consumedMatch.rule().uniqueTag()] = probe + } } return DispatchingFront(this) } @@ -137,8 +139,10 @@ class Dispatcher (val ruleIndex: RuleIndex) { */ internal fun forget(consumedMatch: RuleMatchEx): DispatchingFront { ruletag2probe[consumedMatch.rule().uniqueTag()]?.let { - ruletag2probe[consumedMatch.rule().uniqueTag()] = - it.forget(consumedMatch) + val probe = it.forget(consumedMatch) + if (RULE_MATCHER_PROBE_PERSISTENT) { + ruletag2probe[consumedMatch.rule().uniqueTag()] = probe + } } return DispatchingFront(this) } diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleMatcher.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleMatcher.kt index 688527c4..5fd3bccb 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleMatcher.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleMatcher.kt @@ -16,6 +16,9 @@ package jetbrains.mps.logic.reactor.core +import gnu.trove.list.TIntList +import gnu.trove.list.array.TIntArrayList +import jetbrains.mps.logic.reactor.core.internal.ReteRuleMatcherImpl import jetbrains.mps.logic.reactor.core.internal.RuleMatcherImpl import jetbrains.mps.logic.reactor.program.Rule @@ -33,5 +36,14 @@ interface RuleMatcher { } -fun createRuleMatcher(lookup: RuleLookup, tag: Any): RuleMatcher = RuleMatcherImpl(lookup, tag) +fun createRuleMatcher(lookup: RuleLookup, tag: Any): RuleMatcher = ReteRuleMatcherImpl(lookup, tag) + +val RULE_MATCHER_PROBE_PERSISTENT = false + +// Trove stuff +typealias Signature = TIntList + +fun Signature.copy() = TIntArrayList(this) + +fun IntArray.toSignature() = TIntArrayList(this) diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleMatchingProbe.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleMatchingProbe.kt index d8516687..9e587a44 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleMatchingProbe.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleMatchingProbe.kt @@ -18,6 +18,7 @@ package jetbrains.mps.logic.reactor.core import jetbrains.mps.logic.reactor.evaluation.RuleMatchingProbeState import jetbrains.mps.logic.reactor.program.Rule +import jetbrains.mps.logic.reactor.util.Profiler import java.util.BitSet @@ -38,10 +39,15 @@ interface RuleMatchingProbe : RuleMatchingProbeState { fun expand(occ: Occurrence): RuleMatchingProbe - fun expand(occ: Occurrence, mask: BitSet): RuleMatchingProbe + fun expand(occ: Occurrence, mask: BitSet, profiler: Profiler? = null): RuleMatchingProbe fun contract(occ: Occurrence): RuleMatchingProbe + /** + * The purpose and usages of this method are obscure. + * One of the implementations is a NOP. + */ + @Deprecated("hacky stuff") fun forgetSeen(occ: Occurrence): RuleMatchingProbe fun forgetConsumed(occ: Occurrence): RuleMatchingProbe diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/internal/ReteRuleMatcherImpl.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/internal/ReteRuleMatcherImpl.kt index 632bfdbb..76b0b865 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/internal/ReteRuleMatcherImpl.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/internal/ReteRuleMatcherImpl.kt @@ -18,27 +18,30 @@ package jetbrains.mps.logic.reactor.core.internal import jetbrains.mps.logic.reactor.core.* import jetbrains.mps.logic.reactor.evaluation.ConstraintOccurrence -import jetbrains.mps.logic.reactor.logical.MetaLogical import jetbrains.mps.logic.reactor.program.Rule -import jetbrains.mps.logic.reactor.util.allSetBits -import jetbrains.mps.logic.reactor.util.bitSet -import jetbrains.mps.logic.reactor.util.bitSetOfOnes -import jetbrains.mps.logic.reactor.util.copyApply +import jetbrains.mps.logic.reactor.util.* import java.util.* +import kotlin.collections.ArrayList /** - * An alternative implementation of RuleMatcherImpl. Has similar asymptotic characteristics as the default implementation, - * but in practice is a bit slower. + * An alternative implementation of RuleMatcherImpl. Has similar asymptotic characteristics as the "naïve" implementation. * * Loosely based on "Rete network" algorithm. * + * The implementation of [RuleMatchingProbe] returned from [probe] method is not a persistent object, it's rather + * a mutable object which updates its state through usual update methods that all return the same object. + * * @author Fedor Isakov */ -internal class ReteRuleMatcherImpl(val rule: Rule) : RuleMatcher { +internal class ReteRuleMatcherImpl(private val ruleLookup: RuleLookup, + private val tag: Any) : RuleMatcher +{ - val head = (rule.headKept() + rule.headReplaced()).toList() + val head = (lookupRule().headKept() + lookupRule().headReplaced()).toList() - val propagation = rule.headReplaced().count() == 0 + val propagation = lookupRule().headReplaced().count() == 0 + + fun lookupRule(): Rule = ruleLookup.lookupRuleByTag(tag) ?: throw IllegalStateException("can't lookup rule by tag: '${tag}'") override fun probe(): ReteNetwork = ReteNetwork(head.size) @@ -49,65 +52,40 @@ internal class ReteRuleMatcherImpl(val rule: Rule) : RuleMatcher { assert(headSize > 0) } - val skipOccIndices = BitSet() - val seenOcc2Idx = IdentityHashMap() var nextOccIdx: Int = 0 - val meta2Idx = HashMap, Int>() + var lastGeneration = Generation(arrayListOf(Layer(headSize, InitialNode()))) - var nextMetaIdx: Int = 0 - - val generations = - arrayListOf(Generation( - arrayListOf(Layer(headSize, InitialNode(BitSet()))))) + val consumedSignatures = HashSet() abstract inner class ReteNode { - fun isSubstCompatible (that: ReteNode) : Boolean { - val thisMetaIndices = this.getMeta() - val thatMetaIndices = that.getMeta() - if (thisMetaIndices == null || thatMetaIndices == null) return true + abstract fun subst(): Subst - if (thisMetaIndices.cardinality() != 0 && thatMetaIndices.cardinality() != 0) { - val it = thisMetaIndices.copyApply { and(thatMetaIndices) }.allSetBits() - while (it.hasNext()) { - val shared = it.next() - if (!createOccurrenceMatcher().match(this.getSubst(shared), that.getSubst(shared))) return false - } - } + abstract fun occupiesHeadPosition(headPos: Int) : Boolean - return true - } + abstract fun containsOccurrence(occIdx: Int) : Boolean - abstract fun isCompatible(that: AlphaNode) : Boolean + abstract fun combine(that: AlphaNode, subst: Subst): ReteNode - abstract fun containsOccurrence(skipOccIndices: BitSet) : Boolean - - abstract fun getMeta(): BitSet? - - abstract fun getSubst(metaIdx: Int): Any? - - abstract fun combine(that: AlphaNode): ReteNode - - open fun collectData(occArray: Array, allSubst: Subst): Subst = allSubst + abstract fun collect(occArray: Array) } - inner class InitialNode(val occIndices: BitSet) : ReteNode() { + inner class InitialNode() : ReteNode() { - override fun getMeta(): BitSet? = null + override fun subst(): Subst = emptySubst() - override fun getSubst(metaIdx: Int): Any? = null + override fun occupiesHeadPosition(headPos: Int): Boolean = false - override fun isCompatible(that: AlphaNode): Boolean = !occIndices[that.occIdx] + override fun containsOccurrence(occIdx: Int): Boolean = false - override fun containsOccurrence(skipOccIndices: BitSet) = - !this.occIndices.copyApply { and(skipOccIndices) }.isEmpty + override fun combine(that: AlphaNode, subst: Subst): ReteNode = that - override fun combine(that: AlphaNode): ReteNode = that + override fun collect(occArray: Array) {} } /** @@ -117,37 +95,18 @@ internal class ReteRuleMatcherImpl(val rule: Rule) : RuleMatcher { val posInHead: Int, val subst: Subst) : ReteNode() { - val metaIndices: BitSet? = - if (!subst.isEmpty) bitSet(subst.keys().map { metaLogical -> indexOf(metaLogical) }) else null + val occIndex = indexOf(occurrence) - val occIdx = indexOf(occurrence) + override fun subst(): Subst = subst - private val idx2subst = HashMap() + override fun occupiesHeadPosition(headPos: Int): Boolean = posInHead == headPos - init { - for ((meta, subst) in subst.asMap().entries) { - idx2subst[indexOf(meta)] = subst - } - } + override fun containsOccurrence(occIdx: Int): Boolean = occIndex == occIdx - override fun getMeta(): BitSet? = metaIndices - - override fun isCompatible(that: AlphaNode): Boolean { - // check that the occurrence-position pair is unique - if (this.occIdx == that.occIdx || this.posInHead == that.posInHead) return false - - return isSubstCompatible(that) - } - - override fun containsOccurrence(skipOccIndices: BitSet): Boolean = skipOccIndices[occIdx] - - override fun getSubst(metaIdx: Int): Any? = idx2subst[metaIdx] - - override fun combine(that: AlphaNode): ReteNode = BetaNode(this, that) - - override fun collectData(occArray: Array, allSubst: Subst): Subst { + override fun combine(that: AlphaNode, subst: Subst): ReteNode = BetaNode(this, that, subst) + + override fun collect(occArray: Array) { occArray[posInHead] = occurrence - return subst.fold(allSubst) { acc, (k, v) -> acc.put(k, v) } } } @@ -161,145 +120,225 @@ internal class ReteRuleMatcherImpl(val rule: Rule) : RuleMatcher { val right: AlphaNode + val subst: Subst + val positions: BitSet val occIndices: BitSet - val metaIndices: BitSet? - - constructor(left: AlphaNode, right: AlphaNode) { + constructor(left: AlphaNode, right: AlphaNode, subst: Subst) { this.left = left this.right = right + this.subst = subst this.positions = bitSet(left.posInHead).apply { set(right.posInHead) } - this.occIndices = bitSet(left.occIdx).apply { set(right.occIdx) } - this.metaIndices = left.metaIndices?.let { - if (right.metaIndices != null) it.copyApply { or(right.metaIndices) } else it } ?: - right.metaIndices + this.occIndices = bitSet(left.occIndex).apply { set(right.occIndex) } } - constructor(left: BetaNode, right: AlphaNode) { + constructor(left: BetaNode, right: AlphaNode, subst: Subst) { this.left = left this.right = right + this.subst = subst this.positions = left.positions.copyApply { set(right.posInHead) } - this.occIndices = left.occIndices.copyApply { set(right.occIdx) } - this.metaIndices = left.metaIndices?.let { - if (right.metaIndices != null) it.copyApply { or(right.metaIndices) } else it } ?: - right.metaIndices + this.occIndices = left.occIndices.copyApply { set(right.occIndex) } } - override fun getMeta(): BitSet? = metaIndices + override fun subst(): Subst = subst - override fun isCompatible(that: AlphaNode): Boolean { - // check that the occurrence-position pair is unique - if (occIndices[that.occIdx] || positions[that.posInHead]) return false + override fun occupiesHeadPosition(headPos: Int): Boolean = positions[headPos] - return isSubstCompatible(that) + override fun containsOccurrence(occIdx: Int): Boolean = occIndices[occIdx] + + override fun combine(that: AlphaNode, subst: Subst): ReteNode = BetaNode(this, that, subst) + + override fun collect(occArray: Array) { + left.collect(occArray) + right.collect(occArray) } - - override fun containsOccurrence(skipOccIndices: BitSet): Boolean = - !this.occIndices.copyApply { and(skipOccIndices) }.isEmpty - - override fun getSubst(metaIdx: Int): Any? = right.getSubst(metaIdx) ?: left.getSubst(metaIdx) - - override fun combine(that: AlphaNode): ReteNode = BetaNode(this, that) - - override fun collectData(occArray: Array, allSubst: Subst): Subst = - left.collectData(occArray, right.collectData(occArray, allSubst)) - } - inner class Layer(val vacancies: Int, proto: Layer? = null) { - - val nodesList : MutableList = proto?.nodesList ?: ArrayList(4) - - val startIdx : Int = nodesList.size + /** + * A layer extends incrementally its prototype with new nodes. + */ + inner class Layer(val vacancies: Int, private var block: DelayedBlock? = null, proto: Layer? = null) { val final = (vacancies == 0) + private val occIndices: BitSet = (proto?.complete()?.occIndices?.clone() ?: BitSet()) as BitSet + + private val nodesList : MutableList = proto?.complete()?.nodesList ?: ArrayList(4) + + private var startIdx : Int = nodesList.size + + /** constructs initial layer containing only [InitialNode] */ constructor(vacancies: Int, node: ReteNode) : this (vacancies) { nodesList.add(node) } - fun addNode(n: ReteNode) { - nodesList.add(n) + fun isEmpty(): Boolean = TODO() //nodesList.isEmpty() && (queue?.isEmpty() ?: true) + + fun containsOccurrence(occIdx: Int): Boolean { + return occIndices[occIdx] } - fun nodes() : Iterable = nodesList.subList(startIdx, nodesList.size) + fun ownNodes() : Iterable { + complete() + return nodesList.subList(startIdx, nodesList.size) + } - fun allNodes() : Iterable = nodesList + fun allNodes() : Iterable { + complete() + return nodesList + } + + private fun complete() : Layer { + block?.complete { + nodesList.add(it) + when(it) { + is AlphaNode -> occIndices.set(it.occIndex) + is BetaNode -> occIndices.or(it.occIndices) + } + } + this.block = null + return this + } + } + + /** + * An experimental feature that should allow for delayed update of network nodes. + */ + abstract inner class DelayedBlock { + + abstract fun proto(): Layer? + + abstract fun complete(sink: (ReteNode) -> Unit) + + abstract fun isEmpty(): Boolean + + abstract fun containsOccurrence(occIdx: Int): Boolean } + inner class IntroBlock (val occ: Occurrence, val headPosMask: BitSet, val onTopOf: Layer): DelayedBlock() { + + val occIdx = indexOf(occ) + + override fun proto(): Layer = onTopOf + + override fun complete(sink: (ReteNode) -> Unit) { + + + val it = headPosMask.allSetBits() + while (it.hasNext()) { + val headPos = it.next() + for (n in onTopOf.allNodes()) { + if (n.containsOccurrence(occIdx) || n.occupiesHeadPosition(headPos)) continue + + with(createOccurrenceMatcher(n.subst())) { + if (matches(head[headPos], occ)) { + val intro = AlphaNode(occ, headPos, subst()) + sink(n.combine(intro, subst())) + } + } + } + } + + } + + override fun isEmpty(): Boolean = onTopOf.isEmpty() + + override fun containsOccurrence(occIdx: Int): Boolean = + occIdx == this.occIdx || onTopOf.containsOccurrence(occIdx) + } + + inner class DropBlock (val occ: Occurrence, val from: Layer) : DelayedBlock() { + + val occIdx = indexOf(occ) + + override fun proto(): Layer? = null + + override fun complete(sink: (ReteNode) -> Unit) { + for (n in from.allNodes()) { + if (!n.containsOccurrence(occIdx)) { + sink(n) + } + } + } + + override fun isEmpty(): Boolean = from.isEmpty() + + override fun containsOccurrence(occIdx: Int): Boolean = + occIdx != this.occIdx && from.containsOccurrence(occIdx) + } + + /** + * Collection of [Layer] instances. + * The new layers are prepended to the [layers] list. + * Invariant: the last layer is always the initial one. + */ inner class Generation(val layers: List) { init { assert(layers.isNotEmpty()) } - fun introduce(occIdx: Int, alphaNodes: Collection): Generation { - // propagation history - val initLayer = layers.last() - assert(initLayer.nodesList.size == 1) - val initialNode = initLayer.nodesList.first() as InitialNode - if (propagation && initialNode.occIndices[occIdx]) { - // introducing an already seen constraint - val newLayers = ArrayList(4) - for (la in layers) { - if (la != initLayer) newLayers.add(Layer(la.vacancies, la)) - } - newLayers.add(Layer(headSize, InitialNode(initialNode.occIndices))) + fun introduce(occ: Occurrence, headPosMask: BitSet): Generation { + val occIdx = indexOf(occ) + val reactivated = layers.first().containsOccurrence(occIdx) - return Generation(newLayers) + // propagation history + if (propagation && reactivated) { + // all matches are already in the "final" layer + return this; } else { val newLayers = ArrayList(4) - val occIndices = BitSet() - var last: Layer? = null - for (curr in layers) { - if (!curr.final) { - val newLayer = Layer(curr.vacancies - 1, last) - for (n in curr.allNodes()) { - for (intro in alphaNodes) { - if (n.isCompatible(intro)) { - newLayer.addNode(n.combine(intro)) - occIndices.set(intro.occIdx) - } - } - } + var lastLayer: Layer? = null + for (currLayer in layers) { + if (!currLayer.final) { + val newLayer = Layer(currLayer.vacancies - 1, IntroBlock(occ, headPosMask, currLayer), lastLayer) + newLayers.add(newLayer) + + } else if (reactivated) { + newLayers.add(currLayer) } - last = curr + lastLayer = currLayer } - occIndices.or(initialNode.occIndices) - newLayers.add(Layer(headSize, InitialNode(occIndices))) - + newLayers.add(Layer(headSize, InitialNode())) return Generation(newLayers) } } + fun drop(occ: Occurrence): Generation { + val newLayers = ArrayList(4) + for (currLayer in layers) { + val newLayer = Layer(currLayer.vacancies, DropBlock(occ, currLayer)) + newLayers.add(newLayer) + } + return Generation(newLayers) + } fun matches(): Collection { val topLayer = layers.first() if (topLayer.final) { - + val uniqueSignatures = HashSet() val matches = ArrayList() - for (n in topLayer.nodes()) { - // any excluded occurrences? - if (n.containsOccurrence(skipOccIndices)) continue - - val allSubst : Subst = emptySubst() + for (n in topLayer.ownNodes()) { val occArray = arrayOfNulls(headSize) - n.collectData(occArray, allSubst) + n.collect(occArray) + val signature = occArray.map { it!!.identity }.toIntArray().toSignature() + if (consumedSignatures.contains(signature) || uniqueSignatures.contains(signature)) continue + uniqueSignatures.add(signature) val occList = occArray.toList() as List - val keptCount = rule.headKept().count() + val keptCount = lookupRule().headKept().count() - matches.add(RuleMatchImpl(rule, - allSubst, + matches.add(RuleMatchImpl(lookupRule(), + n.subst(), occList.subList(0, keptCount), occList.subList(keptCount, occList.size))) } @@ -313,64 +352,46 @@ internal class ReteRuleMatcherImpl(val rule: Rule) : RuleMatcher { } - override fun rule(): Rule = rule + override fun rule(): Rule = lookupRule() // for tests only override fun expand(occ: Occurrence): ReteNetwork = - expand(occ, bitSetOfOnes(headSize)) + expand(occ, bitSetOfOnes(headSize), null) - override fun expand(occ: Occurrence, mask: BitSet): ReteNetwork { - // raising from the dead, huh? - val occIdx = indexOf(occ) - skipOccIndices.clear(occIdx) - - val alphaNodes = arrayListOf() - val it = mask.allSetBits() - while (it.hasNext()){ - - val posInHead = it.next() - val matcher = createOccurrenceMatcher(emptySubst()) - if (matcher.matches(head[posInHead], occ)) { - alphaNodes.add(AlphaNode(occ, posInHead, matcher.subst())) - } - - } - - val newGeneration = generations.last().introduce(occIdx, alphaNodes) - generations.add(newGeneration) + override fun expand(occ: Occurrence, mask: BitSet, profiler: Profiler?): ReteNetwork { + this.lastGeneration = lastGeneration.introduce(occ, mask) return this } override fun contract(occ: Occurrence): ReteNetwork { - skipOccIndices.set(indexOf(occ)) + this.lastGeneration = lastGeneration.drop(occ) return this } override fun forgetSeen(occ: Occurrence): ReteNetwork { - TODO("not implemented") //To change body of created functions use File | Settings | File Templates. + return this } override fun forgetConsumed(occ: Occurrence): ReteNetwork { - TODO("not implemented") //To change body of created functions use File | Settings | File Templates. + consumedSignatures.removeIf({ it.contains(occ.identity) }) + return this } - override fun matches(): Collection = generations.last().matches() + override fun matches(): Collection = lastGeneration.matches() override fun consume(ruleMatch: RuleMatchEx): RuleMatchingProbe { - TODO("not implemented") //To change body of created functions use File | Settings | File Templates. + consumedSignatures.add(ruleMatch.signatureArray().toSignature()) + return this } override fun forget(ruleMatch: RuleMatchEx): RuleMatchingProbe { - TODO("not implemented") //To change body of created functions use File | Settings | File Templates. + consumedSignatures.add(ruleMatch.signatureArray().toSignature()) + return this } fun indexOf(occurrence: ConstraintOccurrence): Int = seenOcc2Idx[occurrence] ?: (nextOccIdx++).also { idx -> seenOcc2Idx[occurrence] = idx } - - fun indexOf(metaLogical: MetaLogical<*>): Int = - meta2Idx[metaLogical] ?: (nextMetaIdx++).also { idx -> meta2Idx[metaLogical] = idx } - } } \ No newline at end of file diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/internal/RuleMatcherImpl.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/internal/RuleMatcherImpl.kt index 3ab2d0da..1fccace0 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/internal/RuleMatcherImpl.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/internal/RuleMatcherImpl.kt @@ -19,6 +19,7 @@ package jetbrains.mps.logic.reactor.core.internal import com.github.andrewoma.dexx.collection.Sets import gnu.trove.list.TIntList import gnu.trove.list.array.TIntArrayList +import gnu.trove.set.hash.TIntHashSet import jetbrains.mps.logic.reactor.core.* import jetbrains.mps.logic.reactor.program.Constraint import jetbrains.mps.logic.reactor.program.Rule @@ -27,13 +28,6 @@ import java.util.* import kotlin.collections.ArrayList import com.github.andrewoma.dexx.collection.Set as PersSet -// Trove stuff -typealias Signature = TIntList - -fun Signature.copy() = TIntArrayList(this) - -fun IntArray.toSignature() = TIntArrayList(this) - /** * This implementation of [RuleMatcher] is based on a simple algorithm which enumerates all possible @@ -95,40 +89,52 @@ internal class RuleMatcherImpl(private val ruleLookup: RuleLookup, * Expands the front by creating new leaf nodes that match the occurrence. * Mask specifies possible slots for the occurrence. */ - override fun expand(occ: Occurrence, mask: BitSet): RuleMatchFront { + override fun expand(occ: Occurrence, mask: BitSet, profiler: Profiler?): RuleMatchFront { val reactivated = seenOccurrences.contains(occ.identity) val newSeen = if (reactivated) seenOccurrences else seenOccurrences.add(occ.identity) - val newTrunkNodes = arrayListOf().apply { addAll(trunkNodes) } + val newTrunkNodes = ArrayList(trunkNodes.size).apply { addAll(trunkNodes) } val newLeafNodes = arrayListOf() var newSignatures = leafSignatures val expanded = ArrayList() - for (node in (trunkNodes + BaseMatchNode())) { - val effMask = mask.copyApply { and(node.vacant) } - if ((node is MatchNode && node.hasOccurrence(occ)) || effMask.isEmpty) continue - val it = effMask.allSetBits() - while (it.hasNext()) { - val headIdx = it.next() - val subst = if (node is MatchNode) node.subst else emptySubst() - with (createOccurrenceMatcher(subst)) { - if (matches(head[headIdx], occ)) { - expanded.add(MatchNode(subst(), node, occ, headIdx)) + val effMask = mask.clone() as BitSet + + for (node in (trunkNodes.asSequence() + BaseMatchNode())) { + effMask.clear() + effMask.or(mask) + effMask.and(node.vacant) +// val effMask = mask.copyApply { and(node.vacant) } +// profiler.profile("expand_node1_${occ.constraint.symbol()}") { + + if (!(effMask.isEmpty || node is MatchNode && node.hasOccurrence(occ))) { + + expanded.clear() + + val it = effMask.allSetBits() + while (it.hasNext()) { + val headIdx = it.next() + val subst = if (node is MatchNode) node.subst else emptySubst() + with(createOccurrenceMatcher(subst)) { + if (matches(head[headIdx], occ)) { + expanded.add(MatchNode(subst(), node, occ, headIdx)) + } + } + } + + for (ex in expanded) { + if (ex.leaf) { + // ensure reactivated have effect + val signature = ex.signature() + if (!reactivated && newSignatures.contains(signature)) break + // ...unless propagation (to avoid cycles) + if (reactivated && propagation && consumedSignatures.contains(signature)) break + newLeafNodes.add(ex) + newSignatures = newSignatures.add(signature) + } + newTrunkNodes.add(ex) + } } } - for(ex in expanded) { - if (ex.leaf) { - // ensure reactivated have effect - val signature = ex.signature() - if (!reactivated && newSignatures.contains(signature)) break - // ...unless propagation (to avoid cycles) - if (reactivated && propagation && consumedSignatures.contains(signature)) break - newLeafNodes.add(ex) - newSignatures = newSignatures.add(signature) - } - newTrunkNodes.add(ex) - } - expanded.clear() - } return RuleMatchFront(newTrunkNodes, newLeafNodes, newSignatures, newSeen, consumedSignatures) } @@ -184,6 +190,9 @@ internal class RuleMatcherImpl(private val ruleLookup: RuleLookup, { val leaf = vacant.cardinality() == 0 + val trail: PersSet = + if (parent is MatchNode) parent.trail.add(occurrence.identity) else Sets.of(occurrence.identity) + fun constraint(): Constraint = head[headIndex] // a signature is a (partial) set of constraint occurrences that belong to this node @@ -191,8 +200,8 @@ internal class RuleMatcherImpl(private val ruleLookup: RuleLookup, fold(IntArray(head.size).toSignature()) { s, n -> s.apply { set(n.headIndex, n.occurrence.identity) } } - fun hasOccurrence(occ: Occurrence): Boolean = - fold(false) { flag, n -> flag || n.occurrence === occ } // referential equality ! + fun hasOccurrence(occ: Occurrence): Boolean = trail.contains(occ.identity) +// fold(false) { flag, n -> flag || n.occurrence === occ } // referential equality ! fun matches() : Subst? = fold(createOccurrenceMatcher(emptySubst())) { matcher, n ->