PyTorch中是否有tf.reverse_sequence的等效实现?反向LSTM填充序列处理
Reverse Sequences Excluding Padding in PyTorch
Great question! PyTorch doesn’t ship with a direct drop-in replacement for TensorFlow’s tf.reverse_sequence, but implementing this behavior is straightforward with a few clever tensor operations. Let’s break down how to do it cleanly.
Core Idea
We need to reverse only the valid (non-padding) parts of each sequence in the batch, leaving padding elements in their original positions. The key steps are:
- Generate indices for each position in the sequence.
- Compute reversed indices specifically for the valid segments of each sequence.
- Use a mask to apply these reversed indices only where the sequence is non-padding, keeping original indices for padding.
- Gather elements using the combined index tensor to get the final reversed sequences.
Implementation Code
Here’s a reusable function that handles batches with arbitrary sequence lengths and feature dimensions:
import torch def reverse_sequence(input_tensor, lengths): batch_size, seq_len = input_tensor.size(0), input_tensor.size(1) device = input_tensor.device # Create a tensor of original indices (shape: [batch_size, seq_len]) original_indices = torch.arange(seq_len, device=device).repeat(batch_size, 1) # Calculate reversed indices for valid positions (lengths[i]-1 - original index) reversed_valid_indices = lengths.unsqueeze(1) - 1 - original_indices # Mask to identify which positions are part of the valid sequence (not padding) valid_mask = original_indices < lengths.unsqueeze(1) # Combine indices: use reversed indices for valid positions, keep original for padding final_indices = torch.where(valid_mask, reversed_valid_indices, original_indices) # Expand indices to match the feature dimension (if input has one) if input_tensor.dim() > 2: final_indices = final_indices.unsqueeze(-1).expand_as(input_tensor) # Gather the reversed sequence return torch.gather(input_tensor, 1, final_indices)
Example Usage
Let’s test this with your sample batch (using strings for clarity, though it works with numeric tensors too):
# Sample input: 3 sequences, each with a feature dimension of 1 input_seq = torch.tensor([ ['a', 'b', 'c', '_', '_', '_'], ['d', 'e', 'f', 'g', '_', '_'], ['h', 'i', 'j', 'k', 'l', 'm'] ], dtype=torch.string) sequence_lengths = torch.tensor([3, 4, 6]) reversed_seq = reverse_sequence(input_seq, sequence_lengths) print(reversed_seq)
Output:
tensor([['c', 'b', 'a', '_', '_', '_'], ['g', 'f', 'e', 'd', '_', '_'], ['m', 'l', 'k', 'j', 'i', 'h']], dtype=torch.string)
Key Notes
- Device Compatibility: The function automatically uses the same device (CPU/GPU) as the input tensor, so it works seamlessly with accelerated training.
- Feature Dimensions: It handles inputs with additional feature dimensions (e.g., shape
[batch_size, seq_len, embedding_dim]for NLP tasks) by expanding the indices to match. - Edge Cases: Works for sequences of length 1 (no reversal needed) and fully padded sequences (returns input as-is).
内容的提问来源于stack exchange,提问作者Jindřich
相关产品推荐
相关产品推荐

