KiranAN1988's picture
Initial beta release
2e30a77 verified
Raw History Blame Contribute Delete
976 Bytes
import torch
import torch.nn as nn
class PositionalEmbedding(nn.Module):
def __init__(
self,
context_length,
embedding_dim
):
super().__init__()
self.embedding = nn.Embedding(
context_length,
embedding_dim
)
def forward(self, x):
# Case 1: 1D position tensor passed directly during KV-cached steps, e.g. tensor([past_len, ..., total_len-1])
if x.dim() == 1:
positions = x
# Case 2: 2D input_ids tensor [batch_size, sequence_length]
elif x.dim() == 2:
sequence_length = x.size(1)
positions = torch.arange(
sequence_length,
device=x.device
)
else:
raise ValueError(
f"Expected 1D or 2D tensor, but got input of shape {x.shape}"
)
return self.embedding(
positions
)