Skip to content

Commit e76b6b6

Browse files
Speed up randint array-bounds via chunking Lemire
1 parent 2d64b38 commit e76b6b6

1 file changed

Lines changed: 63 additions & 50 deletions

File tree

‎mkl_random/src/mkl_distributions.cpp‎

Lines changed: 63 additions & 50 deletions
Original file line numberDiff line numberDiff line change
@@ -2262,8 +2262,9 @@ static inline npy_uint64 irk_mulhi(npy_uint64 a, npy_uint64 b, npy_uint64 *lo)
22622262
/*
22632263
* Draw res[i] uniformly from [low[i], hi[i]] (inclusive) using Lemire's
22642264
* multiply-shift method (per-element bounds, same as NumPy).
2265-
* Words are generated in bulk by MKL; the rare rejected elements are
2266-
* gathered into `idx` (allocated lazily) and retried on the next round.
2265+
* Words are generated by MKL in cache-sized chunks; the rare rejected
2266+
* elements are gathered into `idx` (allocated lazily) and retried
2267+
* locally within each chunk.
22672268
* T is the result type, UT its unsigned counterpart,
22682269
* WT the raw-word type (s wraps to 0 for a full-range draw).
22692270
*/
@@ -2274,76 +2275,88 @@ static void irk_rand_bounded_broadcast(irk_state *state,
22742275
const T *low,
22752276
const T *hi)
22762277
{
2277-
npy_intp i = 0;
2278-
npy_intp k = 0;
2279-
npy_intp n_pending = 0;
2280-
npy_intp *idx = nullptr;
2281-
WT *words = nullptr;
2278+
npy_intp *idx = nullptr; /* reject indices */
22822279

22832280
if (len < 1)
22842281
return;
22852282

2286-
/* TODO: possible speedup :
2287-
* generate and consume words in cache-sized chunks
2288-
* instead of one full-length pass */
2289-
words = (WT *)mkl_malloc(len * sizeof(WT), 64);
2283+
/* Optimized path:
2284+
* cache-sized chunks instead of one full-length pass */
2285+
const npy_intp CHUNK_SIZE = 1 << 15; /* ~32K elements per chunk */
2286+
npy_intp chunk_cap = (len < CHUNK_SIZE) ? len : CHUNK_SIZE;
2287+
2288+
WT *words = (WT *)mkl_malloc(chunk_cap * sizeof(WT), 64);
22902289
assert(words != nullptr);
22912290

2292-
irk_uniform_bits_vec(state, len, words);
2291+
/* memoized reject threshold */
2292+
WT last_s = 0, last_t = 0;
22932293

2294-
for (i = 0; i < len; ++i) {
2295-
WT w = (WT)words[i];
2296-
/* diff cast back to UT so narrow types wrap (no signed promotion) */
2297-
UT d = (UT)(((UT)hi[i]) - ((UT)low[i]));
2298-
WT s = (WT)d + 1; /* 0 iff full range (32/64-bit only) */
2299-
WT result = w;
2300-
2301-
if (s != 0) {
2302-
WT lo = 0;
2303-
result = irk_mulhi(w, s, &lo);
2304-
if (lo < s) { /* rare */
2305-
WT t = (WT)(0 - s) % s;
2306-
if (lo < t) {
2307-
if (idx == nullptr) {
2308-
idx =
2309-
(npy_intp *)mkl_malloc(len * sizeof(npy_intp), 64);
2310-
assert(idx != nullptr);
2311-
}
2312-
idx[n_pending++] = i;
2313-
continue;
2314-
}
2315-
}
2316-
}
2317-
res[i] = (T)(((UT)low[i]) + (UT)result);
2318-
}
2319-
2320-
while (n_pending > 0) {
2321-
npy_intp wpos = 0;
2294+
for (npy_intp base = 0; base < len; base += chunk_cap) {
2295+
npy_intp chunk = (len - base < chunk_cap) ? (len - base) : chunk_cap;
2296+
npy_intp n_pending = 0;
23222297

2323-
irk_uniform_bits_vec(state, n_pending, words);
2298+
irk_uniform_bits_vec(state, chunk, words);
23242299

2325-
for (k = 0; k < n_pending; ++k) {
2326-
npy_intp j = idx[k];
2327-
WT w = (WT)words[k];
2300+
for (npy_intp i = 0; i < chunk; ++i) {
2301+
npy_intp j = base + i;
2302+
WT w = (WT)words[i];
2303+
/* diff cast back to UT so narrow types wrap (no signed promotion)
2304+
*/
23282305
UT d = (UT)(((UT)hi[j]) - ((UT)low[j]));
2329-
WT s = (WT)d + 1;
2306+
WT s = (WT)d + 1; /* 0 iff full range (32/64-bit only) */
23302307
WT result = w;
23312308

23322309
if (s != 0) {
23332310
WT lo = 0;
23342311
result = irk_mulhi(w, s, &lo);
2335-
if (lo < s) {
2336-
WT t = (WT)(0 - s) % s;
2337-
if (lo < t) {
2312+
/* t < s always (no lo < s branch) */
2313+
if (s != last_s) { /* recompute threshold */
2314+
last_t = (WT)(0 - s) % s;
2315+
last_s = s;
2316+
}
2317+
if (lo < last_t) { /* rare reject */
2318+
if (idx == nullptr) {
2319+
idx = (npy_intp *)mkl_malloc(
2320+
chunk_cap * sizeof(npy_intp), 64);
2321+
assert(idx != nullptr);
2322+
}
2323+
idx[n_pending++] = j;
2324+
continue;
2325+
}
2326+
}
2327+
res[j] = (T)(((UT)low[j]) + (UT)result);
2328+
}
2329+
2330+
/* retry the chunk's rejects locally with fresh words */
2331+
while (n_pending > 0) {
2332+
npy_intp wpos = 0;
2333+
2334+
irk_uniform_bits_vec(state, n_pending, words);
2335+
2336+
for (npy_intp k = 0; k < n_pending; ++k) {
2337+
npy_intp j = idx[k];
2338+
WT w = (WT)words[k];
2339+
UT d = (UT)(((UT)hi[j]) - ((UT)low[j]));
2340+
WT s = (WT)d + 1;
2341+
WT result = w;
2342+
2343+
if (s != 0) {
2344+
WT lo = 0;
2345+
result = irk_mulhi(w, s, &lo);
2346+
if (s != last_s) {
2347+
last_t = (WT)(0 - s) % s;
2348+
last_s = s;
2349+
}
2350+
if (lo < last_t) {
23382351
/* keep pending; wpos <= k so idx[k] read first */
23392352
idx[wpos++] = j;
23402353
continue;
23412354
}
23422355
}
2356+
res[j] = (T)(((UT)low[j]) + (UT)result);
23432357
}
2344-
res[j] = (T)(((UT)low[j]) + (UT)result);
2358+
n_pending = wpos;
23452359
}
2346-
n_pending = wpos;
23472360
}
23482361

23492362
if (idx != nullptr)

0 commit comments

Comments
 (0)