Neural Networks & LLM Components ยท Easy
PyTorch Batched Embedding Lookup
Look up token embeddings while preserving batch and sequence axes.
Task
Implement batched_embedding_lookup(embedding_table: torch.Tensor, token_ids: torch.Tensor) -> torch.Tensor. Return the embedding vector for every token id in a batched token-id tensor.
Requirements
- embedding_table has shape (vocab_size, embedding_dim).
- token_ids has shape (B, T) and integer dtype.
- Return a tensor with shape (B, T, embedding_dim).
- Preserve the batch and time axes from token_ids.
- Do not use explicit Python loops over batch or time.
- Preserve embedding_table dtype and device in the returned tensor.
Example
batched_embedding_lookup(
torch.tensor([[0.0, 0.0], [1.0, 1.5], [2.0, 2.5]]),
torch.tensor([[1, 2], [0, 1]]),
)
# tensor([[[1.0, 1.5], [2.0, 2.5]],
# [[0.0, 0.0], [1.0, 1.5]]])