ayushblip007/code_completion_v1
1
1import torch2import torch.nn as nn3import pickle4 5# --- Part 1: Re-define the Model Architecture ---6# This class definition must be EXACTLY the same as in your training script.7 8class ResidualLSTMModel(nn.Module):9 def __init__(self, vocab_size, embedding_dim, hidden_units, dropout_prob):10 super(ResidualLSTMModel, self).__init__()11 self.embedding = nn.Embedding(12 num_embeddings=vocab_size,13 embedding_dim=embedding_dim,14 padding_idx=015 )16 self.lstm1 = nn.LSTM(17 input_size=embedding_dim,18 hidden_size=hidden_units,19 num_layers=1,20 batch_first=True21 )22 self.lstm2 = nn.LSTM(23 input_size=hidden_units,24 hidden_size=hidden_units,25 num_layers=1,26 batch_first=True27 )28 self.dropout = nn.Dropout(dropout_prob)29 self.fc = nn.Linear(hidden_units, vocab_size)30 31 def forward(self, x):32 embedded = self.embedding(x)33 out1, _ = self.lstm1(embedded)34 out2, _ = self.lstm2(out1)35 residual_sum = out1 + out236 dropped_out = self.dropout(residual_sum)37 logits = self.fc(dropped_out)38 return logits39 40# --- Part 2: Helper Functions for Processing Text ---41 42def text_to_sequence(text, vocab, max_length):43 """Converts a string of code into a padded tensor."""44 tokens = text.split()45 numericalized = [vocab.get(token, vocab['<UNK>']) for token in tokens]46 47 if len(numericalized) > max_length:48 numericalized = numericalized[:max_length]49 50 pad_id = vocab['<PAD>']51 padding_needed = max_length - len(numericalized)52 padded = numericalized + [pad_id] * padding_needed53 54 return torch.tensor([padded], dtype=torch.long)55 56def sequence_to_text(sequence, vocab):57 """Converts a tensor of token IDs back to a string."""58 id_to_token = {id_val: token for token, id_val in vocab.items()}59 tokens = [id_to_token.get(id_val.item(), '<UNK>') for id_val in sequence if id_val.item() != vocab['<PAD>']]60 return " ".join(tokens)61 62# --- Part 3: Main Prediction Logic ---63 64def predict_next_tokens(model, text, vocab, device, max_length=1000, top_k=5):65 """Predicts the top_k next tokens for a given text input."""66 model.eval()67 with torch.no_grad():68 input_tensor = text_to_sequence(text, vocab, max_length).to(device)69 logits = model(input_tensor)70 71 num_input_tokens = len(text.split())72 last_token_logits = logits[0, num_input_tokens - 1, :]73 74 _, top_k_ids = torch.topk(last_token_logits, top_k, dim=-1)75 top_k_tokens = [sequence_to_text([token_id], vocab) for token_id in top_k_ids]76 77 return top_k_tokens78 79if __name__ == '__main__':80 # --- Configuration ---81 MODEL_PATH = 'model.pt'82 VOCAB_PATH = 'vocab.pkl' # <-- Updated to use .pkl83 MAX_LENGTH = 100084 85 # --- Load everything ---86 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")87 print(f"Using device: {device}")88 89 # Load vocabulary using pickle90 with open(VOCAB_PATH, 'rb') as f: # <-- Use 'rb' for reading bytes91 vocab = pickle.load(f)92 print("Vocabulary loaded.")93 94 # Load the model95 model = torch.load(MODEL_PATH, map_location=device , weights_only=False)96 print("Model loaded.")97 98 # --- Make a Prediction ---99 input_code = "import numpy as" # Example input100 101 print(f"\nInput code: '{input_code}'")102 103 suggestions = predict_next_tokens(model, input_code, vocab, device, max_length=MAX_LENGTH)104 105 print("\nTop 5 suggestions:")106 for i, suggestion in enumerate(suggestions):107 print(f"{i+1}. {suggestion}")