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)
| 316 | |
| 317 | |
| 318 | def _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(" |
no test coverage detected