MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / test_MultinomialRNG

Function test_MultinomialRNG

imperative/python/test/unit/random/test_rng.py:513–656  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

511 get_device_count("xpu") <= 1, reason="xpu counts need > 1",
512)
513def test_MultinomialRNG():
514 # test with replacement
515 num_groups = 2
516 len_probs = 4
517 num_samples = 10000
518 replacement = True
519 m1 = RNG(seed=111, device="xpu0")
520 m2 = RNG(seed=111, device="xpu1")
521 m3 = RNG(seed=222, device="xpu0")
522 input_np = np.array([[1, 2, 3, 4], [0, 7, 2, 1]])
523 probs_np = input_np / input_np.sum(axis=-1, keepdims=True)
524 input = Tensor(input_np, dtype=np.float32)
525 out1 = m1.multinomial(
526 input=input.to("xpu0"), num_samples=num_samples, replacement=replacement
527 )
528 out2 = m2.multinomial(
529 input=input.to("xpu1"), num_samples=num_samples, replacement=replacement
530 )
531 out3 = m3.multinomial(
532 input=input.to("xpu0"), num_samples=num_samples, replacement=replacement
533 )
534 np.testing.assert_allclose(out1.numpy(), out2.numpy(), atol=1e-6)
535 assert out1.device == "xpu0" and out2.device == "xpu1"
536 assert not (out1.numpy() == out3.numpy()).all()
537
538 # test without replacement
539 num_groups = 2
540 len_probs = 4
541 num_samples = 1
542 replacement = False
543 out1 = m1.multinomial(
544 input=input.to("xpu0"), num_samples=num_samples, replacement=replacement
545 )
546 out2 = m2.multinomial(
547 input=input.to("xpu1"), num_samples=num_samples, replacement=replacement
548 )
549 out3 = m3.multinomial(
550 input=input.to("xpu0"), num_samples=num_samples, replacement=replacement
551 )
552 np.testing.assert_allclose(out1.numpy(), out2.numpy(), atol=1e-6)
553 assert out1.device == "xpu0" and out2.device == "xpu1"
554 assert not (out1.numpy() == out3.numpy()).all()
555
556 # test with replacement
557 num_groups = 2
558 len_probs = 4
559 num_samples = 10000
560 replacement = True
561 out = m1.multinomial(
562 input=input.to("xpu0"), num_samples=num_samples, replacement=replacement
563 )
564 out_shp = out.shape
565 expected_shape = (num_groups, num_samples)
566 if isinstance(out_shp, tuple):
567 assert out_shp == expected_shape
568 else:
569 assert all(out.shape.numpy() == np.array(expected_shape))
570 sample_probs = (

Callers

nothing calls this directly

Calls 12

multinomialMethod · 0.95
toMethod · 0.95
RNGClass · 0.90
TensorClass · 0.90
arrayMethod · 0.80
allMethod · 0.80
sumMethod · 0.45
numpyMethod · 0.45
astypeMethod · 0.45
meanMethod · 0.45
varMethod · 0.45
zerosMethod · 0.45

Tested by

no test coverage detected