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

Function _reverse_seq

tensorflow/python/ops/rnn.py:318–357  ·  view source on GitHub ↗

Reverse a list of Tensors up to specified lengths. Args: input_seq: Sequence of seq_len tensors of dimension (batch_size, n_features) or nested tuples of tensors. lengths: A `Tensor` of dimension batch_size, containing lengths for each sequence in the batch. If "None" is spe

(input_seq, lengths)

Source from the content-addressed store, hash-verified

316
317
318def _reverse_seq(input_seq, lengths):
319 """Reverse a list of Tensors up to specified lengths.
320
321 Args:
322 input_seq: Sequence of seq_len tensors of dimension (batch_size, n_features)
323 or nested tuples of tensors.
324 lengths: A `Tensor` of dimension batch_size, containing lengths for each
325 sequence in the batch. If "None" is specified, simply reverses the list.
326
327 Returns:
328 time-reversed sequence
329 """
330 if lengths is None:
331 return list(reversed(input_seq))
332
333 flat_input_seq = tuple(nest.flatten(input_) for input_ in input_seq)
334
335 flat_results = [[] for _ in range(len(input_seq))]
336 for sequence in zip(*flat_input_seq):
337 input_shape = tensor_shape.unknown_shape(rank=sequence[0].get_shape().rank)
338 for input_ in sequence:
339 input_shape.merge_with(input_.get_shape())
340 input_.set_shape(input_shape)
341
342 # Join into (time, batch_size, depth)
343 s_joined = array_ops.stack(sequence)
344
345 # Reverse along dimension 0
346 s_reversed = array_ops.reverse_sequence(s_joined, lengths, 0, 1)
347 # Split again into list
348 result = array_ops.unstack(s_reversed)
349 for r, flat_result in zip(result, flat_results):
350 r.set_shape(input_shape)
351 flat_result.append(r)
352
353 results = [
354 nest.pack_sequence_as(structure=input_, flat_sequence=flat_result)
355 for input_, flat_result in zip(input_seq, flat_results)
356 ]
357 return results
358
359
360@deprecation.deprecated(None, "Please use `keras.layers.Bidirectional("

Callers 1

static_bidirectional_rnnFunction · 0.70

Calls 10

tupleFunction · 0.85
unknown_shapeMethod · 0.80
rangeFunction · 0.70
flattenMethod · 0.45
get_shapeMethod · 0.45
merge_withMethod · 0.45
set_shapeMethod · 0.45
stackMethod · 0.45
unstackMethod · 0.45
appendMethod · 0.45

Tested by

no test coverage detected