numpy version of sample()
(rand,
t,
w_logits,
num_samples,
single_jitter=False,
deterministic_center=False)
| 219 | |
| 220 | |
| 221 | def sample_np(rand, |
| 222 | t, |
| 223 | w_logits, |
| 224 | num_samples, |
| 225 | single_jitter=False, |
| 226 | deterministic_center=False): |
| 227 | """ |
| 228 | numpy version of sample() |
| 229 | """ |
| 230 | eps = np.finfo(np.float32).eps |
| 231 | |
| 232 | # Draw uniform samples. |
| 233 | if not rand: |
| 234 | if deterministic_center: |
| 235 | pad = 1 / (2 * num_samples) |
| 236 | u = np.linspace(pad, 1. - pad - eps, num_samples) |
| 237 | else: |
| 238 | u = np.linspace(0, 1. - eps, num_samples) |
| 239 | u = np.broadcast_to(u, t.shape[:-1] + (num_samples,)) |
| 240 | else: |
| 241 | # `u` is in [0, 1) --- it can be zero, but it can never be 1. |
| 242 | u_max = eps + (1 - eps) / num_samples |
| 243 | max_jitter = (1 - u_max) / (num_samples - 1) - eps |
| 244 | d = 1 if single_jitter else num_samples |
| 245 | u = np.linspace(0, 1 - u_max, num_samples) + \ |
| 246 | np.random.rand(*t.shape[:-1], d) * max_jitter |
| 247 | |
| 248 | return invert_cdf_np(u, t, w_logits) |
| 249 | |
| 250 | |
| 251 | def sample_intervals(rand, |
no test coverage detected