Further optimizing term trie: drop index mask

This commit is contained in:
Fedor Isakov 2024-10-02 15:17:14 +02:00
parent ffa9ebf3f8
commit 4a97d3e89f
3 changed files with 35 additions and 101 deletions

View File

@ -180,7 +180,7 @@ class RuleIndex(rules: Iterable<Rule>, profiler: Profiler? = null) : Iterable<Ru
// all values should be accepted by a meta logical
wildcardSelectors[argIdx].add(ruleBit)
is Term ->
termSelectors[argIdx].put(arg, ruleBit,headPos)
termSelectors[argIdx].put(arg, ruleBit, headPos)
is Any ->
value2indices.getOrPut(arg) { hashSetOf() }.add(ruleBit to headPos)
else ->
@ -236,12 +236,10 @@ class RuleIndex(rules: Iterable<Rule>, profiler: Profiler? = null) : Iterable<Ru
candidateRuleBits.clear()
candidateRuleBits.addAll(wildcardIndices)
val currentSelector = if (refined) selectedRuleBits else null
val argVal = if (arg is Logical<*>) arg.findRoot().value() else arg
when (argVal) {
is Term ->
termIndices.forValuesWithIndex(argVal, currentSelector) { headPos, ruleBit ->
termIndices.forValuesWithIndex(argVal) { headPos, ruleBit ->
candidateRuleBits.add(ruleBit)
slotVotes.getOrPut(ruleBit to headPos) { BitSet() }.set(argIdx)
}

View File

@ -16,12 +16,11 @@
package jetbrains.mps.logic.reactor.util
import gnu.trove.map.hash.TIntIntHashMap
import gnu.trove.map.hash.TIntObjectHashMap
import jetbrains.mps.logic.reactor.logical.VarSymbol
import jetbrains.mps.unification.Term
import java.util.*
import kotlin.collections.ArrayList
import kotlin.collections.HashSet
/**
* @author Fedor Isakov
@ -82,23 +81,23 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
override fun lookupValues(term: Term): Iterable<T> {
val result = ArrayList<T>()
visitMatching(term, root.allIndexMask()) { value, index -> result.add(value) }
visitMatching(term) { value, index -> result.add(value) }
return result
}
override fun lookupValues(term: Term, indexMask: IndexMask): Iterable<T> {
val result = ArrayList<T>()
visitMatching(term, indexMask) { value, index -> result.add(value) }
visitMatching(term) { value, index -> if (indexMask.get(index)) { result.add(value) } }
return result
}
override fun forValuesWithIndex(term: Term, indexMask: IndexMask?, callback: (T, Int) -> Unit) {
visitMatching(term, indexMask) { value, index -> callback(value, index) }
override fun forValuesWithIndex(term: Term, callback: (T, Int) -> Unit) {
visitMatching(term) { value, index -> callback(value, index) }
}
override fun allValues(): Iterable<T> {
val result = ArrayList<T>()
visitAll(root, root.allIndexMask(), { value, index -> result.add(value) })
visitAll(root, null, { value, index -> result.add(value) })
return result
}
@ -112,24 +111,18 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
val seen = IdentityHashMap<Term, Term>()
val nodeStack = arrayDequeOf(root)
val termStack = arrayDequeOf(matchTerm)
val argPool = ArrayList<Term>(16)
while (!termStack.isEmpty()) {
val node = nodeStack.peek()
val term = termStack.pop()
if (index >= 0) {
node.addIndex(index)
}
// dereferece the term only if it hasn't been dereferenced before
val deref = deref(term).let { dt -> seen[dt]?.run { term } ?: dt.apply { seen[dt] = term } }
argPool.addAll(deref.arguments())
for(i in 1..argPool.size) {
termStack.push(argPool[argPool.size - i])
val argSize = deref.arguments().size
for(i in 1..argSize) {
termStack.push(deref.arguments()[argSize - i])
}
val nextNode = node.nextOrDefault(symbolOrWildcard(deref)) { sym -> PathNode(sym, argPool.size) }
val nextNode = node.nextOrDefault(symbolOrWildcard(deref)) { sym -> PathNode(sym, argSize) }
nodeStack.push(nextNode)
argPool.clear()
}
val head = nodeStack.pop()
@ -177,17 +170,17 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
/**
* Call the visitor function with values of the given node and all its direct and indirect successors.
*/
private fun visitAll(node: PathNode<T>, indexMask: IndexMask, visitor: (T, Int) -> Unit) {
node.forEachValueWithIndex(indexMask, visitor)
node.allNext(indexMask).forEach { visitAll(it, indexMask, visitor) }
private fun visitAll(node: PathNode<T>, indexMask: IndexMask?, visitor: (T, Int) -> Unit) {
node.forEachValue { value, index -> if (indexMask?.get(index) ?: true) { visitor(value, index) } }
node.allNext().forEach { visitAll(it, indexMask, visitor) }
}
private fun collectAllLeaves(node: PathNode<T>, indexMask: IndexMask?, leafs: MutableList<PathNode<T>>) {
private fun collectAllLeaves(node: PathNode<T>, leafs: MutableList<PathNode<T>>) {
val nodeStack = arrayListOf(node)
while (nodeStack.isNotEmpty()) {
val n = nodeStack.pop()
if (n.isLeaf()) leafs.add(n)
n.allNext(indexMask).forEach {
n.allNext().forEach {
nodeStack.push(it)
}
}
@ -214,7 +207,7 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
* Given a pattern term, which may contain variables that are treated as wildcards,
* call the passed visitor function with values of all matching nodes and all the nodes that precede them.
*/
private fun visitMatching(pattern: Term, indexMask: IndexMask?, visitor: (T, Int) -> Unit) {
private fun visitMatching(pattern: Term, visitor: (T, Int) -> Unit) {
val seenNonLeaf = IdentityHashMap<Term, Term>()
val canonicTerms = IdentityHashMap<Term, Pair<Term, Any>>()
val visitStack = arrayListOf(listOf(root) to listOf(pattern))
@ -238,13 +231,13 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
if (!ptnTail.isEmpty()) {
// skip the current node
// match the patterns tail with the current node's direct successors
val (wcdTerms, wcdBases) = base.terms2bases(indexMask)
val (wcdTerms, wcdBases) = base.terms2bases()
wcdTerms.forEach { if (it.isLeaf()) allLeaves.add(it) /*it.forEachValueWithIndex(indexMask, visitor)*/ }
visitStack.push(wcdBases to ptnTail)
} else {
// wildcard consumes the rest of the trie
base.allNext(indexMask).forEach { collectAllLeaves(it, indexMask, allLeaves) }
base.allNext().forEach { collectAllLeaves(it, allLeaves) }
}
} else {
@ -274,15 +267,7 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
}
}
if (indexMask != null) {
val iter = indexMask.iterator()
while (iter.hasNext()) {
val index = iter.next()
allLeaves.forEach { it.forEachValueWithIndex(index, visitor) }
}
} else {
allLeaves.forEach { it.forEachValue(visitor) }
}
allLeaves.forEach { it.forEachValue(visitor) }
}
/**
@ -291,44 +276,11 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
private class PathNode<T>(val symbol: Any,
val arity: Int)
{
private val indexBits = BitSet() // map of index cardinalities
private val next = HashMap<Any, PathNode<T>>(8)
private var indexedValues: TIntObjectHashMap<HashSet<T>>? = null
fun allIndexMask(): IndexMask = emptyIndexMask().also { it.addAll(indexBits) }
fun addIndex(index: Int) {
assert(index >= 0)
// val card = if (indexBits.contains(index)) indexBits.get(index) else 0
indexBits.set(index)
}
//
// fun removeIndex(index: Int) {
// assert(index >= 0)
// val card = if (indexBits.contains(index)) indexBits.get(index) else 0
// assert(card > 0)
// if (card > 1) indexBits.set(index) else indexBits.remove(index)
// }
private val next = HashMap<Any, PathNode<T>>(4)
private val indexedValues = arrayListOf<Pair<Int,T>>()
fun forEachValue(callback: (T, Int) -> Unit ) {
indexedValues?.forEachEntry { index, values ->
values.forEach { callback(it, index) }
true
}
}
fun forEachValueWithIndex(indexMask: IndexMask, callback: (T, Int) -> Unit ) {
val iter = indexMask.iterator()
while (iter.hasNext()) {
val index = iter.next()
if (indexedValues?.containsKey(index) ?: false) { indexedValues?.get(index)?.forEach { callback(it, index) } }
}
}
fun forEachValueWithIndex(index: Int, callback: (T, Int) -> Unit ) {
if (indexedValues?.containsKey(index) ?: false) { indexedValues?.get(index)?.forEach { callback(it, index) } }
indexedValues.forEach { (index, value) -> callback(value, index) }
}
fun isLeaf(): Boolean = next.isEmpty()
@ -338,18 +290,9 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
/**
* Returns all trie nodes that are direct successors of this one.
*/
fun allNext(indexMask: IndexMask?): List<PathNode<T>> {
fun allNext(): List<PathNode<T>> {
val res = arrayListOf<PathNode<T>>()
next.values.forEach { n ->
// containsAny ~~ NOT( all( NOT(contains))
if (indexMask != null) {
if (!n.indexBits.stream().allMatch { idx -> !indexMask.contains(idx) }) {
res.add(n)
}
} else {
res.add(n)
}
}
res.addAll(next.values)
return res
}
@ -370,7 +313,7 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
* The size of a list representing a flattened term is equal to
* the sum of arities of all symbols in this list plus 1.
*/
fun terms2bases(indexMask: IndexMask?): Pair<List<PathNode<T>>, List<PathNode<T>>> {
fun terms2bases(): Pair<List<PathNode<T>>, List<PathNode<T>>> {
// for every current node there is a number
// initially 0
@ -380,7 +323,7 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
val terms = ArrayList<PathNode<T>>()
val bases = ArrayList<PathNode<T>>()
val stack = arrayListOf(allNext(indexMask) to 0)
val stack = arrayListOf(allNext() to 0)
while (stack.isNotEmpty()) {
val (nn, count) = stack.pop()
@ -391,7 +334,7 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
} else {
terms.add(n)
stack.push(n.allNext(indexMask) to (newCount - 1))
stack.push(n.allNext() to (newCount - 1))
}
}
}
@ -402,21 +345,12 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
fun addValue(value: T, index: Int) {
assert (index >= 0)
addIndex(index)
if (indexedValues == null) indexedValues = TIntObjectHashMap<HashSet<T>>()
(indexedValues?.get(index) ?: hashSetOf<T>().also { indexedValues?.put(index, it) }).add(value)
indexedValues.add(index to value)
}
fun removeValue(value: T, index: Int): Boolean {
assert (index >= 0)
if (indexedValues?.get(index)?.remove(value) ?: false) {
if (indexedValues?.get(index)?.isEmpty() ?: false) {
indexedValues?.remove(index)
}
// removeIndex(index)
return true
}
return false
return indexedValues.remove(index to value)
}
inline fun nextOrDefault(symbol: Any,
@ -431,10 +365,12 @@ class ClassicIndexedTermTrie<T> : IndexedTermTrie<T> {
}
fun dropNext(drop: PathNode<T>) {
this.next.remove(drop)
this.next.remove(drop.symbol)
}
// override fun toString(): String =
// "${symbol}/${arity} (${values.keys.joinToString(", ")}) > [${next.map { it.symbol }.joinToString(", ")}] "
}
}

View File

@ -34,7 +34,7 @@ interface IndexedTermTrie<T> : TermTrie<T> {
fun lookupValues(term: Term, indexMask: IndexMask): Iterable<T>
fun forValuesWithIndex(term: Term, indexMask: IndexMask?, callback: (T, Int) -> Unit)
fun forValuesWithIndex(term: Term, callback: (T, Int) -> Unit)
fun allValues(indexMask: IndexMask): Iterable<T>