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)
| 1022 | |
| 1023 | |
| 1024 | def _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 |
no test coverage detected