(self)
| 138 | ) |
| 139 | |
| 140 | def test_compute_code_pplx(self): # pylint: disable=no-self-use |
| 141 | vocab_size = 11 |
| 142 | |
| 143 | codes = jnp.array( |
| 144 | [ |
| 145 | [[1, 2], [3, 4], [5, 6], [7, 8], [9, 10], [0, 0]], |
| 146 | [[1, 2], [3, 4], [5, 6], [7, 8], [9, 10], [0, 0]], |
| 147 | ] |
| 148 | ) |
| 149 | paddings = jnp.array([[0, 0, 0, 0, 0, 1], [0, 0, 0, 0, 0, 1]], jnp.bool) |
| 150 | |
| 151 | pplx, entropy = compute_code_pplx( |
| 152 | onehots=jax.nn.one_hot(codes, num_classes=vocab_size, axis=-1), paddings=paddings |
| 153 | ) |
| 154 | |
| 155 | assert_allclose(pplx, 5.000, rtol=1e-06, atol=1e-06) |
| 156 | assert_allclose(entropy, 1.609438, rtol=1e-06, atol=1e-06) |
| 157 | |
| 158 | def test_codebook_initializer(self): |
| 159 | vocab_size, dim_from_all_codebooks, num_groups = 10, 100, 5 |
nothing calls this directly
no test coverage detected