From 4a97d3e89fc7da4a3b86bdb8eb6dc37da5a2c7e4 Mon Sep 17 00:00:00 2001 From: Fedor Isakov Date: Wed, 2 Oct 2024 15:17:14 +0200 Subject: [PATCH] Further optimizing term trie: drop index mask --- .../mps/logic/reactor/core/RuleIndex.kt | 6 +- .../reactor/util/ClassicIndexedTermTrie.kt | 128 +++++------------- .../mps/logic/reactor/util/IndexedTermTrie.kt | 2 +- 3 files changed, 35 insertions(+), 101 deletions(-) diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt index d5359e6d..61dbfc91 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt @@ -180,7 +180,7 @@ class RuleIndex(rules: Iterable, profiler: Profiler? = null) : Iterable - 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, profiler: Profiler? = null) : Iterable) 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) } diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt index 2f5ef88e..1f677d29 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt @@ -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 : IndexedTermTrie { override fun lookupValues(term: Term): Iterable { val result = ArrayList() - 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 { val result = ArrayList() - 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 { val result = ArrayList() - 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 : IndexedTermTrie { val seen = IdentityHashMap() val nodeStack = arrayDequeOf(root) val termStack = arrayDequeOf(matchTerm) - val argPool = ArrayList(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 : IndexedTermTrie { /** * Call the visitor function with values of the given node and all its direct and indirect successors. */ - private fun visitAll(node: PathNode, indexMask: IndexMask, visitor: (T, Int) -> Unit) { - node.forEachValueWithIndex(indexMask, visitor) - node.allNext(indexMask).forEach { visitAll(it, indexMask, visitor) } + private fun visitAll(node: PathNode, 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, indexMask: IndexMask?, leafs: MutableList>) { + private fun collectAllLeaves(node: PathNode, leafs: MutableList>) { 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 : IndexedTermTrie { * 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() val canonicTerms = IdentityHashMap>() val visitStack = arrayListOf(listOf(root) to listOf(pattern)) @@ -238,13 +231,13 @@ class ClassicIndexedTermTrie : IndexedTermTrie { 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 : IndexedTermTrie { } } - 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 : IndexedTermTrie { private class PathNode(val symbol: Any, val arity: Int) { - private val indexBits = BitSet() // map of index cardinalities - - private val next = HashMap>(8) - - private var indexedValues: TIntObjectHashMap>? = 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>(4) + private val indexedValues = arrayListOf>() 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 : IndexedTermTrie { /** * Returns all trie nodes that are direct successors of this one. */ - fun allNext(indexMask: IndexMask?): List> { + fun allNext(): List> { val res = arrayListOf>() - 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 : IndexedTermTrie { * 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>> { + fun terms2bases(): Pair>, List>> { // for every current node there is a number // initially 0 @@ -380,7 +323,7 @@ class ClassicIndexedTermTrie : IndexedTermTrie { val terms = ArrayList>() val bases = ArrayList>() - 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 : IndexedTermTrie { } 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 : IndexedTermTrie { fun addValue(value: T, index: Int) { assert (index >= 0) - addIndex(index) - if (indexedValues == null) indexedValues = TIntObjectHashMap>() - (indexedValues?.get(index) ?: hashSetOf().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 : IndexedTermTrie { } fun dropNext(drop: PathNode) { - this.next.remove(drop) + this.next.remove(drop.symbol) } // override fun toString(): String = // "${symbol}/${arity} (${values.keys.joinToString(", ")}) > [${next.map { it.symbol }.joinToString(", ")}] " + + } } diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt index 26145449..f70b8d43 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt @@ -34,7 +34,7 @@ interface IndexedTermTrie : TermTrie { fun lookupValues(term: Term, indexMask: IndexMask): Iterable - fun forValuesWithIndex(term: Term, indexMask: IndexMask?, callback: (T, Int) -> Unit) + fun forValuesWithIndex(term: Term, callback: (T, Int) -> Unit) fun allValues(indexMask: IndexMask): Iterable