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

PyTorch:多同标签变长输入如何使用pack_padded_sequence及排序

Can I use pack_padded_sequence for three variable-length inputs sharing the same labels?

Absolutely! You can definitely use pack_padded_sequence for this scenario—here's a practical, step-by-step guide to handling the sorting and integration, using PyTorch as an example (since this utility is part of PyTorch's RNN tools):

Core Idea

Since all three inputs (a, b, c) map to the same set of labels, you just need to sort all three input sequences in perfect sync based on their lengths (descending order, which is required for pack_padded_sequence). The key is keeping the sample correspondence intact so predictions align with the original labels later.


Step-by-Step Implementation

1. Calculate Sequence Lengths

First, compute the length of each sequence in your inputs. Since each sample's three inputs belong to the same target, you can pick one input's lengths as your sorting key (e.g., lengths of a)—or use the maximum length per sample if you prefer. Either way, stick to a single key for consistency.

# Assume a, b, c are lists of variable-length tensors (batch-first format)
lengths_a = torch.tensor([len(seq) for seq in a], dtype=torch.int64)

# Optional: Use max length of the three inputs per sample
# lengths = torch.tensor([max(len(a_s), len(b_s), len(c_s)) for a_s, b_s, c_s in zip(a, b, c)], dtype=torch.int64)

2. Sort All Inputs in Sync

pack_padded_sequence requires sequences to be sorted in descending order of length. We'll sort all three inputs using the same indices, and save the original indices so we can restore the order later (this is non-negotiable for matching predictions to labels).

# Get indices that sort lengths_a in descending order
sorted_indices = torch.argsort(lengths_a, descending=True)

# Apply this sort to all three inputs
sorted_a = torch.stack(a)[sorted_indices]
sorted_b = torch.stack(b)[sorted_indices]
sorted_c = torch.stack(c)[sorted_indices]

# Sort the lengths too (needed for packing)
sorted_lengths = lengths_a[sorted_indices]

3. Pack Sequences and Run Through LSTMs

Now you can pack each sorted input and feed them to your independent LSTMs. Each LSTM will process its own variable-length sequence efficiently thanks to the packed format.

from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence

# Pack each input (set batch_first=True if your LSTMs use batch-first tensors)
packed_a = pack_padded_sequence(sorted_a, sorted_lengths, batch_first=True, enforce_sorted=True)
packed_b = pack_padded_sequence(sorted_b, sorted_lengths, batch_first=True, enforce_sorted=True)
packed_c = pack_padded_sequence(sorted_c, sorted_lengths, batch_first=True, enforce_sorted=True)

# Initialize your independent LSTMs
lstm_a = nn.LSTM(input_size=your_input_size_a, hidden_size=your_hidden_size, batch_first=True)
lstm_b = nn.LSTM(input_size=your_input_size_b, hidden_size=your_hidden_size, batch_first=True)
lstm_c = nn.LSTM(input_size=your_input_size_c, hidden_size=your_hidden_size, batch_first=True)

# Forward pass through each LSTM
output_a, (h_n_a, _) = lstm_a(packed_a)
output_b, (h_n_b, _) = lstm_b(packed_b)
output_c, (h_n_c, _) = lstm_c(packed_c)

# Optional: Unpack outputs if you need padded sequences for further processing
unpacked_a, _ = pad_packed_sequence(output_a, batch_first=True)

4. Fuse Results and Restore Original Order

After getting outputs from each LSTM, fuse them (e.g., concatenate final hidden states, add them, or use attention-based fusion). Then use the inverse of the sorted indices to put everything back in the original order—this ensures your fused outputs match the original label order.

# Example fusion: Concatenate the final hidden states of each LSTM
fused_output = torch.cat([h_n_a.squeeze(0), h_n_b.squeeze(0), h_n_c.squeeze(0)], dim=1)

# Get inverse indices to restore the original batch order
original_indices = torch.argsort(sorted_indices)
restored_fused = fused_output[original_indices]

# Now you can use restored_fused for prediction, which aligns perfectly with your original labels

Critical Tips

  • Never sort inputs independently: If you sort a, b, and c separately, you'll break the link between each sample's inputs and its label—always use the same sorted indices for all three.
  • enforce_sorted=True: This default setting requires sequences to be pre-sorted. It's more efficient to sort upfront than letting PyTorch do it for you (which happens if you set this to False).
  • Length flexibility: Even if a, b, and c have different lengths for the same sample, using a single length key works because pack_padded_sequence only cares about the length of the specific sequence being packed. Each LSTM will process its own input's actual length correctly.

内容的提问来源于stack exchange,提问作者John M.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 08:22:31