MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / ctc_state_log_probs

Function ctc_state_log_probs

tensorflow/python/ops/ctc_ops.py:470–507  ·  view source on GitHub ↗

Computes CTC alignment initial and final state log probabilities. Create the initial/final state values directly as log values to avoid having to take a float64 log on tpu (which does not exist). Args: seq_lengths: int tensor of shape [batch_size], seq lengths in the batch. max_seq_l

(seq_lengths, max_seq_length)

Source from the content-addressed store, hash-verified

468
469
470def ctc_state_log_probs(seq_lengths, max_seq_length):
471 """Computes CTC alignment initial and final state log probabilities.
472
473 Create the initial/final state values directly as log values to avoid
474 having to take a float64 log on tpu (which does not exist).
475
476 Args:
477 seq_lengths: int tensor of shape [batch_size], seq lengths in the batch.
478 max_seq_length: int, max sequence length possible.
479
480 Returns:
481 initial_state_log_probs, final_state_log_probs
482 """
483
484 batch_size = _get_dim(seq_lengths, 0)
485 num_label_states = max_seq_length + 1
486 num_duration_states = 2
487 num_states = num_duration_states * num_label_states
488 log_0 = math_ops.cast(
489 math_ops.log(math_ops.cast(0, dtypes.float64) + 1e-307), dtypes.float32)
490
491 initial_state_log_probs = array_ops.one_hot(
492 indices=array_ops.zeros([batch_size], dtype=dtypes.int32),
493 depth=num_states,
494 on_value=0.0,
495 off_value=log_0,
496 axis=1)
497
498 label_final_state_mask = array_ops.one_hot(
499 seq_lengths, depth=num_label_states, axis=0)
500 duration_final_state_mask = array_ops.ones(
501 [num_duration_states, 1, batch_size])
502 final_state_mask = duration_final_state_mask * label_final_state_mask
503 final_state_log_probs = (1.0 - final_state_mask) * log_0
504 final_state_log_probs = array_ops.reshape(final_state_log_probs,
505 [num_states, batch_size])
506
507 return initial_state_log_probs, array_ops.transpose(final_state_log_probs)
508
509
510def _ilabel_to_state(labels, num_labels, ilabel_log_probs):

Callers 1

ctc_loss_and_gradFunction · 0.85

Calls 6

onesMethod · 0.80
reshapeMethod · 0.80
transposeMethod · 0.80
_get_dimFunction · 0.70
castMethod · 0.45
logMethod · 0.45

Tested by

no test coverage detected