Optimizing new matching algorithm. Keep an index of occurrence slots for each rule.

This commit is contained in:
Fedor Isakov 2018-08-04 14:24:37 +02:00
parent d8c2310f98
commit 75c6dff551
5 changed files with 153 additions and 131 deletions

View File

@ -62,10 +62,8 @@ class Dispatcher (val ruleIndex: RuleIndex) {
fun activated(active: ConstraintOccurrence): DispatchFringe {
return DispatchFringe(this,
ruleIndex.forOccurrence(active).mapNotNull { rule ->
rule2fringe[rule]
}.map { fringe ->
fringe.expand(active)
ruleIndex.forOccurrenceWithMask(active).mapNotNull { (rule, mask) ->
rule2fringe[rule]?.expand(active, mask)
})
}

View File

@ -19,12 +19,12 @@ package jetbrains.mps.logic.reactor.core
import jetbrains.mps.logic.reactor.evaluation.ConstraintOccurrence
import jetbrains.mps.logic.reactor.logical.Logical
import jetbrains.mps.logic.reactor.logical.MetaLogical
import jetbrains.mps.logic.reactor.program.Constraint
import jetbrains.mps.logic.reactor.program.ConstraintSymbol
import jetbrains.mps.logic.reactor.program.Handler
import jetbrains.mps.logic.reactor.program.Rule
import jetbrains.mps.logic.reactor.program.*
import jetbrains.mps.logic.reactor.util.allSetBits
import jetbrains.mps.unification.Term
import java.util.*
import kotlin.collections.ArrayList
import kotlin.collections.HashMap
/**
* @author Fedor Isakov
@ -32,159 +32,162 @@ import java.util.*
class RuleIndex(handlers: Iterable<Handler>) : Iterable<Rule> {
private val primarySymbol2valueIndex = HashMap<ConstraintSymbol, ValueIndex>()
private val auxSymbol2valueIndex = HashMap<ConstraintSymbol, ValueIndex>()
private val symbol2index = HashMap<ConstraintSymbol, ArgumentRuleIndex>()
private val tag2rule = LinkedHashMap<String, Rule>()
val rules = ArrayList<Rule>()
private val slotIndices = ArrayList<SlotIndex>()
init {
buildIndex(handlers)
}
fun byTag(tag: String): Rule? = tag2rule[tag]
fun forOccurrence(occ: ConstraintOccurrence): Iterable<Rule> =
primarySymbol2valueIndex.get(occ.constraint().symbol())?.select(occ) ?:
(auxSymbol2valueIndex.get(occ.constraint().symbol())?.select(occ) ?: emptyList())
fun forOccurrence(occ: ConstraintOccurrence): Iterable<Rule> {
val ruleIndices = symbol2index[occ.constraint().symbol()]?.select(occ) ?: return emptyList()
return ruleIndices.allSetBits().map { idx -> rules[idx] }
}
/**
* Returns a pair of rule and bit mask with 1's marking matching slots.
*/
fun forOccurrenceWithMask(occ: ConstraintOccurrence): Iterable<Pair<Rule, BitSet>> {
val ruleIndices = symbol2index[occ.constraint().symbol()]?.select(occ) ?: return emptyList()
return ruleIndices.allSetBits().mapNotNull { idx -> slotIndices[idx][occ]?.let { mask -> rules[idx] to mask }}
}
override fun iterator(): Iterator<Rule> = tag2rule.values.iterator()
private fun buildIndex(handlers: Iterable<Handler>) {
// first, init the primary symbols value index
handlers
.flatMap { h -> h.primarySymbols() }
.forEach { symbol -> primarySymbol2valueIndex.getOrPut(symbol) { ValueIndex(symbol) } }
var pos = 0
for (h in handlers) {
val primSymbols = h.primarySymbols().toSet()
for (r in h.rules()) {
if (tag2rule.containsKey(r.tag())) {
throw IllegalStateException("duplicate rule tag ${r.tag()}")
}
tag2rule[r.tag()] = r
for (c in (r.headKept() + r.headReplaced())) {
val symbol = c.symbol()
if (symbol in primSymbols) {
primarySymbol2valueIndex[symbol]?.update(r, c)
for (rule in h.rules()) {
if (tag2rule.containsKey(rule.tag())) throw IllegalStateException("duplicate rule tag ${rule.tag()}")
} else {
auxSymbol2valueIndex.getOrPut(symbol) { ValueIndex(symbol) }.update(r, c)
}
tag2rule[rule.tag()] = rule
rules.add(rule)
val head = rule.headKept() + rule.headReplaced()
val symbol2mask = SlotIndex()
for ((bit, cst) in head.withIndex()) {
symbol2index.getOrPut(cst.symbol()) { ArgumentRuleIndex(cst.symbol()) }.update(pos, cst)
symbol2mask.update(cst, bit)
}
slotIndices.add(symbol2mask)
pos += 1
}
}
}
}
class SlotIndex() {
/**
* Represents a list of rules with at least one constraint matching the 'symbol'.
* The list is indexed by all values extracted from the constraint arguments.
* The list can be selected based on the argument(s) of a constraint occurrence.
*/
private class ValueIndex(val symbol: ConstraintSymbol) {
val symbol2mask = HashMap<Symbol, BitSet>()
val rules = ArrayList<Rule>()
val anySelectors = ArrayList<MutableMap<Any, BitSet>>()
val termSelectors = ArrayList<TermTrie<Int>>()
val wildcardSelectors = ArrayList<BitSet>()
init {
for (idx in 1..symbol.arity()) {
anySelectors.add(HashMap())
termSelectors.add(TermTrie())
wildcardSelectors.add(BitSet())
fun update(cst: Constraint, pos: Int) {
symbol2mask.getOrPut(cst.symbol()) { BitSet() }.set(pos)
}
operator fun get(occ: ConstraintOccurrence): BitSet? =
symbol2mask[occ.constraint().symbol()] // guaranteed to != null
}
fun update(rule: Rule, cst: Constraint) {
var bit = rules.indexOf(rule)
if (bit < 0) {
bit = rules.size
rules.add(rule)
/**
* Represents a list of rules with at least one constraint matching the 'symbol'.
* The list is indexed by all values extracted from the constraint arguments.
* The list can be selected based on the argument(s) of a constraint occurrence.
*/
inner class ArgumentRuleIndex(val symbol: ConstraintSymbol) {
val anySelectors = ArrayList<MutableMap<Any, BitSet>>()
val termSelectors = ArrayList<TermTrie<Int>>()
val wildcardSelectors = ArrayList<BitSet>()
init {
for (idx in 1..symbol.arity()) {
anySelectors.add(HashMap())
termSelectors.add(TermTrie())
wildcardSelectors.add(BitSet())
}
}
for ((idx, arg) in cst.arguments().withIndex()) {
val value2indices = anySelectors[idx]
when (arg) {
is MetaLogical<*> -> {
// all values should be accepted by a meta logical
wildcardSelectors[idx].set(bit)
}
is Term -> {
termSelectors.set(idx, termSelectors[idx].put(arg, bit))
}
is Any -> {
val bitset = value2indices.getOrPut(arg) { BitSet() }
bitset.set(bit)
}
else -> {
/* never happens */
throw NullPointerException()
fun update(bit: Int, cst: Constraint) {
for ((idx, arg) in cst.arguments().withIndex()) {
val value2indices = anySelectors[idx]
when (arg) {
is MetaLogical<*> ->
// all values should be accepted by a meta logical
wildcardSelectors[idx].set(bit)
is Term ->
termSelectors.set(idx, termSelectors[idx].put(arg, bit))
is Any ->
value2indices.getOrPut(arg) { BitSet() }.also { it.set(bit) }
else ->
throw NullPointerException() // never happens
}
}
}
}
fun select(occ: ConstraintOccurrence): Iterable<Rule> {
if (occ.constraint().symbol() != symbol) throw IllegalArgumentException()
/**
* Returns bit set where 1's indicate the indices of matching rules.
*/
fun select(occ: ConstraintOccurrence): BitSet {
if (occ.constraint().symbol() != symbol) throw IllegalArgumentException()
// initially select all rules
val upToBit = rules.size
val ruleIndices = BitSet(rules.size)
ruleIndices.set(0, upToBit)
// initially select all rules
val upToBit = rules.size
val ruleIndices = BitSet(rules.size)
ruleIndices.set(0, upToBit)
for ((idx, arg) in occ.arguments().withIndex()) {
val value2indices = anySelectors[idx]
val termIndices = termSelectors[idx]
val wildcardIndices = wildcardSelectors[idx]
if (arg is Logical<*> && !arg.isBound) {
// ALL values must be selected for a free logical
continue
}
val argVal = if (arg is Logical<*>) arg.findRoot().value() else arg
when (argVal) {
is Term -> {
val bits = wildcardIndices.clone() as BitSet
for (i in termIndices.lookupValues(argVal)) {
bits.set(i)
}
ruleIndices.and(bits)
for ((idx, arg) in occ.arguments().withIndex()) {
val value2indices = anySelectors[idx]
val termIndices = termSelectors[idx]
val wildcardIndices = wildcardSelectors[idx]
if (arg is Logical<*> && !arg.isBound) {
// ALL values must be selected for a free logical
continue
}
is Any -> {
if (value2indices.containsKey(argVal)) {
// ensure only rules with either matching values or wildcard arguments get selected
val argVal = if (arg is Logical<*>) arg.findRoot().value() else arg
when (argVal) {
is Term -> {
val bits = wildcardIndices.clone() as BitSet
bits.or(value2indices[argVal])
for (i in termIndices.lookupValues(argVal)) {
bits.set(i)
}
ruleIndices.and(bits)
}
is Any -> {
if (value2indices.containsKey(argVal)) {
// ensure only rules with either matching values or wildcard arguments get selected
val bits = wildcardIndices.clone() as BitSet
bits.or(value2indices[argVal])
ruleIndices.and(bits)
} else {
ruleIndices.and(wildcardIndices)
} else {
ruleIndices.and(wildcardIndices)
}
}
else -> {
/* never happens */
throw NullPointerException()
}
}
else -> {
/* never happens */
throw NullPointerException()
if (ruleIndices.isEmpty) {
break
}
}
if (ruleIndices.isEmpty) {
break
}
return ruleIndices
}
val result = ArrayList<Rule>(ruleIndices.cardinality())
var next = ruleIndices.nextSetBit(0)
while (next >= 0) {
result.add(rules.get(next))
next = ruleIndices.nextSetBit(next+1)
}
return result
}
}
}

View File

@ -61,7 +61,11 @@ class RuleMatcher(val rule: Rule) {
rn is ActiveFringeNode && rn.complete && rn.genId == genId }.map { rn ->
(rn as ActiveFringeNode).toMatchRule() }
fun expand(occ: ConstraintOccurrence): MatchFringe {
/**
* Expands the fringe by creating new leaf nodes that match the occurrence.
* Mask specifies possible slots for the occurrence.
*/
fun expand(occ: ConstraintOccurrence, mask: BitSet? = null): MatchFringe {
if (seen.contains(occ)) {
// constraint occurrence is reactivated
// there are nodes having occ in their paths
@ -74,12 +78,15 @@ class RuleMatcher(val rule: Rule) {
fn.unrelatedOrCopy(occ, genId + 1)
else
fn
}.appendAllTo(emptyConsList())
}.toConsList()
return MatchFringe(newNodes, seen, genId + 1)
} else {
val newNodes = nodes.asSequence().flatMap { it.expand(occ, genId + 1) }
return MatchFringe(newNodes.appendAllTo(nodes), seen.add(occ), genId + 1)
val newNodes = nodes.asSequence().filter { fn ->
mask == null || fn.matchesVacant(mask)
}.flatMap { fn ->
fn.expand(occ, genId + 1) }
return MatchFringe(newNodes.prependTo(nodes), seen.add(occ), genId + 1)
}
}
@ -87,6 +94,7 @@ class RuleMatcher(val rule: Rule) {
val newNodes = nodes.asSequence().mapNotNull { it.unrelatedOrNull(occ) }
return MatchFringe(newNodes.toConsList(), seen.remove(occ), genId + 1)
}
}
open inner class FringeNode(val subst: Subst,
@ -97,10 +105,8 @@ class RuleMatcher(val rule: Rule) {
* If the occurrence is already in the path, return empty sequence.
*/
fun expand(occ: ConstraintOccurrence, genId: Int): Sequence<ActiveFringeNode> {
val fringeNode = unrelatedOrNull(occ)
if (fringeNode == null) return emptySequence()
return fringeNode.vacant.allSetBits().map { idx ->
val unrelated = unrelatedOrNull(occ) ?: return emptySequence()
return unrelated.vacant.allSetBits().map { idx ->
idx to match(head[idx] !!, occ, subst) }.asSequence().mapNotNull { (idx, newSubst) ->
newSubst?.let { ActiveFringeNode(this, occ, idx, genId, it) }
}
@ -111,6 +117,8 @@ class RuleMatcher(val rule: Rule) {
*/
open fun unrelatedOrNull(occ: ConstraintOccurrence): FringeNode? = this
fun matchesVacant(mask: BitSet) = !mask.copyApply { and(vacant) }.isEmpty
/**
* Matches constraint and occurrence.
* Recursively processes all arguments, including terms.

View File

@ -22,6 +22,9 @@ import java.util.*
* @author Fedor Isakov
*/
inline fun BitSet.copyApply (f: BitSet.() -> Unit): BitSet =
(clone() as BitSet).apply(f)
fun bitSetOfOnes(size: Int): BitSet =
BitSet(size).also { it.set(0, size) }

View File

@ -35,10 +35,20 @@ fun <E> consListOf(vararg args: E): ConsList<E> {
return builder.build()
}
fun <E> Sequence<E>.appendAllTo(toList: ConsList<E>): ConsList<E> =
fold(toList) { list, e -> list.append(e) }
fun <E> Sequence<E>.prependTo(toList: ConsList<E>): ConsList<E> =
fold(toList) { list, e -> list.prepend(e) }
fun <E> Sequence<E>.toConsList(): ConsList<E> = appendAllTo(emptyConsList())
fun <E> Sequence<E>.toConsList(): ConsList<E> {
val builder = ConsList.factory<E>().newBuilder()
val var2 = iterator()
while (var2.hasNext()) {
val e = var2.next()
builder.add(e)
}
return builder.build()
}
fun <E> ConsList<E>.removeAt(idx: Int): ConsList<E> {
if (idx < 0) throw IllegalArgumentException("index < 0")