Sum of Values by Prefix
Problem
Design a store of string keys with integer values and two operations. put(key, value) sets the value of a key, replacing the old value if the key is already stored. total(prefix) returns the sum of the values of all keys that start with the prefix, or 0 if there are none. Both operations should cost time proportional to the length of their argument.
Examples
Input: put("apple", 3), total("ap"), put("app", 2), total("ap"), total("apple")
Output: 3, 5, 3
Why: after both puts, "ap" begins both keys and "apple" begins only one
Input: put("apple", 3), put("app", 2), put("apple", 5), total("ap")
Output: 7
Why: the second put of "apple" replaces 3 with 5 rather than adding to it
Input: put("apple", 3), total("b")
Output: 0
Why: edge case, no key starts with the prefix
Hints
0 / 3
Adding up every matching key at query time costs time for every key under the prefix. Try to have the answer ready before the question is asked.
In a prefix tree, the keys that start with a prefix are exactly the keys below that prefix's node. Each node can store the running sum of everything below it.
On put, work out how much the key's value changes, which is the new value minus the old one or zero, and add that difference to every node on the key's path. On total, walk to the prefix's node and return its stored sum.
Solution
Each node of the prefix tree keeps the sum of the values of every key whose path passes through it, so the answer for a prefix is simply the number stored at the prefix's node. A put changes that sum on exactly the nodes along the key's path, and only by the difference between the new value and the old one, which a separate dictionary remembers; this is what makes overwriting a key correct instead of double counting it. Both operations touch one node per character, so each costs O(L) for an argument of length L, and the tree uses O(total characters stored) space.
class PrefixSums:
def __init__(self):
self.root, self.value = {"#": 0}, {}
def put(self, key, val):
delta = val - self.value.get(key, 0) # overwrite, do not double count
self.value[key] = val
node = self.root
node["#"] += delta # the empty prefix covers all
for ch in key:
node = node.setdefault(ch, {"#": 0})
node["#"] += delta # every prefix of key gains delta
def total(self, prefix):
node = self.root
for ch in prefix:
if ch not in node:
return 0
node = node[ch]
return node["#"]
s = PrefixSums(); s.put("apple", 3)
print(s.total("ap")) # -> 3
s.put("app", 2)
print(s.total("ap"), s.total("apple")) # -> 5 3
s.put("apple", 5)
print(s.total("ap"), s.total("b")) # -> 7 0Stuck on the idea rather than the code? Prefix Search covers it.