PyTorch:多同标签变长输入如何使用pack_padded_sequence及排序
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, andcseparately, 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 toFalse).- Length flexibility: Even if
a,b, andchave different lengths for the same sample, using a single length key works becausepack_padded_sequenceonly 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.

