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

Function _forward_backward_log

tensorflow/python/ops/ctc_ops.py:1024–1108  ·  view source on GitHub ↗

Forward-backward algorithm computed in log domain. Args: state_trans_log_probs: tensor of shape [states, states] or if different transition matrix per batch [batch_size, states, states] initial_state_log_probs: tensor of shape [batch_size, states] final_state_log_probs: tensor o

(state_trans_log_probs, initial_state_log_probs,
                          final_state_log_probs, observed_log_probs,
                          sequence_length)

Source from the content-addressed store, hash-verified

1022
1023
1024def _forward_backward_log(state_trans_log_probs, initial_state_log_probs,
1025 final_state_log_probs, observed_log_probs,
1026 sequence_length):
1027 """Forward-backward algorithm computed in log domain.
1028
1029 Args:
1030 state_trans_log_probs: tensor of shape [states, states] or if different
1031 transition matrix per batch [batch_size, states, states]
1032 initial_state_log_probs: tensor of shape [batch_size, states]
1033 final_state_log_probs: tensor of shape [batch_size, states]
1034 observed_log_probs: tensor of shape [frames, batch_size, states]
1035 sequence_length: tensor of shape [batch_size]
1036
1037 Returns:
1038 forward backward log probabilites: tensor of shape [frames, batch, states]
1039 log_likelihood: tensor of shape [batch_size]
1040
1041 Raises:
1042 ValueError: If state_trans_log_probs has unknown or incorrect rank.
1043 """
1044
1045 if state_trans_log_probs.shape.ndims == 2:
1046 perm = [1, 0]
1047 elif state_trans_log_probs.shape.ndims == 3:
1048 perm = [0, 2, 1]
1049 else:
1050 raise ValueError(
1051 "state_trans_log_probs rank must be known and == 2 or 3, is: %s" %
1052 state_trans_log_probs.shape.ndims)
1053
1054 bwd_state_trans_log_probs = array_ops.transpose(state_trans_log_probs, perm)
1055 batch_size = _get_dim(observed_log_probs, 1)
1056
1057 def _forward(state_log_prob, obs_log_prob):
1058 state_log_prob = array_ops.expand_dims(state_log_prob, axis=1) # Broadcast.
1059 state_log_prob += state_trans_log_probs
1060 state_log_prob = math_ops.reduce_logsumexp(state_log_prob, axis=-1)
1061 state_log_prob += obs_log_prob
1062 log_prob_sum = math_ops.reduce_logsumexp(
1063 state_log_prob, axis=-1, keepdims=True)
1064 state_log_prob -= log_prob_sum
1065 return state_log_prob
1066
1067 fwd = _scan(
1068 _forward, observed_log_probs, initial_state_log_probs, inclusive=True)
1069
1070 def _backward(accs, elems):
1071 """Calculate log probs and cumulative sum masked for sequence length."""
1072 state_log_prob, cum_log_sum = accs
1073 obs_log_prob, mask = elems
1074 state_log_prob += obs_log_prob
1075 state_log_prob = array_ops.expand_dims(state_log_prob, axis=1) # Broadcast.
1076 state_log_prob += bwd_state_trans_log_probs
1077 state_log_prob = math_ops.reduce_logsumexp(state_log_prob, axis=-1)
1078
1079 log_prob_sum = math_ops.reduce_logsumexp(
1080 state_log_prob, axis=-1, keepdims=True)
1081 state_log_prob -= log_prob_sum

Callers 1

ctc_loss_and_gradFunction · 0.85

Calls 5

_scanFunction · 0.85
transposeMethod · 0.80
_get_dimFunction · 0.70
logMethod · 0.45
expand_dimsMethod · 0.45

Tested by

no test coverage detected