You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.12 05:07:42