All problems
EasyNeural Networks & LLM Componentsv2026-07-07

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]]])
solution.pySign in to save
PyTorch runs in an isolated cloud sandbox. Sign in to keep it warm between runs.
Run your code to see public test results.