PyTorch中Embedding、LSTM、Linear层输入及Seq2Seq批量训练输入准备
Let's walk through exactly how to structure inputs for each component of your Seq2Seq encoder-decoder model in PyTorch—especially for batch training. I'll break down each layer's input/output requirements, cover critical batch-specific details like padding and sequence ordering, and wrap it up with a working example.
First, let's set the baseline: for batch training, you'll have two core tensor inputs:
- Source sequences: Shape
(batch_size, src_seq_len)— integer tensors where each element is a word index from your source vocabulary. - Target sequences: Shape
(batch_size, tgt_seq_len)— integer tensors for your target vocabulary, usually prefixed with a<sos>(start-of-sequence) token for teacher-forcing training.
1. nn.Embedding Layer
This layer turns word indices into dense vector representations.
- Input shape:
(batch_size, seq_len)(works for both source and target sequences) - Output shape:
(batch_size, seq_len, embed_dim)— each word index is replaced by itsembed_dim-dimensional vector. - Batch-specific notes:
- Sequences in a batch rarely have the same length, so use
torch.nn.utils.rnn.pad_sequence(withbatch_first=True) to pad all sequences to the length of the longest one in the batch. - Always track the actual length of each sequence (before padding) — this is crucial for optimizing LSTM performance later.
- Sequences in a batch rarely have the same length, so use
2. nn.LSTM Layer
PyTorch's LSTM defaults to a sequence-first input format ((seq_len, batch_size, input_size)), which is different from the embedding layer's batch-first output. We'll split this into encoder and decoder LSTMs since their inputs differ slightly.
Encoder LSTM
The encoder processes the source sequence to produce context states for the decoder.
- Input:
- Take the embedding output (
(batch_size, src_seq_len, embed_dim)), transpose it to(src_seq_len, batch_size, embed_dim)using.permute(1, 0, 2). - For padded sequences, wrap the transposed embedding with
pack_padded_sequence(using the actual sequence lengths you tracked). This tells the LSTM to ignore padded tokens, improving efficiency and accuracy.
- Take the embedding output (
- Outputs:
h_n: Hidden state of the last LSTM layer — shape(num_layers * num_directions, batch_size, hidden_size). For a standard unidirectional encoder, this simplifies to(num_layers, batch_size, hidden_size).c_n: Cell state of the last LSTM layer — same shape ash_n.output(optional): The full sequence of LSTM outputs. If you usedpack_padded_sequence, this will be aPackedSequence; usepad_packed_sequenceto convert it back to(src_seq_len, batch_size, hidden_size).
Decoder LSTM
The decoder generates the target sequence using the encoder's context states.
- Input (Training Mode - Teacher Forcing):
- Use the padded target sequence's embedding output, transposed to
(tgt_seq_len, batch_size, embed_dim). Teacher forcing feeds the actual target tokens (instead of predicted ones) to the decoder, speeding up training. - Initialize the decoder's hidden and cell states with the encoder's final
h_nandc_n(make sure the number of layers matches between encoder and decoder!).
- Use the padded target sequence's embedding output, transposed to
- Input (Inference Mode):
- Feed one token at a time: input shape is
(1, batch_size, embed_dim)(the<sos>token first, then each predicted token in sequence).
- Feed one token at a time: input shape is
- Outputs:
h_n/c_n: Updated hidden/cell states after processing each token — same shape as encoder outputs.output: Sequence of decoder outputs — shape(tgt_seq_len, batch_size, hidden_size)(for training) or(1, batch_size, hidden_size)(for inference).
3. nn.Linear Layer
This layer maps the decoder's LSTM outputs to probabilities over the target vocabulary.
- Input:
- Take the decoder's output sequence (
(tgt_seq_len, batch_size, hidden_size)), then reshape it to(batch_size * tgt_seq_len, hidden_size)(flatten the sequence and batch dimensions into a single axis). This matches the Linear layer's requirement for 2D inputs.
- Take the decoder's output sequence (
- Output:
- Shape
(batch_size * tgt_seq_len, tgt_vocab_size)— each position has logits for every word in the target vocabulary. - Reshape back to
(batch_size, tgt_seq_len, tgt_vocab_size)to easily compute cross-entropy loss (PyTorch'sCrossEntropyLossworks with(N, C)inputs, but reshaping to batch-first sequence format makes it easier to align with target labels).
- Shape
- Tokenize & Index: Convert raw text sequences into integer tensors of word indices.
- Pad Sequences: Use
pad_sequenceto make all sequences in the batch the same length (batch-first shape). - Track Lengths: Save the original length of each sequence (before padding) for the encoder's packed sequences.
- Optional (But Recommended): Sort sequences by length in descending order — this optimizes
pack_padded_sequenceperformance. - Transpose for LSTM: Convert embedding outputs from batch-first to sequence-first format for LSTM input.
import torch import torch.nn as nn from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence class Encoder(nn.Module): def __init__(self, src_vocab_size, embed_dim, hidden_size, num_layers=1): super().__init__() self.embedding = nn.Embedding(src_vocab_size, embed_dim) self.lstm = nn.LSTM(embed_dim, hidden_size, num_layers, batch_first=False) def forward(self, src_input, src_lengths): # src_input: (batch_size, src_seq_len) embed_out = self.embedding(src_input) # (batch_size, src_seq_len, embed_dim) # Convert to sequence-first format for LSTM embed_out = embed_out.permute(1, 0, 2) # (src_seq_len, batch_size, embed_dim) # Pack padded sequences to ignore padding tokens packed_embed = pack_padded_sequence(embed_out, src_lengths, enforce_sorted=False) _, (hidden, cell) = self.lstm(packed_embed) return hidden, cell class Decoder(nn.Module): def __init__(self, tgt_vocab_size, embed_dim, hidden_size, num_layers=1): super().__init__() self.embedding = nn.Embedding(tgt_vocab_size, embed_dim) self.lstm = nn.LSTM(embed_dim, hidden_size, num_layers, batch_first=False) self.fc = nn.Linear(hidden_size, tgt_vocab_size) def forward(self, tgt_input, hidden, cell): # tgt_input: (batch_size, tgt_seq_len) embed_out = self.embedding(tgt_input) # (batch_size, tgt_seq_len, embed_dim) embed_out = embed_out.permute(1, 0, 2) # (tgt_seq_len, batch_size, embed_dim) output, (hidden, cell) = self.lstm(embed_out, (hidden, cell)) # Flatten for Linear layer output_flat = output.reshape(-1, output.size(2)) # (batch_size*tgt_seq_len, hidden_size) logits = self.fc(output_flat) # (batch_size*tgt_seq_len, tgt_vocab_size) # Reshape back to batch-first sequence format logits = logits.reshape(tgt_input.size(0), tgt_input.size(1), -1) # (batch_size, tgt_seq_len, tgt_vocab_size) return logits, hidden, cell class Seq2Seq(nn.Module): def __init__(self, encoder, decoder): super().__init__() self.encoder = encoder self.decoder = decoder def forward(self, src_input, src_lengths, tgt_input): hidden, cell = self.encoder(src_input, src_lengths) logits, _, _ = self.decoder(tgt_input, hidden, cell) return logits # Test the model if __name__ == "__main__": src_vocab_size = 1000 tgt_vocab_size = 1200 embed_dim = 256 hidden_size = 512 num_layers = 1 encoder = Encoder(src_vocab_size, embed_dim, hidden_size, num_layers) decoder = Decoder(tgt_vocab_size, embed_dim, hidden_size, num_layers) model = Seq2Seq(encoder, decoder) # Simulate batch input batch_size = 8 src_seq_len = 15 tgt_seq_len = 20 src_input = torch.randint(0, src_vocab_size, (batch_size, src_seq_len)) # Simulate variable-length sequences (actual lengths before padding) src_lengths = torch.tensor([15, 13, 10, 8, 15, 12, 9, 14]) tgt_input = torch.randint(0, tgt_vocab_size, (batch_size, tgt_seq_len)) logits = model(src_input, src_lengths, tgt_input) print(f"Logits shape: {logits.shape}") # Expected: (8, 20, 1200)
内容的提问来源于stack exchange,提问作者Silpara

