PyTorch中如何将Transformer编码器的序列输出转为二分类单值?
Hey there! I get your frustration—Transformer encoders are great for sequence data, but when you just need a single output instead of a sequence, it's not always obvious how to bridge that gap. Let's walk through a few practical, widely-used approaches tailored to your PyTorch setup:
1. Use a CLS Token (Classification Token)
This is the standard approach from models like BERT, and it's perfect for sequence-level classification tasks like yours. Here's how it works:
- Add a special CLS token to the start of your input sequence before passing it to the encoder. This token is learned by the model to capture the global context of the entire sequence.
- After the encoder processes the sequence, you only take the output corresponding to this CLS token, then feed it through a linear layer to get your single logit for binary classification.
Code Example:
import torch import torch.nn as nn class TransformerBinaryClassifier(nn.Module): def __init__(self, input_dim=2, embed_dim=8, num_heads=2, num_layers=2): super().__init__() # Embedding layer to map 2D input to 8D embeddings self.embedding = nn.Linear(input_dim, embed_dim) # Learnable CLS token self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim)) # Transformer encoder encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=num_heads) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) # Final linear layer to get single logit self.fc = nn.Linear(embed_dim, 1) def forward(self, x): # x shape: (batch_size, seq_len=1000, input_dim=2) batch_size = x.size(0) # Embed input: (batch_size, 1000, 8) x_embed = self.embedding(x) # Expand CLS token to match batch size: (batch_size, 1, 8) cls_tokens = self.cls_token.expand(batch_size, -1, -1) # Prepend CLS token to sequence: (batch_size, 1001, 8) x_with_cls = torch.cat([cls_tokens, x_embed], dim=1) # Transformer encoder expects (seq_len, batch_size, embed_dim), so permute x_encoder_input = x_with_cls.permute(1, 0, 2) # Encoder output: (1001, batch_size, 8) encoder_output = self.encoder(x_encoder_input) # Take CLS token output: (batch_size, 8) cls_output = encoder_output[0, :, :] # Final logit: (batch_size, 1) logit = self.fc(cls_output) return logit
2. Global Pooling (Average/Max Pooling)
If you don't want to add an extra token, global pooling is a simple yet effective way to aggregate all sequence outputs into a single vector. You can use either:
- Average Pooling: Take the mean of all time-step outputs across the sequence dimension.
- Max Pooling: Take the maximum value across all time-steps for each feature dimension.
Both methods compress your (batch_size, 1000, 8) encoder output into (batch_size, 8), which you can then pass through a linear layer to get your single value.
Code Example (Average Pooling):
class TransformerBinaryClassifier(nn.Module): def __init__(self, input_dim=2, embed_dim=8, num_heads=2, num_layers=2): super().__init__() self.embedding = nn.Linear(input_dim, embed_dim) encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=num_heads) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) self.fc = nn.Linear(embed_dim, 1) def forward(self, x): # x shape: (batch_size, 1000, 2) x_embed = self.embedding(x) # (batch_size, 1000, 8) x_encoder_input = x_embed.permute(1, 0, 2) # (1000, batch_size, 8) encoder_output = self.encoder(x_encoder_input) # (1000, batch_size, 8) # Permute back to (batch_size, 1000, 8) then average over sequence pooled = torch.mean(encoder_output.permute(1, 0, 2), dim=1) # (batch_size, 8) logit = self.fc(pooled) # (batch_size, 1) return logit
3. Attention-Based Pooling
For more flexibility, you can add a lightweight attention layer to let the model automatically learn which time-steps are most important for your binary classification task. This works by computing attention weights for each time-step, then taking a weighted sum of the encoder outputs.
Code Example:
class TransformerBinaryClassifier(nn.Module): def __init__(self, input_dim=2, embed_dim=8, num_heads=2, num_layers=2): super().__init__() self.embedding = nn.Linear(input_dim, embed_dim) encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=num_heads) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) # Attention layer to compute weights for each time-step self.attention = nn.Linear(embed_dim, 1) self.fc = nn.Linear(embed_dim, 1) def forward(self, x): # x shape: (batch_size, 1000, 2) x_embed = self.embedding(x) # (batch_size, 1000, 8) x_encoder_input = x_embed.permute(1, 0, 2) # (1000, batch_size, 8) encoder_output = self.encoder(x_encoder_input) # (1000, batch_size, 8) encoder_output = encoder_output.permute(1, 0, 2) # (batch_size, 1000, 8) # Compute attention weights attn_weights = torch.softmax(self.attention(encoder_output), dim=1) # (batch_size, 1000, 1) # Weighted sum of encoder outputs pooled = torch.sum(attn_weights * encoder_output, dim=1) # (batch_size, 8) logit = self.fc(pooled) # (batch_size, 1) return logit
Quick Notes on Loss Function
- If your model outputs logits (like all the examples above), use
nn.BCEWithLogitsLoss—it combines sigmoid and binary cross-entropy into one numerically stable operation. - If you apply
torch.sigmoid()to the output first, usenn.BCELossinstead.
All these approaches do the same core thing: aggregate the sequence-level information from the Transformer encoder into a single vector, then map it to your binary classification logit. They're all valid alternatives to the RNN "take last step" approach, and which one you pick depends on how much control you want over the aggregation process!
内容的提问来源于stack exchange,提问作者jeffery_the_wind

