如何在PyTorch中高效解码嵌入?Python2替代@运算符咨询
Hey there! Great job figuring out that torch.mm is the right replacement for Python 2's lack of the @ matrix multiplication operator—you're already on the right track. Let's break down better, more professional ways to handle embedding decoding, since your current loop-based approach can get slow with larger vocabularies and might fail due to floating-point precision issues.
First: Confirming your torch.mm fix is correct
In Python 3, the @ operator for PyTorch tensors maps directly to matrix multiplication, which is exactly what torch.mm() does for 2D tensors. Your line:
out_unembedded = torch.mm(embedded_output.view(-1, hidden_size), embedding.weight.transpose(0, 1))
is totally the right way to replicate the Python 3 code in Python 2. This will efficiently compute the unembedded logits just like the original tutorial.
Better Embedding Decoding (Replacing the Loop)
Your current loop checks for exact element-wise equality between the target embedding and each row in the embedding weight matrix. This works for toy examples, but in real training, embedding vectors are continuous floating-point values—exact matches are rare, and looping over every row is inefficient for large vocabularies. Here are two better approaches:
1. L2 Distance Matching (Find the Closest Embedding)
Calculate the L2 (Euclidean) distance between your target embedding and all embeddings in the weight matrix, then pick the index with the smallest distance:
import torch import torch.nn as nn from torch.autograd import Variable word2id = {'hello': 0, 'world': 1, 'I': 2, 'am': 3,'writing': 4,'pytorch': 5} embeds = nn.Embedding(6, 3) id2word = {v: k for k, v in word2id.iteritems()} # Get target embedding target_word = 'am' word_embed = embeds(Variable(torch.LongTensor([word2id[target_word]]))) # Calculate L2 distance between target and all embeddings distances = torch.norm(embeds.weight - word_embed.data, dim=1) # Find index of smallest distance closest_idx = torch.argmin(distances).item() print("Closest index: {}".format(closest_idx)) print("Decoded word: {}".format(id2word[closest_idx]))
2. Matrix Multiplication for Similarity (Logit-Based Decoding)
This is the approach the tutorial hints at—when you compute out_unembedded, you're essentially calculating the similarity (logits) between the model's output and each embedding. For a single target embedding, you can do the same to find the most similar embedding:
# Compute similarity scores (inner product) between target and all embeddings similarity_scores = torch.mm(word_embed.data, embeds.weight.transpose(0, 1)) # Find index of highest score best_idx = torch.argmax(similarity_scores).item() print("Best matching index: {}".format(best_idx)) print("Decoded word: {}".format(id2word[best_idx]))
If you normalize the embeddings first (using nn.Embedding(..., norm_type=2) or manually normalizing), this becomes equivalent to cosine similarity, which is a common metric for embedding matching.
Why These Are Better
- Efficiency: Both methods use vectorized operations that run on GPU (if available) instead of Python loops, which are orders of magnitude slower for large vocabularies.
- Robustness: Floating-point precision issues won't break these methods—they don't rely on exact equality.
- Scalability: They work seamlessly even if your vocabulary has thousands or tens of thousands of words.
内容的提问来源于stack exchange,提问作者Pingjiang Li

