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

Function test_beta_op

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

Source from the content-addressed store, hash-verified

123 get_device_count("xpu") <= 2, reason="xpu counts need > 2",
124)
125def 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(

Callers

nothing calls this directly

Calls 10

get_global_rng_seedFunction · 0.90
BetaRNGClass · 0.85
strFunction · 0.85
CompNodeClass · 0.85
applyFunction · 0.50
sqrtMethod · 0.45
fabsMethod · 0.45
meanMethod · 0.45
numpyMethod · 0.45
varMethod · 0.45

Tested by

no test coverage detected