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

Function body

tensorflow/python/ops/rnn.py:1186–1243  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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)

Callers 15

_BuildLoopMethod · 0.70
while_loopFunction · 0.70
wrapped_bodyFunction · 0.70
body_wrapperFunction · 0.50
_py_for_stmtFunction · 0.50
while_bodyFunction · 0.50
while_body_actualFunction · 0.50
true_fnFunction · 0.50
reduce_bodyFunction · 0.50
while_stmtFunction · 0.50
aug_bodyFunction · 0.50
_py_while_stmtFunction · 0.50

Calls 3

_copy_some_throughFunction · 0.70
loop_fnFunction · 0.50
writeMethod · 0.45

Tested by

no test coverage detected