| 511 | get_device_count("xpu") <= 1, reason="xpu counts need > 1", |
| 512 | ) |
| 513 | def 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 = ( |