1. You want to calculate the l2-nearest neighbor to y from a set of x1, x2, ... xn. Where each point is a d-dimensional vector.
2. Naive approach will take O(nd) computations. For very large d, say 1e4, this is very expensive
Multi-arm bandits idea:
3. Estimate the distance of y to each point (xi) by sampling only along log(d) dimensions.
4. Keep the nearest n/2 points (throw away the rest)
5. Repeat step 3 and 4 until you have one point left
6. The total run time = nlog(d) + nlog(d)/2 + nlog(d)/4.... = O (nlog(d)log(n))
This is a huge improvement over O(nd) for large d (which is usually the case today)
This paper (https://ar5iv.org/abs/1805.08321) from Stanford pioneered this idea for nearest neighbors and many such household problems. It also theoretically proves that the above approach gives the right answer and empirically gets huge (100x) speed boost on large datasets.
The k-medoid paper is based off of this work :)