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

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.

Seq2Seq 编解码器各层输入输出详解(批量训练场景)

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 its embed_dim-dimensional vector.
  • Batch-specific notes:
    • Sequences in a batch rarely have the same length, so use torch.nn.utils.rnn.pad_sequence (with batch_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.

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.
  • 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 as h_n.
    • output (optional): The full sequence of LSTM outputs. If you used pack_padded_sequence, this will be a PackedSequence; use pad_packed_sequence to 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_n and c_n (make sure the number of layers matches between encoder and decoder!).
  • 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).
  • 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.
  • 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's CrossEntropyLoss works with (N, C) inputs, but reshaping to batch-first sequence format makes it easier to align with target labels).
Full Batch Input Preparation Workflow
  1. Tokenize & Index: Convert raw text sequences into integer tensors of word indices.
  2. Pad Sequences: Use pad_sequence to make all sequences in the batch the same length (batch-first shape).
  3. Track Lengths: Save the original length of each sequence (before padding) for the encoder's packed sequences.
  4. Optional (But Recommended): Sort sequences by length in descending order — this optimizes pack_padded_sequence performance.
  5. Transpose for LSTM: Convert embedding outputs from batch-first to sequence-first format for LSTM input.
Working Example Code
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 03:46:03