Internal while loop body for raw_rnn. Args: time: time scalar. elements_finished: batch-size vector. current_input: possibly nested tuple of input tensors. emit_ta: possibly nested tuple of output TensorArrays. state: possibly nested tuple of state tens
(time, elements_finished, current_input, emit_ta, state,
loop_state)
| 1184 | return math_ops.logical_not(math_ops.reduce_all(elements_finished)) |
| 1185 | |
| 1186 | def body(time, elements_finished, current_input, emit_ta, state, |
| 1187 | loop_state): |
| 1188 | """Internal while loop body for raw_rnn. |
| 1189 | |
| 1190 | Args: |
| 1191 | time: time scalar. |
| 1192 | elements_finished: batch-size vector. |
| 1193 | current_input: possibly nested tuple of input tensors. |
| 1194 | emit_ta: possibly nested tuple of output TensorArrays. |
| 1195 | state: possibly nested tuple of state tensors. |
| 1196 | loop_state: possibly nested tuple of loop state tensors. |
| 1197 | |
| 1198 | Returns: |
| 1199 | Tuple having the same size as Args but with updated values. |
| 1200 | """ |
| 1201 | (next_output, cell_state) = cell(current_input, state) |
| 1202 | |
| 1203 | nest.assert_same_structure(state, cell_state) |
| 1204 | nest.assert_same_structure(cell.output_size, next_output) |
| 1205 | |
| 1206 | next_time = time + 1 |
| 1207 | (next_finished, next_input, next_state, emit_output, |
| 1208 | next_loop_state) = loop_fn(next_time, next_output, cell_state, |
| 1209 | loop_state) |
| 1210 | |
| 1211 | nest.assert_same_structure(state, next_state) |
| 1212 | nest.assert_same_structure(current_input, next_input) |
| 1213 | nest.assert_same_structure(emit_ta, emit_output) |
| 1214 | |
| 1215 | # If loop_fn returns None for next_loop_state, just reuse the |
| 1216 | # previous one. |
| 1217 | loop_state = loop_state if next_loop_state is None else next_loop_state |
| 1218 | |
| 1219 | def _copy_some_through(current, candidate): |
| 1220 | """Copy some tensors through via array_ops.where.""" |
| 1221 | |
| 1222 | def copy_fn(cur_i, cand_i): |
| 1223 | # TensorArray and scalar get passed through. |
| 1224 | if isinstance(cur_i, tensor_array_ops.TensorArray): |
| 1225 | return cand_i |
| 1226 | if cur_i.shape.rank == 0: |
| 1227 | return cand_i |
| 1228 | # Otherwise propagate the old or the new value. |
| 1229 | with ops.colocate_with(cand_i): |
| 1230 | return array_ops.where(elements_finished, cur_i, cand_i) |
| 1231 | |
| 1232 | return nest.map_structure(copy_fn, current, candidate) |
| 1233 | |
| 1234 | emit_output = _copy_some_through(zero_emit, emit_output) |
| 1235 | next_state = _copy_some_through(state, next_state) |
| 1236 | |
| 1237 | emit_ta = nest.map_structure(lambda ta, emit: ta.write(time, emit), |
| 1238 | emit_ta, emit_output) |
| 1239 | |
| 1240 | elements_finished = math_ops.logical_or(elements_finished, next_finished) |
| 1241 | |
| 1242 | return (next_time, elements_finished, next_input, emit_ta, next_state, |
| 1243 | loop_state) |
no test coverage detected