From 3c96e3d36c129e680b9d9e8ab5bbc85d0cfef099 Mon Sep 17 00:00:00 2001 From: Fedor Isakov Date: Sun, 16 Jul 2017 14:43:46 +0200 Subject: [PATCH] Rewrite TermTrie to not use recursion. Optimize put/lookup for speed. --- .../mps/logic/reactor/core/TermTrie.kt | 195 +++++++++--------- 1 file changed, 101 insertions(+), 94 deletions(-) diff --git a/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermTrie.kt b/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermTrie.kt index c6d6bdab..6780fdaf 100644 --- a/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermTrie.kt +++ b/reactor/Core/src/jetbrains/mps/logic/reactor/core/TermTrie.kt @@ -17,7 +17,6 @@ package jetbrains.mps.logic.reactor.core import com.github.andrewoma.dexx.collection.ConsList -import com.github.andrewoma.dexx.collection.ConsList.empty import com.github.andrewoma.dexx.collection.Map as PersMap import com.github.andrewoma.dexx.collection.Maps import jetbrains.mps.logic.reactor.util.* @@ -57,7 +56,7 @@ class TermTrie() { } - private lateinit var root: PathNode + private var root: PathNode init { root = PathNode(WILDCARD, 1) @@ -67,13 +66,13 @@ class TermTrie() { this.root = setRoot } - fun put(term: Term, value: T): TermTrie = TermTrie(putValue(root, value, IdHashSet(), term, empty())) + fun put(term: Term, value: T): TermTrie = TermTrie(putValue(term, value)) - fun remove(term: Term, value: T): TermTrie = TermTrie(removeValue(root, value, IdHashSet(), term, empty())) + fun remove(term: Term, value: T): TermTrie = TermTrie(removeValue(term, value)) fun lookupValues(term: Term): Iterable { val result = ArrayList() - visitMatching(root, IdHashSet(), term, empty()) { value -> result.add(value) } + visitMatching(term) { value -> result.add(value) } return result } @@ -83,89 +82,106 @@ class TermTrie() { return result } - private fun putValue(node: PathNode, value: T, seen: IdHashSet, term: Term, tail: ConsList): PathNode { - val deref = if (seen.contains(deref(term))) term else deref(term) - val arguments = deref.arguments() - val newTail = arguments.reversed().fold(tail) { list, t -> list.prepend(t) } + private fun putValue(matchTerm: Term, value: T): PathNode { + val seen = IdentityHashMap() + var nodeStack: ConsList> = cons(root) + val termList = arrayListOf(matchTerm) - val nextNode = node.nextOrDefault(symbolOrWildcard(deref)) { sym -> PathNode(sym, arguments.size) } + while (!termList.isEmpty()) { + val node = nodeStack.first() !! + val term = termList.removeAt(termList.size - 1) - //invariant: terms arity is fixed - assert(nextNode.arity == arguments.size) + // 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 } } + val arguments = deref.arguments().toList() + for(i in 1..arguments.size) { + termList.add(arguments[arguments.size - i]) + } - return if (newTail.isEmpty()) { - node.putNext(nextNode.addValue(value)) + val nextNode = node.nextOrDefault(symbolOrWildcard(deref)) { sym -> PathNode(sym, arguments.size) } + nodeStack = nodeStack.prepend(nextNode) + } + val head = nodeStack.first() !! + return nodeStack.tail().fold(head.addValue(value)) { nextNode, node -> node.putNext(nextNode) } + } + + private fun removeValue(matchTerm: Term, value: T): PathNode { + val seen = IdentityHashMap() + var nodeStack: ConsList> = cons(root) + val termList = arrayListOf(matchTerm) + + while (!termList.isEmpty()) { + val node = nodeStack.first() !! + val term = termList.removeAt(termList.size - 1) + + // 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 } } + val arguments = deref.arguments().toList() + for(i in 1..arguments.size) { + termList.add(arguments[arguments.size - i]) + } + + val nextNode = node.next(symbolOrWildcard(deref)) + if (nextNode == null) { + // term not found + return root + } + nodeStack = nodeStack.prepend(nextNode) + } + + val head = nodeStack.first() !! + val newHead = head.removeValue(value) + return if (head !== newHead) { + nodeStack.tail().fold(newHead) { nextNode, node -> node.putNext(nextNode) } } else { - node.putNext(putValue(nextNode, value, seen.add(term), newTail.first()!!, newTail.drop(1))) + // value not found + root } } - private fun removeValue(node: PathNode, value: T, seen: IdHashSet, term: Term, tail: ConsList): PathNode { - val deref = if (seen.contains(deref(term))) term else deref(term) - val arguments = deref.arguments() - val newTail = arguments.reversed().fold(tail) { list, t -> list.prepend(t) } + private fun visitMatching(matchTerm: Term, visitor: (T) -> Unit) { + val seen = IdentityHashMap() + val visitList = arrayListOf(root.to(cons(matchTerm))) - return node.next(symbolOrWildcard(deref))?.let { nextNode -> + while (!visitList.isEmpty()) { + val (node, termList) = visitList.removeAt(visitList.size - 1) + if (!termList.isEmpty) { + val term = termList.first() !! + val termTail = termList.tail() - //invariant: terms arity is fixed - assert(nextNode.arity == arguments.size) + // 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 } } - return if (newTail.isEmpty) { - val newNext = nextNode.removeValue(value) - if (newNext != nextNode) { - if (!newNext.hasValues()) { - node.removeNext(newNext) + val sym = symbolOrWildcard(deref) + if (sym == WILDCARD) { + if (!termTail.isEmpty) { + node.skipAllNext().forEach { n -> visitList.add(n.to(termTail)) } } else { - node.putNext(newNext) + node.allNext().forEach { visitAll(it, visitor) } } } else { - node - } - - } else { - val newNext = removeValue(nextNode, value, seen.add(term), newTail.first()!!, newTail.drop(1)) - if (newNext !== nextNode) { - if (!newNext.hasNext()) { - node.removeNext(nextNode) - - } else { - node.putNext(newNext) + node.next(WILDCARD)?.let { nn -> + nn.values().forEach(visitor) + if (!termTail.isEmpty) { + visitList.add(nn.to(termTail)) + } } + node.next(sym)?.let { nn -> + nn.values().forEach(visitor) - } else { - node - } - } - - } ?: node - } - - private fun visitMatching(node: PathNode, seen: IdHashSet, term: Term, tail: ConsList, visitor: (T) -> Unit) { - val deref = if (seen.contains(deref(term))) term else deref(term) - val sym = symbolOrWildcard(deref) - if (sym == WILDCARD) { - if (!tail.isEmpty) { - node.skipAllNext().forEach { visitMatching(it, seen.add(term), tail.first()!!, tail.drop(1), visitor) } - - } else { - node.allNext().forEach { visitAll(it, visitor) } - } - - } else { - node.next(sym)?.let { nn -> - nn.values().forEach(visitor) - val newTail = deref.arguments().reversed().fold(tail) { list, t -> list.prepend(t) } - if (!newTail.isEmpty) { - visitMatching(nn, seen.add(term), newTail.first()!!, newTail.drop(1), visitor) - } - } - node.next(WILDCARD)?.let { nn -> - nn.values().forEach(visitor) - if (!tail.isEmpty) { - visitMatching(nn, seen.add(term), tail.first()!!, tail.drop(1), visitor) + var newTail = termTail + val arguments = deref.arguments().toList() + for(i in 1..arguments.size) { + newTail = newTail.prepend(arguments[arguments.size - i]) + } + + if (!newTail.isEmpty) { + visitList.add(nn.to(newTail)) + } + } } } } @@ -193,37 +209,27 @@ class TermTrie() { } } - private class PathNode(val symbol: Any, val arity: Int) { + private class PathNode(val symbol: Any, + val arity: Int, + val next: PersMap>, + val values: IdHashSet) + { - private lateinit var next: PersMap> + constructor(symbol: Any, arity: Int) : + this(symbol, arity, Maps.of(), emptySet()) - private lateinit var values: IdHashSet - - init { - this.next = Maps.of() - this.values = emptySet() - } - - private constructor(symbol: Any, arity: Int, next: PersMap>, values: IdHashSet) : - this(symbol, arity) - { - this.next = next - this.values = values - } - - private constructor(copyFrom: PathNode) : - this(copyFrom.symbol, copyFrom.arity, copyFrom.next, copyFrom.values) - - private constructor(copyFrom: PathNode, setValues: IdHashSet) : + constructor(copyFrom: PathNode, setValues: IdHashSet) : this(copyFrom.symbol, copyFrom.arity, copyFrom.next, setValues) - private constructor(copyFrom: PathNode, setNext: PersMap>) : + constructor(copyFrom: PathNode, setNext: PersMap>) : this(copyFrom.symbol, copyFrom.arity, setNext, copyFrom.values) fun values(): Iterable = values fun hasValues(): Boolean = !values.isEmpty + fun hasValue(value: T): Boolean = values.contains(value) + fun next(symbol: Any): PathNode? = next[symbol] fun hasNext(): Boolean = !next.isEmpty @@ -248,10 +254,11 @@ class TermTrie() { fun addValue(value: T): PathNode = PathNode(this, values.add(value)) - fun removeValue(value: T): PathNode = PathNode(this, values.remove(value)) + fun removeValue(value: T): PathNode = + if (values.contains(value)) PathNode(this, values.remove(value)) else this - fun nextOrDefault(symbol: Any, - default: (sym: Any) -> PathNode): PathNode + inline fun nextOrDefault(symbol: Any, + default: (sym: Any) -> PathNode): PathNode { return next[symbol] ?: default(symbol) }