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

Function _add_codebook_summaries

axlearn/common/quantizer.py:276–310  ·  view source on GitHub ↗

Helper function to compute codebook distribution statistics and add to summaries. The statistics are from all frames, not only on those masked frames in self-supervised training. # ToDo(zhiyunlu): Add support to take an additional mask input. Args: context: Module invocation co

(*, context: InvocationContext, onehots: Tensor, paddings: Tensor)

Source from the content-addressed store, hash-verified

274
275
276def _add_codebook_summaries(*, context: InvocationContext, onehots: Tensor, paddings: Tensor):
277 """Helper function to compute codebook distribution statistics and add to summaries.
278
279 The statistics are from all frames, not only on those masked frames in self-supervised training.
280 # ToDo(zhiyunlu): Add support to take an additional mask input.
281
282 Args:
283 context: Module invocation context to add summaries to.
284 onehots: onehot of BaseQuantizer.Output.ids.
285 paddings: 0/1 tensor of shape [batch_size, seq_len], where 0 is valid position.
286 """
287 coverage = compute_code_coverage(onehots=onehots, paddings=paddings)
288 pplx, entropy = compute_code_pplx(onehots=onehots, paddings=paddings)
289 batch_size = paddings.shape[0]
290
291 num_frames = jnp.sum(safe_not(paddings))
292 context.add_summary(
293 "codebook/num_frames",
294 WeightedSummary(num_frames.astype(jnp.float32) / batch_size, batch_size),
295 )
296 # Mean coverage of all codebooks.
297 context.add_summary(
298 "codebook/coverage",
299 WeightedSummary(coverage, jnp.maximum(1, num_frames)),
300 )
301 # Mean perplexity of all codebooks.
302 context.add_summary(
303 "codebook/pplx",
304 WeightedSummary(pplx, jnp.maximum(1, num_frames)),
305 )
306 # Mean entropy of all codebooks.
307 context.add_summary(
308 "codebook/entropy",
309 WeightedSummary(entropy, jnp.maximum(1, num_frames)),
310 )
311
312
313class RandomVectorQuantizer(BaseQuantizer):

Callers 3

forwardMethod · 0.85
forwardMethod · 0.85
forwardMethod · 0.85

Calls 6

safe_notFunction · 0.90
WeightedSummaryClass · 0.90
compute_code_coverageFunction · 0.85
compute_code_pplxFunction · 0.85
astypeMethod · 0.80
add_summaryMethod · 0.45

Tested by

no test coverage detected