| 123 | get_device_count("xpu") <= 2, reason="xpu counts need > 2", |
| 124 | ) |
| 125 | def test_beta_op(): |
| 126 | set_global_seed(1024) |
| 127 | _alpha, _beta = 2, 0.8 |
| 128 | _expected_mean = _alpha / (_alpha + _beta) |
| 129 | _expected_std = np.sqrt( |
| 130 | _alpha * _beta / ((_alpha + _beta) ** 2 * (_alpha + _beta + 1)) |
| 131 | ) |
| 132 | |
| 133 | alpha = F.full([8, 9, 11, 12], value=_alpha, dtype="float32") |
| 134 | beta = F.full([8, 9, 11, 12], value=_beta, dtype="float32") |
| 135 | op = BetaRNG(seed=get_global_rng_seed()) |
| 136 | (output,) = apply(op, alpha, beta) |
| 137 | assert np.fabs(output.numpy().mean() - _expected_mean) < 1e-1 |
| 138 | assert np.fabs(np.sqrt(output.numpy().var()) - _expected_std) < 1e-1 |
| 139 | assert str(output.device) == str(CompNode("xpux")) |
| 140 | |
| 141 | cn = CompNode("xpu2") |
| 142 | seed = 233333 |
| 143 | h = new_rng_handle(cn, seed) |
| 144 | alpha = F.full([8, 9, 11, 12], value=_alpha, dtype="float32", device=cn) |
| 145 | beta = F.full([8, 9, 11, 12], value=_beta, dtype="float32", device=cn) |
| 146 | op = BetaRNG(seed=seed, handle=h) |
| 147 | (output,) = apply(op, alpha, beta) |
| 148 | delete_rng_handle(h) |
| 149 | assert np.fabs(output.numpy().mean() - _expected_mean) < 1e-1 |
| 150 | assert np.fabs(np.sqrt(output.numpy().var()) - _expected_std) < 1e-1 |
| 151 | assert str(output.device) == str(cn) |
| 152 | |
| 153 | |
| 154 | @pytest.mark.skipif( |