From 964b2cea29a5918c66cd8b5705d39bcbd71d4b17 Mon Sep 17 00:00:00 2001 From: Fedor Isakov Date: Mon, 1 Jul 2024 15:28:20 +0200 Subject: [PATCH] Minor optimizations and refactoring (TIntSet -> BitSet) --- .../mps/logic/reactor/core/RuleIndex.kt | 16 ++++---- .../core/internal/ReteRuleMatcherImpl.kt | 2 +- .../mps/logic/reactor/util/BitSet.kt | 38 +++++++++++++++---- .../reactor/util/ClassicIndexedTermTrie.kt | 34 ++++++++++------- .../mps/logic/reactor/util/IndexedTermTrie.kt | 7 ++-- reactor/Test/test/TestIndexedTermTrie.kt | 9 ++++- 6 files changed, 70 insertions(+), 36 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 2e03a993..357b1dea 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/RuleIndex.kt @@ -28,8 +28,8 @@ import java.util.* import kotlin.collections.ArrayList import kotlin.collections.HashMap -typealias RuleBits = IndexMask -fun ruleBitsOf() = indexMaskOf() +typealias RuleBits = BitSet +fun emptyRuleBits() = BitSet() /** * A container for [Rule] instances with the ability to look up by [ConstraintOccurrence]. @@ -47,8 +47,7 @@ class RuleIndex(): Iterable, RuleLookup private val symbol2index = HashMap() - // TODO: is linked hashmap really necessary? might be inefficient - private val tag2rule = LinkedHashMap() + private val tag2rule = HashMap() private val tag2bit = TObjectIntHashMap() @@ -111,6 +110,7 @@ class RuleIndex(): Iterable, RuleLookup } } + @Deprecated(message = "update functionality to be dropped") fun updateIndexFromRules(rules: Iterable) { val removedTags = allRules.map { it.rule.uniqueTag() }.toHashSet() rules.map { it.uniqueTag() }.forEach{ removedTags.remove(it) } @@ -206,7 +206,7 @@ class RuleIndex(): Iterable, RuleLookup */ inner class ArgumentRuleIndex(val symbol: ConstraintSymbol) { - val symbolSelector = ruleBitsOf() + val symbolSelector = emptyRuleBits() // value -> List of Pairs of rule bits and head positions val anySelectors = ArrayList>>>() @@ -217,7 +217,7 @@ class RuleIndex(): Iterable, RuleLookup for (idx in 1..symbol.arity()) { anySelectors.add(HashMap()) termSelectors.add(indexedTermTrie()) - wildcardSelectors.add(ruleBitsOf()) + wildcardSelectors.add(emptyRuleBits()) } } @@ -268,7 +268,7 @@ class RuleIndex(): Iterable, RuleLookup val slotVotes = HashMap, BitSet>() val wildcardSlots = BitSet() - val selectedRuleBits = ruleBitsOf() + val selectedRuleBits = emptyRuleBits() for ((argIdx, arg) in occ.arguments().withIndex()) { if (arg is Logical<*> && !arg.isBound) { @@ -277,7 +277,7 @@ class RuleIndex(): Iterable, RuleLookup continue } - val candidateRuleBits = ruleBitsOf() + val candidateRuleBits = emptyRuleBits() val value2indices = anySelectors[argIdx] val termIndices = termSelectors[argIdx] val wildcardIndices = wildcardSelectors[argIdx] 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 d94f2c29..40e0c715 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 @@ -330,7 +330,7 @@ internal class ReteRuleMatcherImpl(private var ruleLookup: RuleLookup?, continue } - val it = headPosMask.allSetBits() + val it = headPosMask.iterator() while (it.hasNext()) { val headPos = it.next() if (n.occupiesHeadPosition(headPos)) continue diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/util/BitSet.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/util/BitSet.kt index 0ec16eca..cc23abed 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/util/BitSet.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/BitSet.kt @@ -17,6 +17,7 @@ package jetbrains.mps.logic.reactor.util import gnu.trove.iterator.TIntIterator +import gnu.trove.set.TIntSet import java.util.* import kotlin.NoSuchElementException @@ -33,16 +34,37 @@ fun bitSetOfOnes(size: Int): BitSet = fun bitSet(setBit: Int): BitSet = BitSet().apply { set(setBit) } -fun bitSet(setBits: Iterable): BitSet= - BitSet().apply { setBits.forEach { set(it) } } +//fun bitSet(setBits: Iterable): BitSet= +// BitSet().apply { setBits.forEach { set(it) } } +// +//fun BitSet.setBit(bit: Int): BitSet = +// (clone() as BitSet).apply { set(bit) } +// +//fun BitSet.clearBit(bit: Int): BitSet = +// (clone() as BitSet).apply { clear(bit) } -fun BitSet.setBit(bit: Int): BitSet = - (clone() as BitSet).apply { set(bit) } +fun BitSet.contains(bit: Int) = get(bit) +fun BitSet.add(bit: Int) = set(bit) +fun BitSet.remove(bit: Int) = clear(bit) +fun BitSet.addAll(that: BitSet) = or(that) +fun BitSet.addAll(that: TIntSet) { + val iter = that.iterator() + while (iter.hasNext()) { + set(iter.next()) + } +} +fun BitSet.retainAll(that: BitSet) = and(that) +fun BitSet.retainAll(that: TIntSet) { + var bit = nextSetBit(0) + while (bit >= 0) { + if (!that.contains(bit)) { + clear(bit) + } + bit = nextSetBit(bit + 1) + } +} -fun BitSet.clearBit(bit: Int): BitSet = - (clone() as BitSet).apply { clear(bit) } - -fun BitSet.allSetBits(): TIntIterator = object : TIntIterator { +fun BitSet.iterator(): TIntIterator = object : TIntIterator { var next = nextSetBit(0) 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 cc0333b0..bae4c3cf 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/ClassicIndexedTermTrie.kt @@ -55,10 +55,6 @@ class ClassicIndexedTermTrie : IndexedTermTrie { private val WILDCARD = object : Any() { override fun toString() = "WILDCARD" } - - private val KEYHOLDER = object : Any() { - override fun toString() = "KEYHOLDER" - } } @@ -69,7 +65,7 @@ class ClassicIndexedTermTrie : IndexedTermTrie { } override fun put(term: Term, value: T) { - putValue(term, -1, value) + putValue(term, 0, value) } override fun put(term: Term, index: Int, value: T) { @@ -77,7 +73,7 @@ class ClassicIndexedTermTrie : IndexedTermTrie { } override fun remove(term: Term, value: T) { - removeValue(term, -1, value) + removeValue(term, 0, value) } override fun remove(term: Term, index: Int, value: T) { @@ -121,7 +117,9 @@ class ClassicIndexedTermTrie : IndexedTermTrie { val node = nodeStack.peek() val term = termStack.pop() - node.addIndex(index) + 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 } } @@ -276,10 +274,13 @@ class ClassicIndexedTermTrie : IndexedTermTrie { } } - indexMask?.forEach { index -> - allLeaves.forEach { it.forEachValueWithIndex(index, visitor) } - true - } ?: run { + 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) } } } @@ -296,14 +297,16 @@ class ClassicIndexedTermTrie : IndexedTermTrie { private val indexedValues = TIntObjectHashMap>() - fun allIndexMask(): IndexMask = indexCardinalities.keySet() + fun allIndexMask(): IndexMask = emptyIndexMask().also { it.addAll(indexCardinalities.keySet()) } fun addIndex(index: Int) { + assert(index >= 0) val card = if (indexCardinalities.contains(index)) indexCardinalities.get(index) else 0 indexCardinalities.put(index, card + 1) } fun removeIndex(index: Int) { + assert(index >= 0) val card = if (indexCardinalities.contains(index)) indexCardinalities.get(index) else 0 assert(card > 0) if (card > 1) indexCardinalities.put(index, card - 1) else indexCardinalities.remove(index) @@ -317,9 +320,10 @@ class ClassicIndexedTermTrie : IndexedTermTrie { } fun forEachValueWithIndex(indexMask: IndexMask, callback: (T, Int) -> Unit ) { - indexMask.forEach { index -> + val iter = indexMask.iterator() + while (iter.hasNext()) { + val index = iter.next() if (indexedValues.containsKey(index)) { indexedValues.get(index).forEach { callback(it, index) } } - true } } @@ -397,11 +401,13 @@ class ClassicIndexedTermTrie : IndexedTermTrie { fun addValue(value: T, index: Int) { + assert (index >= 0) addIndex(index) (indexedValues.get(index) ?: hashSetOf().also { indexedValues.put(index, it) }).add(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) 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 94341f5b..26145449 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/util/IndexedTermTrie.kt @@ -16,16 +16,15 @@ package jetbrains.mps.logic.reactor.util -import gnu.trove.set.TIntSet -import gnu.trove.set.hash.TIntHashSet import jetbrains.mps.unification.Term +import java.util.BitSet /** * @author Fedor Isakov */ -typealias IndexMask = TIntSet -fun indexMaskOf(vararg indices: Int): TIntSet = TIntHashSet(intArrayOf(*indices)) +typealias IndexMask = BitSet +fun emptyIndexMask() = BitSet() interface IndexedTermTrie : TermTrie { diff --git a/reactor/Test/test/TestIndexedTermTrie.kt b/reactor/Test/test/TestIndexedTermTrie.kt index a8526d4b..a67b9959 100644 --- a/reactor/Test/test/TestIndexedTermTrie.kt +++ b/reactor/Test/test/TestIndexedTermTrie.kt @@ -1,4 +1,7 @@ -import jetbrains.mps.logic.reactor.util.indexMaskOf +import gnu.trove.set.TIntSet +import gnu.trove.set.hash.TIntHashSet +import jetbrains.mps.logic.reactor.util.addAll +import jetbrains.mps.logic.reactor.util.emptyIndexMask import jetbrains.mps.logic.reactor.util.indexedTermTrie import jetbrains.mps.logic.reactor.util.termTrie import jetbrains.mps.unification.test.MockTermsParser @@ -6,6 +9,7 @@ import jetbrains.mps.unification.test.MockTermsParser.* import org.junit.Assert import org.junit.Assert.* import org.junit.Test +import java.util.* /* * Copyright 2014-2021 JetBrains s.r.o. @@ -26,6 +30,9 @@ import org.junit.Test /** * @author Fedor Isakov */ + +fun indexMaskOf(vararg indices: Int) = emptyIndexMask().also { it.addAll(TIntHashSet(intArrayOf(*indices))) } + class TestIndexedTermTrie { @Test