Views
No views yet
1import torch
2import torch.nn as nn
3from transformers import AutoConfig, AutoTokenizer, AutoModel
4from huggingface_hub import hf_hub_download
5import json
6from types import SimpleNamespace
7
8#model architecture - needed since this is a custom model
9class ProjectionModel(nn.Module):
10 def __init__(self, config):
11 super(ProjectionModel, self).__init__()
12 self.config = config
13 self.c_code_encoder = AutoModel.from_pretrained("microsoft/codebert-base")
14 self.pseudocode_encoder = AutoModel.from_pretrained("microsoft/codebert-base")
15
16 #Projection network with 1 hidden layer
17 self.projection = nn.Sequential(
18 nn.Linear(config.embedding_dim, config.hidden_dim),
19 nn.ReLU(), # First activation function
20 nn.Linear(config.hidden_dim, config.hidden_dim), #Hidden layer
21 nn.ReLU(), # Second activation function
22 nn.Linear(config.hidden_dim, config.embedding_dim) #Output layer projecting back to the original embedding space
23 )
24
25 def forward(self, c_code_inputs, pseudocode_inputs):
26 #Encode C code and pseudocode
27 c_code_embedding = self.c_code_encoder(**c_code_inputs).last_hidden_state.mean(dim=1)
28 pseudocode_embedding = self.pseudocode_encoder(**pseudocode_inputs).last_hidden_state.mean(dim=1)
29
30 #Apply the projection network to the pseudocode embeddings
31 projected_pseudocode_embedding = self.projection(pseudocode_embedding)
32
33 return c_code_embedding, projected_pseudocode_embedding
34
35model_name = "aircrypto/code-llama-7b-projection-largev2.11"
36config_file = hf_hub_download(repo_id=model_name, filename="config.json")
37
38with open(config_file, 'r') as f:
39 config_dict = json.load(f)
40config = SimpleNamespace(**config_dict)
41
42model = ProjectionModel(config)
43model_path = hf_hub_download(repo_id=model_name, filename="pytorch_model.bin")
44state_dict = torch.load(model_path, map_location="cpu")
45model.load_state_dict(state_dict)
46
47device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
48model = model.to(device)
49
50print("Model loaded successfully!")
51
52tokenizer = AutoTokenizer.from_pretrained("aircrypto/code-llama-7b-projection-largev2.11")
53
54print("Tokenizer loaded successfully!")