Spaces:
Sleeping
Sleeping
| import torch | |
| import torch.nn.functional as F | |
| from torch_geometric.nn import GCNConv | |
| import re | |
| import string | |
| import os | |
| import pickle | |
| # Define the BiGNN model | |
| class BiGNN(torch.nn.Module): | |
| def __init__(self, input_dim, hidden_dim, output_dim): | |
| super(BiGNN, self).__init__() | |
| self.conv1 = GCNConv(input_dim, hidden_dim) | |
| self.conv2 = GCNConv(hidden_dim, output_dim) | |
| def forward(self, data): | |
| x, edge_index = data.x, data.edge_index | |
| x = F.relu(self.conv1(x, edge_index)) | |
| x = F.dropout(x, training=self.training) | |
| x = self.conv2(x, edge_index) | |
| return F.log_softmax(x, dim=1) | |
| # Simple text preprocessing | |
| def clean_text(text): | |
| text = text.lower() | |
| text = re.sub(r'\[.*?\]', '', text) | |
| text = re.sub(r'https?://\S+|www\.\S+', '', text) | |
| text = re.sub(r'<.*?>+', '', text) | |
| text = re.sub(f'[{re.escape(string.punctuation)}]', '', text) | |
| text = re.sub(r'\n', ' ', text) | |
| text = re.sub(r'\w*\d\w*', '', text) | |
| text = re.sub(' +', ' ', text) | |
| return text.strip() | |
| # Load pre-trained components | |
| def load_model_and_vectorizer(model_path='question_api/artifacts/biggn_model.pt', | |
| vectorizer_path='question_api/artifacts/vectorizer.pkl'): | |
| with open(vectorizer_path, 'rb') as f: | |
| vectorizer = pickle.load(f) | |
| input_dim = vectorizer.transform(["sample"]).shape[1] # Feature size | |
| model = BiGNN(input_dim=input_dim, hidden_dim=16, output_dim=3) | |
| model.load_state_dict(torch.load(model_path, map_location=torch.device('cpu'))) | |
| model.eval() | |
| return model, vectorizer | |
| # Classify content into [0 = discard, 1 = MCQ-worthy, 2 = Theory-worthy] | |
| def classify_sentences(sentences): | |
| model, vectorizer = load_model_and_vectorizer() | |
| results = [] | |
| for sentence in sentences: | |
| cleaned = clean_text(sentence) | |
| vector = vectorizer.transform([cleaned]) | |
| features = torch.tensor(vector.toarray(), dtype=torch.float32) | |
| data = type('Data', (object,), {})() # Dummy PyG-like object | |
| data.x = features | |
| data.edge_index = torch.tensor([[0], [0]], dtype=torch.long) | |
| output = model(data) | |
| pred = torch.argmax(output, dim=1).item() | |
| results.append((sentence, pred)) | |
| return results | |