Computes pplx and entropy of the quantized codes distribution.
(onehots: Tensor, paddings: Tensor)
| 70 | |
| 71 | |
| 72 | def compute_code_pplx(onehots: Tensor, paddings: Tensor) -> tuple[Tensor, Tensor]: |
| 73 | """Computes pplx and entropy of the quantized codes distribution.""" |
| 74 | histogram = compute_code_histogram(onehots, paddings) |
| 75 | normalizer = jnp.sum(safe_not(paddings)) |
| 76 | # [num_codebooks, codebook_size]. |
| 77 | probs = histogram / jnp.maximum(normalizer, 1.0) |
| 78 | log_probs = jnp.log(jnp.maximum(1.0e-30, probs)) |
| 79 | # [num_codebooks]. |
| 80 | sum_plogp = jnp.sum(log_probs * probs, axis=-1) |
| 81 | pplx = jnp.mean(jnp.exp(-sum_plogp)) |
| 82 | entropy = jnp.log(pplx) |
| 83 | return pplx, entropy |
| 84 | |
| 85 | |
| 86 | def compute_code_coverage(onehots: Tensor, paddings: Tensor) -> Tensor: |