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

Method _input_stats_summaries

axlearn/audio/decoder_asr.py:239–271  ·  view source on GitHub ↗

Computes input lengths stats. Args: input_batch: See forward method signature. target_paddings: See _compute_target_paddings method return value. is_valid_example: A 0/1 Tensor of shape [batch_size], 1 if the example is a valid input to th

(
        self, input_batch: Nested[Tensor], *, target_paddings: Tensor, is_valid_example: Tensor
    )

Source from the content-addressed store, hash-verified

237 raise NotImplementedError(type(self))
238
239 def _input_stats_summaries(
240 self, input_batch: Nested[Tensor], *, target_paddings: Tensor, is_valid_example: Tensor
241 ) -> dict[str, Union[WeightedSummary, Tensor]]:
242 """Computes input lengths stats.
243
244 Args:
245 input_batch: See forward method signature.
246 target_paddings: See _compute_target_paddings method return value.
247 is_valid_example: A 0/1 Tensor of shape [batch_size], 1 if the example is
248 a valid input to the loss computation.
249
250 Returns:
251 A dictionary of input stats summaries.
252 """
253 valid_frames = (1.0 - input_batch["paddings"]) * is_valid_example[:, None]
254 valid_labels = (1.0 - target_paddings) * is_valid_example[:, None]
255
256 total_source_lengths = jnp.sum(valid_frames)
257 total_target_lengths = jnp.sum(valid_labels)
258 total_num_examples = jnp.maximum(is_valid_example.sum(), 1.0)
259 total_num_frames = jnp.maximum(jnp.size(input_batch["paddings"]), 1)
260 input_stats = {
261 "input_stats/average_target_length": WeightedSummary(
262 total_target_lengths / total_num_examples, total_num_examples
263 ),
264 "input_stats/average_source_length": WeightedSummary(
265 total_source_lengths / total_num_examples, total_num_examples
266 ),
267 "input_stats/frame_packing_efficiency": WeightedSummary(
268 total_source_lengths / total_num_frames, total_num_frames
269 ),
270 }
271 return input_stats
272
273
274class CTCDecoderModel(BaseASRDecoderModel):

Callers 3

forwardMethod · 0.80
forwardMethod · 0.80
forwardMethod · 0.80

Calls 1

WeightedSummaryClass · 0.90

Tested by

no test coverage detected