Team Ai
Apppublic

ayushblip007/code_completion_v1

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
1likes
main.py107 linesDownload Raw Back to root
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}")