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

如何理解PyTorch中nn.Embedding的num_embeddings与embedding_dim参数?

Understanding PyTorch's nn.Embedding: Sample Count vs Vocabulary Size

Hey there! Let's break down this confusion around nn.Embedding—it's a super common mix-up, so you're definitely not alone in wondering about this.

First, let's clarify the two key parameters and what they actually control:

  • num_embeddings: This defines the size of your vocabulary/dictionary, meaning the valid range of indices you can pass in is 0 to num_embeddings - 1 (since PyTorch uses 0-based indexing for embeddings). When you set num_embeddings=10, you're telling the layer you have 10 unique embeddings, indexed from 0 to 9.
  • The shape of your input tensor (like the 13 in your b tensor, which is (13,3)): This is just the number of samples in your batch (or sequence length, depending on your use case). This value has no connection to num_embeddings at all.

Let's walk through your example

Your tensor b has 13 samples, but every index inside those samples (1, 2, 3, 4, 5, 6, 7, 8, 9) falls within the valid 0-9 range set by num_embeddings=10. That's why there's no error—PyTorch only cares that each individual index is valid, not how many samples you're passing through the layer.

If you tried to pass an index that's outside the 0-9 range, that's when you'd get an error. For example:

# This will throw an index out of bounds error
bad_input = torch.LongTensor([[10]])  # 10 is >= num_embeddings=10
embedding(bad_input)

Quick recap to solidify this:

  • Sample count (first dimension of input tensor): Can be any positive integer—PyTorch handles batches of any size here.
  • Index values inside the input tensor: Must be between 0 and num_embeddings - 1 inclusive. Any value outside this range will trigger an error.

That's why your test tensors a, b, and c all work fine—all their indices are within the valid range, regardless of how many samples each tensor has.

内容的提问来源于stack exchange,提问作者Sean

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 08:30:06