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)
| 468 | |
| 469 | |
| 470 | def 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 | |
| 510 | def _ilabel_to_state(labels, num_labels, ilabel_log_probs): |