import torch import torch.nn as nn # Parameters INPUT_SIZE = 10 # Numbers 0–9 HIDDEN_SIZE = 16 OUTPUT_SIZE = 10 # Same as input SEQ_LEN = 3 # Encoder: RNN that reads the input class Encoder(nn.Module): def __init__(self): super().__init__() self.embedding = nn.Embedding(INPUT_SIZE, HIDDEN_SIZE) self.rnn = nn.GRU(HIDDEN_SIZE, HIDDEN_SIZE) def forward(self, input_seq): embedded = self.embedding(input_seq).unsqueeze(1) outputs, hidden = self.rnn(embedded) return hidden # Decoder: RNN that writes the output, one token at a time class Decoder(nn.Module): def __init__(self): super().__init__() self.embedding = nn.Embedding(OUTPUT_SIZE, HIDDEN_SIZE) self.rnn = nn.GRU(HIDDEN_SIZE, HIDDEN_SIZE) self.fc = nn.Linear(HIDDEN_SIZE, OUTPUT_SIZE) def forward(self, input_token, hidden): embedded = self.embedding(input_token).unsqueeze(0) output, hidden = self.rnn(embedded, hidden) logits = self.fc(output.squeeze(0)) return logits, hidden # Example: run once (no training) encoder = Encoder() decoder = Decoder() input_seq = torch.tensor([1, 2, 3], dtype=torch.long) target_seq = torch.tensor([3, 2, 1], dtype=torch.long) # Encode context = encoder(input_seq) # Decode step by step decoded_tokens = [] decoder_input = torch.tensor([0]) # Start token (use 0 here) hidden = context for i in range(SEQ_LEN): logits, hidden = decoder(decoder_input, hidden) prediction = logits.argmax(1) decoded_tokens.append(prediction.item()) decoder_input = prediction print("Predicted output:", decoded_tokens)