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. Ifidis already present, its weight becomesweight. RaiseValueErrorifweight <= 0.remove(id): delete the item. RaiseKeyErrorif it isn't present.sample() -> id: return a present id at random, each with probabilityweight / total weight. RaiseIndexErrorwhen 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
kwithprefix[k-1] <= r < prefix[k]for a uniformrin[0, total)gives probabilityweight[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.