Premium problem98. Embedding Lookup

Easy Locked

embedding_matrix is (V, D) and token_ids holds integer indices of any shape. Return the embeddings, shaped token_ids.shape + (D,).

No Python loop -- indexing a tensor with an integer tensor already does this, and preserves the index tensor's shape.

Input

embedding_matrix =
tensor([[ 0.,  1.,  2.],
        [ 3.,  4.,  5.],
        [ 6.,  7.,  8.],
        [ 9., 10., 11.]])
token_ids =
tensor([[0, 2],
        [1, 3]])

Output

tensor([[[ 0.,  1.,  2.],
         [ 6.,  7.,  8.]],

        [[ 3.,  4.,  5.],
         [ 9., 10., 11.]]])

Premium problem

This one's part of Premium. Unlock the full PyTorch track plus every other premium problem on the site.

Implement solve(...)