MCPcopy Create free account
hub / github.com/apple/axlearn / compute_code_pplx

Function compute_code_pplx

axlearn/common/quantizer.py:72–83  ·  view source on GitHub ↗

Computes pplx and entropy of the quantized codes distribution.

(onehots: Tensor, paddings: Tensor)

Source from the content-addressed store, hash-verified

70
71
72def 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
86def compute_code_coverage(onehots: Tensor, paddings: Tensor) -> Tensor:

Callers 3

test_quantizeMethod · 0.90
_add_codebook_summariesFunction · 0.85

Calls 2

safe_notFunction · 0.90
compute_code_histogramFunction · 0.85

Tested by 2

test_quantizeMethod · 0.72