~/problems / Range queries / Fenwick tree (BIT)

Weighted random sampler with insert and remove

hard ~45 min Citadel

Build WeightedSampler(rng=None). rng is a random.Random to draw from (make your own when it's None). It holds items, each an id (any hashable) with a positive integer weight:

  • insert(id, weight): add the item. If id is already present, its weight becomes weight. Raise ValueError if weight <= 0.
  • remove(id): delete the item. Raise KeyError if it isn't present.
  • sample() -> id: return a present id at random, each with probability weight / total weight. Raise IndexError when empty. Sampling doesn't remove anything.
  • total() -> int: the current sum of weights. __len__(): the number of items.

Inserts, removes and samples come interleaved, up to 10**5 of each, and the item set changes between every sample. So sample can't scan all items (O(n)), and you can't rebuild a prefix-sum array after every update either. Aim for O(log n) per operation (amortised is fine). Use only rng for randomness, never the global random module, so that identically seeded samplers make identical draws.

s = WeightedSampler(random.Random(1))
s.insert("a", 1)
s.insert("b", 3)
s.total()      # 4
s.sample()     # "b" about 75% of the time, "a" about 25%
s.remove("b")
s.sample()     # "a"

Discussion

  • Prove that picking the slot k with prefix[k-1] <= r < prefix[k] for a uniform r in [0, total) gives probability weight[k] / total.
  • Why doesn't a heap or a sorted container help here?
Show hint

Put the items in numbered slots (reusing freed slots) and draw r uniformly from [0, total): the answer is the slot where the running sum of weights first exceeds r. You need a structure that updates one slot's weight and finds that slot in O(log n) each.

Topic: Fenwick tree (BIT). Point update + prefix sum in O(log n) with i & -i.

Read the visual guide
0:00
Ctrl ' run · Ctrl ↵ submit
esc