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 )