题目
Given a blacklist B containing unique integers from [0, N), write a function to return a uniform random integer from [0, N) which is NOT in B.
Optimize it such that it minimizes the call to system’s Math.random().
答案
以下答案有一个test case过不了
class Solution {
List<Integer> pool = new ArrayList<>();
Set<Integer> blist = new HashSet<>();
Random rand = new Random();
int max;
int last_rand;
public Solution(int N, int[] blacklist) {
max = N;
last_rand = 0;
for(int i = 0; i < blacklist.length; i++) {
blist.add(blacklist[i]);
}
}
public int pick() {
int rand_num = rand.nextInt(max);
if(rand_num == 0 || rand_num == 1) {
//System.out.println("debug");
}
if(blist.contains(rand_num)) {
if(pool.size() == 0) {
// Generate a random number until it's not in black list
while(blist.contains(rand_num)) {
rand_num = rand.nextInt(max);
}
last_rand = rand_num;
pool.add(rand_num);
return rand_num;
}
// If black list contains the number that we just generated, we will use a number in the pool
// First, scale last_rand from range [0, N) to range(0, size of pool]
int index = (int)(((double)last_rand / max) * pool.size());
return pool.get(index);
}
else {
last_rand = rand_num;
pool.add(rand_num);
return rand_num;
}
}
}
另想了一个根据Random Pick with Weight改编的答案
既然有黑名单,那我们可以曲线救国,根据黑名单把[0, N)这个区间划分成k个子区间
用数组subs来存放这些区间,每个子区间的size为它们的权重
所以区间的权重总和为weight_sum
我们从[1, weight_sum]之间挑选一个随机数i, 即可以通过二分搜索推导出这个随机数i落在哪个区间中的具体哪个数字。详细代码看下面
class Solution {
int[] cumsum;
Random rand = new Random();
List<int[]> list;
public Solution(int N, int[] blacklist) {
Arrays.sort(blacklist);
list = new ArrayList<>();
int start = 0;
for(int b : blacklist) {
// Add interval [start, b - 1] if b - 1 >= start
// Otherwise, it's not a valid interval, then we assume start = b + 1
if(b - 1 >= start) {
list.add(new int[]{start, b - 1});
}
start = b + 1;
}
if(N - 1 >= start) {
list.add(new int[]{start, N - 1});
}
cumsum = new int[list.size()];
cumsum[0] = list.get(0)[1] - list.get(0)[0] + 1;
for(int i = 1; i < list.size(); i++) {
cumsum[i] = list.get(i)[1] - list.get(i)[0] + 1 + cumsum[i - 1];
}
}
public int pick() {
int left = 0, right = cumsum.length;
int target = rand.nextInt(cumsum[cumsum.length - 1]) + 1;
// We want to find the index into cumsum array, such that target falls into
while(left < right) {
int mid = (left + right) / 2;
int mid_val = cumsum[mid];
if(mid_val >= target) {
right = mid;
}
else {
left = mid + 1;
}
}
target -= 1;
if(left == 0) {
return list.get(left)[0] + (target - 0);
}
else {
return list.get(left)[0] + (target - cumsum[left - 1]);
}
}
}