Refactor and redesign ReteRuleMatcher for better memory footprint.

Drop layer prototype, make layer mutable.
Keep nodes in a layer in a linked list with fail-safe iterator.
Introduce update queue for update blocks.
Update blocks run incrementally unless reset.
Maintain trail of introduced and dropped occurrences.
Generation iterates over nodes in the last (final) layer incrementally,
unless reset.
Introduce new test.
This commit is contained in:
Fedor Isakov 2019-11-15 16:49:50 +01:00
parent 96fbab18b4
commit 9fab2e9d34
2 changed files with 362 additions and 143 deletions

View File

@ -16,12 +16,13 @@
package jetbrains.mps.logic.reactor.core.internal
import gnu.trove.map.hash.TIntObjectHashMap
import gnu.trove.set.hash.TIntHashSet
import jetbrains.mps.logic.reactor.core.*
import jetbrains.mps.logic.reactor.evaluation.ConstraintOccurrence
import jetbrains.mps.logic.reactor.program.Rule
import jetbrains.mps.logic.reactor.util.*
import java.util.*
import kotlin.NoSuchElementException
import kotlin.collections.ArrayList
/**
@ -35,11 +36,11 @@ import kotlin.collections.ArrayList
* @author Fedor Isakov
*/
typealias Footprint = TIntHashSet
typealias Trail = TIntHashSet
fun footprintOf(): Footprint = TIntHashSet()
fun trailOf(): Trail = TIntHashSet()
fun Signature.toFootprint() = TIntHashSet(this)
fun Signature.toTrail() = TIntHashSet(this)
internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
private val tag: Any) : RuleMatcher
@ -65,13 +66,15 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
// rules with empty head are not allowed
assert(headSize > 0)
}
var lastGeneration = Generation(arrayListOf(Layer(headSize, InitialNode())))
var lastGeneration = Generation(Layer(InitialNode()))
val consumedSignatures = HashSet<Signature>()
abstract inner class ReteNode
{
abstract val depth: Int
abstract val signature: Signature
abstract fun subst(): Subst
@ -88,6 +91,8 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
inner class InitialNode() : ReteNode() {
override val depth = 0
override val signature: Signature = noSignature()
override fun subst(): Subst = emptySubst()
@ -108,6 +113,9 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
val posInHead: Int,
val subst: Subst) : ReteNode()
{
override val depth = 1
override val signature = IntArray(headSize).toSignature().apply { set(posInHead, occurrence.identity) }
override fun subst(): Subst = subst
@ -117,7 +125,7 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
override fun containsOccurrence(occ: Occurrence): Boolean = occ === occurrence // reference eq!
override fun combine(that: AlphaNode, subst: Subst): ReteNode = BetaNode(this, that, subst)
override fun collect(occArray: Array<Occurrence?>) {
occArray[posInHead] = occurrence
}
@ -131,6 +139,9 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
override val signature: Signature
override val depth: Int
get() = left.depth + right.depth
val left: ReteNode
val right: AlphaNode
@ -167,213 +178,270 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
override fun containsOccurrence(occ: Occurrence): Boolean = signature.contains(occ.identity)
override fun combine(that: AlphaNode, subst: Subst): ReteNode = BetaNode(this, that, subst)
override fun collect(occArray: Array<Occurrence?>) {
left.collect(occArray)
right.collect(occArray)
}
}
/**
* A layer extends incrementally its prototype with new nodes.
* A layer contains nodes with shared properties:
* - number of "vacancies" in the matching slots;
* - slot position occupied by the last introduced occurrence.
*
* Node list is built incrementally from update blocks.
*/
inner class Layer(val vacancies: Int, private var block: DelayedBlock? = null, proto: Layer? = null) {
inner class Layer() {
val final = (vacancies == 0)
private val introTrail: Trail = trailOf()
private val footprint: Footprint by lazy(LazyThreadSafetyMode.NONE) {
(proto?.complete()?.footprint ?: footprintOf())
}
private val droppedTrail: Trail = trailOf()
private val nodesList : MutableList<ReteNode> by lazy(LazyThreadSafetyMode.NONE) {
val protoNodes = proto?.complete()?.nodesList ?: ArrayList(4)
this.startIdx = protoNodes.size
protoNodes
}
private val nodeList = LoopLinkedList<ReteNode>()
// guaranteed to be initialized to proper value before this var is accessed
private var startIdx : Int = -1
private val updateQueue: MutableList<UpdateBlock> = LinkedList()
var primed = false
/** constructs initial layer containing only [InitialNode] */
constructor(vacancies: Int, node: ReteNode) : this (vacancies) {
nodesList.add(node)
constructor(node: ReteNode) : this(){
nodeList.add(node)
}
fun isEmpty() = nodeList.isEmpty()
fun containsOccurrence(occ: Occurrence): Boolean {
return footprint.contains(occ.identity)
return introTrail.contains(occ.identity)
}
fun forgetContains(occ: Occurrence): Boolean {
return footprint.remove(occ.identity)
droppedTrail.remove(occ.identity)
// FIXME this breaks the internal invariant
return introTrail.remove(occ.identity)
}
fun ownNodes() : Iterable<ReteNode> {
complete()
return nodesList.subList(startIdx, nodesList.size)
}
fun iterate(): Iterator<ReteNode> = nodeList.iterator()
fun allNodes(droppedFootprint: Footprint?) : Iterable<ReteNode> {
complete(droppedFootprint)
return nodesList
}
fun copyIterator(it: Iterator<ReteNode>): Iterator<ReteNode> =
nodeList.iterator((it as LoopLinkedList<ReteNode>.Iterator).current)
private fun complete(droppedFootprint: Footprint? = null) : Layer {
block?.complete(droppedFootprint) {
nodesList.add(it)
when (it) {
is AlphaNode -> footprint.addAll(it.signature)
is BetaNode -> footprint.addAll(it.signature)
}
fun nextNode(it: Iterator<ReteNode>): ReteNode? {
if (!it.hasNext()) {
prime()
}
return if(it.hasNext()) it.next() else null
}
fun prime(): Layer {
if (!primed) {
for (block in updateQueue) {
update(block)
}
this.primed = true
}
this.block = null
return this
}
}
/**
* An experimental feature that should allow for delayed update of network nodes.
*/
abstract inner class DelayedBlock(val occurrence: Occurrence) {
fun reset() {
for (block in updateQueue) {
block.reset()
}
this.primed = false
}
abstract fun proto(): Layer?
fun update(block: UpdateBlock): Boolean {
return block.update(nodeList, introTrail, droppedTrail)
}
abstract fun complete(droppedFootprint: Footprint?, sink: (ReteNode) -> Unit)
fun queueUpdate(block: UpdateBlock): Layer {
updateQueue.removeIf { it.occurrence === block.occurrence }
updateQueue.add(block)
return this
}
abstract fun containsOccurrence(occ: Occurrence): Boolean
fun runUpdate(block: UpdateBlock): Layer {
updateQueue.removeIf { it.occurrence === block.occurrence }
update(block)
return this
}
}
inner class IntroBlock (occurrence: Occurrence, val headPosMask: BitSet, val onTopOf: Layer) :
DelayedBlock(occurrence)
{
abstract inner class UpdateBlock(val occurrence: Occurrence) {
override fun proto(): Layer = onTopOf
abstract fun reset()
abstract fun update(nodes: LoopLinkedList<ReteNode>, introTrail: Trail, droppedTrail: Trail): Boolean
}
inner class IntroBlock (occurrence: Occurrence,
val headPosMask: BitSet,
val onTopOf: Layer) : UpdateBlock(occurrence) {
var nodesIt: Iterator<ReteNode> = onTopOf.iterate()
override fun reset() {
// FIXME when enabled the following code causes major slowdown on certain programs
// nodesIt = onTopOf.iterate()
}
override fun update(nodes: LoopLinkedList<ReteNode>, introTrail: Trail, droppedTrail: Trail): Boolean {
if (droppedTrail.contains(occurrence.identity)) return false
var found = false
while (true) {
val n = onTopOf.nextNode(nodesIt)
if (n == null) break
override fun complete(droppedFootprint: Footprint?, sink: (ReteNode) -> Unit) {
if (droppedFootprint?.contains(occurrence.identity) ?: false) return
for (n in onTopOf.allNodes(droppedFootprint)) {
if (n.containsOccurrence(occurrence)) continue
val it = headPosMask.allSetBits()
while (it.hasNext()) {
val headPos = it.next()
if (n.occupiesHeadPosition(headPos)) continue
with(createOccurrenceMatcher(n.subst())) {
if (matches(head[headPos], occurrence)) {
val intro = AlphaNode(occurrence, headPos, subst())
sink(n.combine(intro, subst()))
val it = headPosMask.allSetBits()
while (it.hasNext()) {
val headPos = it.next()
if (n.occupiesHeadPosition(headPos)) continue
with(createOccurrenceMatcher(n.subst())) {
if (matches(head[headPos], occurrence)) {
val intro = AlphaNode(occurrence, headPos, subst())
nodes.add(n.combine(intro, subst()))
found = true
}
}
}
}
}
override fun containsOccurrence(occ: Occurrence): Boolean =
occurrence === occ || onTopOf.containsOccurrence(occ) // reference eq!
}
inner class DropBlock (occurrence: Occurrence, val from: Layer) : DelayedBlock(occurrence) {
override fun proto(): Layer? = null
override fun complete(droppedFootprint: Footprint?, sink: (ReteNode) -> Unit) {
if (droppedFootprint?.contains(occurrence.identity) ?: false) return
droppedFootprint?.add(occurrence.identity)
for (n in from.allNodes(droppedFootprint)) {
if (!n.containsOccurrence(occurrence)) {
sink(n)
if (found) {
introTrail.add(occurrence.identity)
}
}
return found
}
}
inner class DropBlock (occurrence: Occurrence) : UpdateBlock(occurrence) {
override fun reset() {
// NOP
}
override fun update(nodes: LoopLinkedList<ReteNode>, introTrail: Trail, droppedTrail: Trail): Boolean {
val it = nodes.iterator()
while (it.hasNext()) {
val n = it.next()
if (n.containsOccurrence(occurrence)) {
it.remove()
}
}
introTrail.remove(occurrence.identity)
droppedTrail.add(occurrence.identity)
return false
}
override fun containsOccurrence(occ: Occurrence): Boolean =
occurrence !== occ && from.containsOccurrence(occ) // reference neq!
}
/**
* Collection of [Layer] instances.
* The new layers are prepended to the [layers] list.
* Invariant: the last layer is always the initial one.
* The new layers are appended to the [layers] list.
* Invariant: the first layer is always the initial one.
*/
inner class Generation(val layers: List<Layer>) {
inner class Generation private constructor (prev: Generation?) {
val layers: MutableList<Layer> = prev?.layers ?: ArrayList<Layer>(4)
lateinit var nodesIt: Iterator<ReteNode>
private var __matches: Collection<RuleMatchImpl>? = null
val matches: Collection<RuleMatchImpl>
get() {
if (__matches == null) {
this.__matches = calcMatches()
}
return __matches !!
}
init {
assert(layers.isNotEmpty())
if (prev != null) {
this.nodesIt = prev.nodesIt
}
}
constructor(initLayer: Layer) : this(null) {
layers.add(initLayer)
for (i in 1..headSize) {
layers.add(Layer())
}
nodesIt = layers.last().iterate()
}
fun reset(): Generation {
this.nodesIt = layers.last().iterate()
return this
}
fun introduce(occurrence: Occurrence, headPosMask: BitSet): Generation {
val reactivated = layers.first().containsOccurrence(occurrence)
val reactivated = layers.any { it.containsOccurrence(occurrence) }
// propagation history
if (propagation && reactivated) {
// all matches are already in the "final" layer
return this;
if (propagation && reactivated) return nextGeneration().reset()
} else {
val newLayers = ArrayList<Layer>(4)
var lastLayer: Layer? = null
for (currLayer in layers) {
if (!currLayer.final) {
newLayers.add(Layer(
currLayer.vacancies - 1,
IntroBlock(occurrence, headPosMask, currLayer),
lastLayer))
} else if (reactivated) {
newLayers.add(currLayer)
}
lastLayer = currLayer
var firstAffected = -1
for ((idx, layer) in layers.withIndex()) {
if (idx > 0 && headPosMask.get(idx - 1)) {
if (firstAffected < 0) firstAffected = idx
layer.queueUpdate(IntroBlock(occurrence, headPosMask, layers[idx - 1]))
}
newLayers.add(Layer(headSize, InitialNode()))
return Generation(newLayers)
}
for (i in firstAffected until layers.size) {
layers[i].reset()
}
return nextGeneration()
}
fun drop(occurrence: Occurrence): Generation {
val newLayers = ArrayList<Layer>(4)
for (currLayer in layers) {
val newLayer = Layer(currLayer.vacancies, DropBlock(occurrence, currLayer))
newLayers.add(newLayer)
for (layer in layers) {
layer.runUpdate(DropBlock(occurrence))
}
return Generation(newLayers)
return nextGeneration().reset()
}
/*
* Allows to avoid triggering reactivation logic in the next call of "introduce" for this occurrence.
*/
fun forgetIntroduced(occurrence: Occurrence) {
layers.first().forgetContains(occurrence)
fun forgetIntroduced(occurrence: Occurrence): Generation {
layers.last().forgetContains(occurrence)
return nextGeneration()
}
fun matches(): Collection<RuleMatchImpl> {
val topLayer = layers.first()
if (topLayer.final) {
val uniqueSignatures = HashSet<Signature>()
val matches = ArrayList<RuleMatchImpl>()
for (n in topLayer.ownNodes()) {
val signature = n.signature
if (consumedSignatures.contains(signature) || uniqueSignatures.contains(signature)) continue
uniqueSignatures.add(signature)
fun nextGeneration(): Generation = Generation(this)
val occArray = arrayOfNulls<Occurrence>(headSize).apply(n::collect)
val occList = occArray.toList() as List<Occurrence>
val keptCount = lookupRule().headKept().count()
fun calcMatches(): Collection<RuleMatchImpl> {
val topLayer = layers.last()
val uniqueSignatures = HashSet<Signature>()
val matches = ArrayList<RuleMatchImpl>()
while(true) {
val n = topLayer.nextNode(nodesIt)
if (n == null) break
val signature = n.signature
if (consumedSignatures.contains(signature) || uniqueSignatures.contains(signature)) continue
uniqueSignatures.add(signature)
matches.add(RuleMatchImpl(lookupRule(),
n.subst(),
occList.subList(0, keptCount),
occList.subList(keptCount, occList.size)))
}
return matches
val occArray = arrayOfNulls<Occurrence>(headSize).apply(n::collect)
val occList = occArray.toList() as List<Occurrence>
val keptCount = lookupRule().headKept().count()
matches.add(RuleMatchImpl(lookupRule(),
n.subst(),
occList.subList(0, keptCount),
occList.subList(keptCount, occList.size)))
}
else {
return emptyList()
}
return matches
}
}
@ -404,7 +472,9 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
return this
}
override fun matches(): Collection<RuleMatchImpl> = lastGeneration.matches()
override fun matches(): Collection<RuleMatchImpl> {
return lastGeneration.matches
}
override fun consume(ruleMatch: RuleMatchEx): RuleMatchingProbe {
consumedSignatures.add(ruleMatch.signatureArray().toSignature())
@ -418,4 +488,101 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup,
}
}
}
/**
* Simple linked list with fail-safe iterator.
*/
class LoopLinkedList<T> : Iterable<T> {
open inner class Joint {
lateinit var prev: Joint
lateinit var next: Joint
open fun remove() {
prev.next = this.next
next.prev = this.prev
}
fun prepend(that: Joint) {
insert(prev, that, this)
}
fun append(that: Joint) {
insert(this, that, next)
}
fun insert(pred: Joint, node: Joint, succ: Joint) {
if (pred === node || succ === node) throw IllegalStateException()
pred.next = node
node.prev = pred
succ.prev = node
node.next = succ
}
fun clear() {
this.next = this
this.prev = this
}
fun isDegenerate() : Boolean {
return this.prev === this && this.next === this
}
}
inner class InitialJoint: Joint() {
init {
prev = this
next = this
}
override fun remove() {
throw IllegalStateException()
}
}
inner class DataJoint(val value: T) : Joint() {}
val initial = InitialJoint()
fun add(value: T) {
initial.prepend(DataJoint(value))
}
fun clear() {
initial.clear()
}
fun isEmpty(): Boolean = initial.isDegenerate()
override fun iterator(): MutableIterator<T> = iterator(initial)
fun iterator(start: Joint): MutableIterator<T> = Iterator(start)
inner class Iterator : MutableIterator<T> {
lateinit var current : LoopLinkedList<T>.Joint
constructor(start: Joint) {
this.current = start
}
override fun hasNext(): Boolean = current.next !is InitialJoint
override fun next(): T {
this.current = current.next
if (current !is DataJoint) throw NoSuchElementException()
return (current as DataJoint).value
}
override fun remove() {
val prev = current.prev
current.remove()
this.current = prev
}
}
}

View File

@ -293,6 +293,10 @@ class TestRuleMatcher {
matches().size shouldBe 0 }.run {
expand(bara) }.apply {
// the constraint that has been dropped remains dropped
matches().size shouldBe 0
expand(occurrence("bar", "a")) }.apply {
matches().size shouldBe 1
with(matches().first()) {
@ -306,6 +310,52 @@ class TestRuleMatcher {
}
}
@Test
fun testAmbiguousRules() {
val (X, Y) = metaLogical<String>("X", "Y")
val x = X.logical()
val y = Y.logical()
with(programWithRules(
rule("rule1",
headReplaced(
constraint("bar", X, Y)
),
headKept(
constraint("foo", X),
constraint("foo", Y)
),
body(
constraint("expected")
)),
rule("rule2",
headReplaced(
constraint("bar", X, Y)
),
headKept(
constraint("foo", X)
),
body(
constraint("unexpected")
))))
{
val foox = occurrence("foo", x)
val fooy = occurrence("foo", y)
with(ruleMatcher().probe()) {
expand(occurrence("foo", "schmoo")) }.run {
expand(foox) }.run {
expand(fooy) }.run {
expand(fooy) }.run {
expand(fooy) }.run {
expand(occurrence("bar", x, y)) }.apply {
matches().size shouldBe 1
matches().first().rule().tag() shouldBe "rule1"
}
}
}
@Test
fun testMetaLogical() {
@ -500,10 +550,11 @@ class TestRuleMatcher {
constraint("qux")
))))
{
val occ2reactivate = occurrence("bazz", term("k", logicalVar(yLogical)))
val bazzk = occurrence("bazz", term("k", logicalVar(yLogical)))
with(ruleMatcher().probe()) {
expand(occ2reactivate) }.apply {
expand(bazzk) }.apply {
matches().size shouldBe 0
yLogical.findRoot().value() shouldBe null
}.run {
@ -519,7 +570,7 @@ class TestRuleMatcher {
// test reactivate occurrence
yLogical.findRoot().set(parseTerm("h"))
}.run {
expand(occ2reactivate) }.apply {
expand(bazzk) }.apply {
matches().size shouldBe 1
with(matches().first()) {
@ -599,6 +650,7 @@ class TestRuleMatcher {
with(matches().first()) {
rule().uniqueTag().toString() shouldBe "rule1"
logicalContext().variable(X).value() shouldBe "a"
}
}.run {