codekingpro/portable-devtools
115k
1import os
2import sys
3from os.path import join
4
5import pytest
6
7import numpy as np
8from numpy.random import (
9 MT19937,
10 PCG64,
11 PCG64DXSM,
12 SFC64,
13 Generator,
14 Philox,
15 RandomState,
16 SeedSequence,
17 default_rng,
18)
19from numpy.random._common import interface
20from numpy.testing import (
21 assert_allclose,
22 assert_array_equal,
23 assert_equal,
24 assert_raises,
25)
26
27try:
28 import cffi # noqa: F401
29
30 MISSING_CFFI = False
31except ImportError:
32 MISSING_CFFI = True
33
34try:
35 import ctypes # noqa: F401
36
37 MISSING_CTYPES = False
38except ImportError:
39 MISSING_CTYPES = False
40
41if sys.flags.optimize > 1:
42 # no docstrings present to inspect when PYTHONOPTIMIZE/Py_OptimizeFlag > 1
43 # cffi cannot succeed
44 MISSING_CFFI = True
45
46
47pwd = os.path.dirname(os.path.abspath(__file__))
48
49
50def assert_state_equal(actual, target):
51 for key in actual:
52 if isinstance(actual[key], dict):
53 assert_state_equal(actual[key], target[key])
54 elif isinstance(actual[key], np.ndarray):
55 assert_array_equal(actual[key], target[key])
56 else:
57 assert actual[key] == target[key]
58
59
60def uint32_to_float32(u):
61 return ((u >> np.uint32(8)) * (1.0 / 2**24)).astype(np.float32)
62
63
64def uniform32_from_uint64(x):
65 x = np.uint64(x)
66 upper = np.array(x >> np.uint64(32), dtype=np.uint32)
67 lower = np.uint64(0xffffffff)
68 lower = np.array(x & lower, dtype=np.uint32)
69 joined = np.column_stack([lower, upper]).ravel()
70 return uint32_to_float32(joined)
71
72
73def uniform32_from_uint53(x):
74 x = np.uint64(x) >> np.uint64(16)
75 x = np.uint32(x & np.uint64(0xffffffff))
76 return uint32_to_float32(x)
77
78
79def uniform32_from_uint32(x):
80 return uint32_to_float32(x)
81
82
83def uniform32_from_uint(x, bits):
84 if bits == 64:
85 return uniform32_from_uint64(x)
86 elif bits == 53:
87 return uniform32_from_uint53(x)
88 elif bits == 32:
89 return uniform32_from_uint32(x)
90 else:
91 raise NotImplementedError
92
93
94def uniform_from_uint(x, bits):
95 if bits in (64, 63, 53):
96 return uniform_from_uint64(x)
97 elif bits == 32:
98 return uniform_from_uint32(x)
99
100
101def uniform_from_uint64(x):
102 return (x >> np.uint64(11)) * (1.0 / 9007199254740992.0)
103
104
105def uniform_from_uint32(x):
106 out = np.empty(len(x) // 2)
107 for i in range(0, len(x), 2):
108 a = x[i] >> 5
109 b = x[i + 1] >> 6
110 out[i // 2] = (a * 67108864.0 + b) / 9007199254740992.0
111 return out
112
113
114def uniform_from_dsfmt(x):
115 return x.view(np.double) - 1.0
116
117
118def gauss_from_uint(x, n, bits):
119 if bits in (64, 63):
120 doubles = uniform_from_uint64(x)
121 elif bits == 32:
122 doubles = uniform_from_uint32(x)
123 else: # bits == 'dsfmt'
124 doubles = uniform_from_dsfmt(x)
125 gauss = []
126 loc = 0
127 x1 = x2 = 0.0
128 while len(gauss) < n:
129 r2 = 2
130 while r2 >= 1.0 or r2 == 0.0:
131 x1 = 2.0 * doubles[loc] - 1.0
132 x2 = 2.0 * doubles[loc + 1] - 1.0
133 r2 = x1 * x1 + x2 * x2
134 loc += 2
135
136 f = np.sqrt(-2.0 * np.log(r2) / r2)
137 gauss.append(f * x2)
138 gauss.append(f * x1)
139
140 return gauss[:n]
141
142
143def test_seedsequence():
144 from numpy.random.bit_generator import (
145 ISeedSequence,
146 ISpawnableSeedSequence,
147 SeedlessSeedSequence,
148 )
149
150 s1 = SeedSequence(range(10), spawn_key=(1, 2), pool_size=6)
151 s1.spawn(10)
152 s2 = SeedSequence(**s1.state)
153 assert_equal(s1.state, s2.state)
154 assert_equal(s1.n_children_spawned, s2.n_children_spawned)
155
156 # The interfaces cannot be instantiated themselves.
157 assert_raises(TypeError, ISeedSequence)
158 assert_raises(TypeError, ISpawnableSeedSequence)
159 dummy = SeedlessSeedSequence()
160 assert_raises(NotImplementedError, dummy.generate_state, 10)
161 assert len(dummy.spawn(10)) == 10
162
163
164def test_generator_spawning():
165 """ Test spawning new generators and bit_generators directly.
166 """
167 rng = np.random.default_rng()
168 seq = rng.bit_generator.seed_seq
169 new_ss = seq.spawn(5)
170 expected_keys = [seq.spawn_key + (i,) for i in range(5)]
171 assert [c.spawn_key for c in new_ss] == expected_keys
172
173 new_bgs = rng.bit_generator.spawn(5)
174 expected_keys = [seq.spawn_key + (i,) for i in range(5, 10)]
175 assert [bg.seed_seq.spawn_key for bg in new_bgs] == expected_keys
176
177 new_rngs = rng.spawn(5)
178 expected_keys = [seq.spawn_key + (i,) for i in range(10, 15)]
179 found_keys = [rng.bit_generator.seed_seq.spawn_key for rng in new_rngs]
180 assert found_keys == expected_keys
181
182 # Sanity check that streams are actually different:
183 assert new_rngs[0].uniform() != new_rngs[1].uniform()
184
185
186def test_non_spawnable():
187 from numpy.random.bit_generator import ISeedSequence
188
189 class FakeSeedSequence:
190 def generate_state(self, n_words, dtype=np.uint32):
191 return np.zeros(n_words, dtype=dtype)
192
193 ISeedSequence.register(FakeSeedSequence)
194
195 rng = np.random.default_rng(FakeSeedSequence())
196
197 with pytest.raises(TypeError, match="The underlying SeedSequence"):
198 rng.spawn(5)
199
200 with pytest.raises(TypeError, match="The underlying SeedSequence"):
201 rng.bit_generator.spawn(5)
202
203
204class Base:
205 dtype = np.uint64
206 data2 = data1 = {}
207
208 @classmethod
209 def setup_class(cls):
210 cls.bit_generator = PCG64
211 cls.bits = 64
212 cls.dtype = np.uint64
213 cls.seed_error_type = TypeError
214 cls.invalid_init_types = []
215 cls.invalid_init_values = []
216
217 @classmethod
218 def _read_csv(cls, filename):
219 with open(filename) as csv:
220 seed = csv.readline()
221 seed = seed.split(',')
222 seed = [int(s.strip(), 0) for s in seed[1:]]
223 data = []
224 for line in csv:
225 data.append(int(line.split(',')[-1].strip(), 0))
226 return {'seed': seed, 'data': np.array(data, dtype=cls.dtype)}
227
228 def test_raw(self):
229 bit_generator = self.bit_generator(*self.data1['seed'])
230 uints = bit_generator.random_raw(1000)
231 assert_equal(uints, self.data1['data'])
232
233 bit_generator = self.bit_generator(*self.data1['seed'])
234 uints = bit_generator.random_raw()
235 assert_equal(uints, self.data1['data'][0])
236
237 bit_generator = self.bit_generator(*self.data2['seed'])
238 uints = bit_generator.random_raw(1000)
239 assert_equal(uints, self.data2['data'])
240
241 def test_random_raw(self):
242 bit_generator = self.bit_generator(*self.data1['seed'])
243 uints = bit_generator.random_raw(output=False)
244 assert uints is None
245 uints = bit_generator.random_raw(1000, output=False)
246 assert uints is None
247
248 def test_gauss_inv(self):
249 n = 25
250 rs = RandomState(self.bit_generator(*self.data1['seed']))
251 gauss = rs.standard_normal(n)
252 assert_allclose(gauss,
253 gauss_from_uint(self.data1['data'], n, self.bits))
254
255 rs = RandomState(self.bit_generator(*self.data2['seed']))
256 gauss = rs.standard_normal(25)
257 assert_allclose(gauss,
258 gauss_from_uint(self.data2['data'], n, self.bits))
259
260 def test_uniform_double(self):
261 rs = Generator(self.bit_generator(*self.data1['seed']))
262 vals = uniform_from_uint(self.data1['data'], self.bits)
263 uniforms = rs.random(len(vals))
264 assert_allclose(uniforms, vals)
265 assert_equal(uniforms.dtype, np.float64)
266
267 rs = Generator(self.bit_generator(*self.data2['seed']))
268 vals = uniform_from_uint(self.data2['data'], self.bits)
269 uniforms = rs.random(len(vals))
270 assert_allclose(uniforms, vals)
271 assert_equal(uniforms.dtype, np.float64)
272
273 def test_uniform_float(self):
274 rs = Generator(self.bit_generator(*self.data1['seed']))
275 vals = uniform32_from_uint(self.data1['data'], self.bits)
276 uniforms = rs.random(len(vals), dtype=np.float32)
277 assert_allclose(uniforms, vals)
278 assert_equal(uniforms.dtype, np.float32)
279
280 rs = Generator(self.bit_generator(*self.data2['seed']))
281 vals = uniform32_from_uint(self.data2['data'], self.bits)
282 uniforms = rs.random(len(vals), dtype=np.float32)
283 assert_allclose(uniforms, vals)
284 assert_equal(uniforms.dtype, np.float32)
285
286 def test_repr(self):
287 rs = Generator(self.bit_generator(*self.data1['seed']))
288 assert 'Generator' in repr(rs)
289 assert f'{id(rs):#x}'.upper().replace('X', 'x') in repr(rs)
290
291 def test_str(self):
292 rs = Generator(self.bit_generator(*self.data1['seed']))
293 assert 'Generator' in str(rs)
294 assert str(self.bit_generator.__name__) in str(rs)
295 assert f'{id(rs):#x}'.upper().replace('X', 'x') not in str(rs)
296
297 def test_pickle(self):
298 import pickle
299
300 bit_generator = self.bit_generator(*self.data1['seed'])
301 state = bit_generator.state
302 bitgen_pkl = pickle.dumps(bit_generator)
303 reloaded = pickle.loads(bitgen_pkl)
304 reloaded_state = reloaded.state
305 assert_array_equal(Generator(bit_generator).standard_normal(1000),
306 Generator(reloaded).standard_normal(1000))
307 assert bit_generator is not reloaded
308 assert_state_equal(reloaded_state, state)
309
310 ss = SeedSequence(100)
311 aa = pickle.loads(pickle.dumps(ss))
312 assert_equal(ss.state, aa.state)
313
314 def test_pickle_preserves_seed_sequence(self):
315 # GH 26234
316 # Add explicit test that bit generators preserve seed sequences
317 import pickle
318
319 bit_generator = self.bit_generator(*self.data1['seed'])
320 ss = bit_generator.seed_seq
321 bg_plk = pickle.loads(pickle.dumps(bit_generator))
322 ss_plk = bg_plk.seed_seq
323 assert_equal(ss.state, ss_plk.state)
324 assert_equal(ss.pool, ss_plk.pool)
325
326 bit_generator.seed_seq.spawn(10)
327 bg_plk = pickle.loads(pickle.dumps(bit_generator))
328 ss_plk = bg_plk.seed_seq
329 assert_equal(ss.state, ss_plk.state)
330 assert_equal(ss.n_children_spawned, ss_plk.n_children_spawned)
331
332 def test_invalid_state_type(self):
333 bit_generator = self.bit_generator(*self.data1['seed'])
334 with pytest.raises(TypeError):
335 bit_generator.state = {'1'}
336
337 def test_invalid_state_value(self):
338 bit_generator = self.bit_generator(*self.data1['seed'])
339 state = bit_generator.state
340 state['bit_generator'] = 'otherBitGenerator'
341 with pytest.raises(ValueError):
342 bit_generator.state = state
343
344 def test_invalid_init_type(self):
345 bit_generator = self.bit_generator
346 for st in self.invalid_init_types:
347 with pytest.raises(TypeError):
348 bit_generator(*st)
349
350 def test_invalid_init_values(self):
351 bit_generator = self.bit_generator
352 for st in self.invalid_init_values:
353 with pytest.raises((ValueError, OverflowError)):
354 bit_generator(*st)
355
356 def test_benchmark(self):
357 bit_generator = self.bit_generator(*self.data1['seed'])
358 bit_generator._benchmark(1)
359 bit_generator._benchmark(1, 'double')
360 with pytest.raises(ValueError):
361 bit_generator._benchmark(1, 'int32')
362
363 @pytest.mark.skipif(MISSING_CFFI, reason='cffi not available')
364 def test_cffi(self):
365 bit_generator = self.bit_generator(*self.data1['seed'])
366 cffi_interface = bit_generator.cffi
367 assert isinstance(cffi_interface, interface)
368 other_cffi_interface = bit_generator.cffi
369 assert other_cffi_interface is cffi_interface
370
371 @pytest.mark.skipif(MISSING_CTYPES, reason='ctypes not available')
372 def test_ctypes(self):
373 bit_generator = self.bit_generator(*self.data1['seed'])
374 ctypes_interface = bit_generator.ctypes
375 assert isinstance(ctypes_interface, interface)
376 other_ctypes_interface = bit_generator.ctypes
377 assert other_ctypes_interface is ctypes_interface
378
379 def test_getstate(self):
380 bit_generator = self.bit_generator(*self.data1['seed'])
381 state = bit_generator.state
382 alt_state = bit_generator.__getstate__()
383 assert isinstance(alt_state, tuple)
384 assert_state_equal(state, alt_state[0])
385 assert isinstance(alt_state[1], SeedSequence)
386
387class TestPhilox(Base):
388 @classmethod
389 def setup_class(cls):
390 cls.bit_generator = Philox
391 cls.bits = 64
392 cls.dtype = np.uint64
393 cls.data1 = cls._read_csv(
394 join(pwd, './data/philox-testset-1.csv'))
395 cls.data2 = cls._read_csv(
396 join(pwd, './data/philox-testset-2.csv'))
397 cls.seed_error_type = TypeError
398 cls.invalid_init_types = []
399 cls.invalid_init_values = [(1, None, 1), (-1,), (None, None, 2 ** 257 + 1)]
400
401 def test_set_key(self):
402 bit_generator = self.bit_generator(*self.data1['seed'])
403 state = bit_generator.state
404 keyed = self.bit_generator(counter=state['state']['counter'],
405 key=state['state']['key'])
406 assert_state_equal(bit_generator.state, keyed.state)
407
408
409class TestPCG64(Base):
410 @classmethod
411 def setup_class(cls):
412 cls.bit_generator = PCG64
413 cls.bits = 64
414 cls.dtype = np.uint64
415 cls.data1 = cls._read_csv(join(pwd, './data/pcg64-testset-1.csv'))
416 cls.data2 = cls._read_csv(join(pwd, './data/pcg64-testset-2.csv'))
417 cls.seed_error_type = (ValueError, TypeError)
418 cls.invalid_init_types = [(3.2,), ([None],), (1, None)]
419 cls.invalid_init_values = [(-1,)]
420
421 def test_advance_symmetry(self):
422 rs = Generator(self.bit_generator(*self.data1['seed']))
423 state = rs.bit_generator.state
424 step = -0x9e3779b97f4a7c150000000000000000
425 rs.bit_generator.advance(step)
426 val_neg = rs.integers(10)
427 rs.bit_generator.state = state
428 rs.bit_generator.advance(2**128 + step)
429 val_pos = rs.integers(10)
430 rs.bit_generator.state = state
431 rs.bit_generator.advance(10 * 2**128 + step)
432 val_big = rs.integers(10)
433 assert val_neg == val_pos
434 assert val_big == val_pos
435
436 def test_advange_large(self):
437 rs = Generator(self.bit_generator(38219308213743))
438 pcg = rs.bit_generator
439 state = pcg.state["state"]
440 initial_state = 287608843259529770491897792873167516365
441 assert state["state"] == initial_state
442 pcg.advance(sum(2**i for i in (96, 64, 32, 16, 8, 4, 2, 1)))
443 state = pcg.state["state"]
444 advanced_state = 135275564607035429730177404003164635391
445 assert state["state"] == advanced_state
446
447
448class TestPCG64DXSM(Base):
449 @classmethod
450 def setup_class(cls):
451 cls.bit_generator = PCG64DXSM
452 cls.bits = 64
453 cls.dtype = np.uint64
454 cls.data1 = cls._read_csv(join(pwd, './data/pcg64dxsm-testset-1.csv'))
455 cls.data2 = cls._read_csv(join(pwd, './data/pcg64dxsm-testset-2.csv'))
456 cls.seed_error_type = (ValueError, TypeError)
457 cls.invalid_init_types = [(3.2,), ([None],), (1, None)]
458 cls.invalid_init_values = [(-1,)]
459
460 def test_advance_symmetry(self):
461 rs = Generator(self.bit_generator(*self.data1['seed']))
462 state = rs.bit_generator.state
463 step = -0x9e3779b97f4a7c150000000000000000
464 rs.bit_generator.advance(step)
465 val_neg = rs.integers(10)
466 rs.bit_generator.state = state
467 rs.bit_generator.advance(2**128 + step)
468 val_pos = rs.integers(10)
469 rs.bit_generator.state = state
470 rs.bit_generator.advance(10 * 2**128 + step)
471 val_big = rs.integers(10)
472 assert val_neg == val_pos
473 assert val_big == val_pos
474
475 def test_advange_large(self):
476 rs = Generator(self.bit_generator(38219308213743))
477 pcg = rs.bit_generator
478 state = pcg.state
479 initial_state = 287608843259529770491897792873167516365
480 assert state["state"]["state"] == initial_state
481 pcg.advance(sum(2**i for i in (96, 64, 32, 16, 8, 4, 2, 1)))
482 state = pcg.state["state"]
483 advanced_state = 277778083536782149546677086420637664879
484 assert state["state"] == advanced_state
485
486
487class TestMT19937(Base):
488 @classmethod
489 def setup_class(cls):
490 cls.bit_generator = MT19937
491 cls.bits = 32
492 cls.dtype = np.uint32
493 cls.data1 = cls._read_csv(join(pwd, './data/mt19937-testset-1.csv'))
494 cls.data2 = cls._read_csv(join(pwd, './data/mt19937-testset-2.csv'))
495 cls.seed_error_type = ValueError
496 cls.invalid_init_types = []
497 cls.invalid_init_values = [(-1,)]
498
499 def test_seed_float_array(self):
500 assert_raises(TypeError, self.bit_generator, np.array([np.pi]))
501 assert_raises(TypeError, self.bit_generator, np.array([-np.pi]))
502 assert_raises(TypeError, self.bit_generator, np.array([np.pi, -np.pi]))
503 assert_raises(TypeError, self.bit_generator, np.array([0, np.pi]))
504 assert_raises(TypeError, self.bit_generator, [np.pi])
505 assert_raises(TypeError, self.bit_generator, [0, np.pi])
506
507 def test_state_tuple(self):
508 rs = Generator(self.bit_generator(*self.data1['seed']))
509 bit_generator = rs.bit_generator
510 state = bit_generator.state
511 desired = rs.integers(2 ** 16)
512 tup = (state['bit_generator'], state['state']['key'],
513 state['state']['pos'])
514 bit_generator.state = tup
515 actual = rs.integers(2 ** 16)
516 assert_equal(actual, desired)
517 tup = tup + (0, 0.0)
518 bit_generator.state = tup
519 actual = rs.integers(2 ** 16)
520 assert_equal(actual, desired)
521
522
523class TestSFC64(Base):
524 @classmethod
525 def setup_class(cls):
526 cls.bit_generator = SFC64
527 cls.bits = 64
528 cls.dtype = np.uint64
529 cls.data1 = cls._read_csv(
530 join(pwd, './data/sfc64-testset-1.csv'))
531 cls.data2 = cls._read_csv(
532 join(pwd, './data/sfc64-testset-2.csv'))
533 cls.seed_error_type = (ValueError, TypeError)
534 cls.invalid_init_types = [(3.2,), ([None],), (1, None)]
535 cls.invalid_init_values = [(-1,)]
536
537 def test_legacy_pickle(self):
538 # Pickling format was changed in 2.0.x
539 import gzip
540 import pickle
541
542 expected_state = np.array(
543 [
544 9957867060933711493,
545 532597980065565856,
546 14769588338631205282,
547 13
548 ],
549 dtype=np.uint64
550 )
551
552 base_path = os.path.split(os.path.abspath(__file__))[0]
553 pkl_file = os.path.join(base_path, "data", "sfc64_np126.pkl.gz")
554 with gzip.open(pkl_file) as gz:
555 sfc = pickle.load(gz)
556
557 assert isinstance(sfc, SFC64)
558 assert_equal(sfc.state["state"]["state"], expected_state)
559
560
561class TestDefaultRNG:
562 def test_seed(self):
563 for args in [(), (None,), (1234,), ([1234, 5678],)]:
564 rg = default_rng(*args)
565 assert isinstance(rg.bit_generator, PCG64)
566
567 def test_passthrough(self):
568 bg = Philox()
569 rg = default_rng(bg)
570 assert rg.bit_generator is bg
571 rg2 = default_rng(rg)
572 assert rg2 is rg
573 assert rg2.bit_generator is bg
574
575 @pytest.mark.thread_unsafe(
576 reason="np.random.set_bit_generator affects global state"
577 )
578 def test_coercion_RandomState_Generator(self):
579 # use default_rng to coerce RandomState to Generator
580 rs = RandomState(1234)
581 rg = default_rng(rs)
582 assert isinstance(rg.bit_generator, MT19937)
583 assert rg.bit_generator is rs._bit_generator
584
585 # RandomState with a non MT19937 bit generator
586 _original = np.random.get_bit_generator()
587 bg = PCG64(12342298)
588 np.random.set_bit_generator(bg)
589 rs = np.random.mtrand._rand
590 rg = default_rng(rs)
591 assert rg.bit_generator is bg
592
593 # vital to get global state back to original, otherwise
594 # other tests start to fail.
595 np.random.set_bit_generator(_original)
596 