ashritaa-chinnari123's picture
Upload 8 files
38b27cd verified
Raw
History Blame Contribute Delete
5.92 kB
"""
Optional: LSTM baseline for BUG DETECTION (binary classification: buggy vs clean).
This exists so your project report can include a comparison table like:
Model Accuracy F1
----------------------------------------
LSTM (from scratch) 71% 0.68
GRU (from scratch) 73% 0.70
CodeT5 (transformer) 89% 0.87
Why it's weaker: LSTM/GRU process code token-by-token in sequence, so
relationships between distant tokens (e.g. a variable used far from where
it's defined, or a missing bracket 40 tokens later) are harder to learn.
Transformers see all tokens at once via self-attention, so they capture
those long-range dependencies much better -- which is exactly what matters
for code.
This is a self-contained PyTorch script -- no Hugging Face needed for this part.
"""
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from collections import Counter
import re
# ---------------------------------------------------------------------------
# 1. Tokenizer (very simple, word/symbol level -- fine for a baseline)
# ---------------------------------------------------------------------------
def simple_tokenize(code: str):
return re.findall(r"\w+|[^\s\w]", code)
class Vocab:
def __init__(self, token_lists, min_freq=1):
counter = Counter(tok for toks in token_lists for tok in toks)
self.itos = ["<pad>", "<unk>"] + [
tok for tok, freq in counter.items() if freq >= min_freq
]
self.stoi = {tok: i for i, tok in enumerate(self.itos)}
def encode(self, tokens, max_len):
ids = [self.stoi.get(tok, 1) for tok in tokens][:max_len]
ids += [0] * (max_len - len(ids))
return ids
def __len__(self):
return len(self.itos)
# ---------------------------------------------------------------------------
# 2. Dataset
# ---------------------------------------------------------------------------
class CodeBugDataset(Dataset):
"""
Expects a list of (code_string, label) pairs, label = 1 if buggy else 0.
Replace `load_your_data()` with loading from CodeXGLUE / your own CSV.
"""
def __init__(self, samples, vocab, max_len=128):
self.samples = samples
self.vocab = vocab
self.max_len = max_len
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
code, label = self.samples[idx]
tokens = simple_tokenize(code)
ids = self.vocab.encode(tokens, self.max_len)
return torch.tensor(ids, dtype=torch.long), torch.tensor(label, dtype=torch.float)
# ---------------------------------------------------------------------------
# 3. Model: swap nn.LSTM for nn.GRU or nn.RNN to compare all three
# ---------------------------------------------------------------------------
class RecurrentBugClassifier(nn.Module):
def __init__(self, vocab_size, embed_dim=128, hidden_dim=128, cell_type="LSTM"):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
cell_type = cell_type.upper()
if cell_type == "LSTM":
self.rnn = nn.LSTM(embed_dim, hidden_dim, batch_first=True, bidirectional=True)
elif cell_type == "GRU":
self.rnn = nn.GRU(embed_dim, hidden_dim, batch_first=True, bidirectional=True)
elif cell_type == "RNN":
self.rnn = nn.RNN(embed_dim, hidden_dim, batch_first=True, bidirectional=True)
else:
raise ValueError("cell_type must be one of: LSTM, GRU, RNN")
self.classifier = nn.Sequential(
nn.Linear(hidden_dim * 2, 64),
nn.ReLU(),
nn.Dropout(0.3),
nn.Linear(64, 1),
)
def forward(self, x):
embedded = self.embedding(x) # (batch, seq_len, embed_dim)
output, _ = self.rnn(embedded) # (batch, seq_len, hidden_dim*2)
pooled = output.mean(dim=1) # mean pooling over time steps
logits = self.classifier(pooled).squeeze(-1)
return logits
# ---------------------------------------------------------------------------
# 4. Training loop
# ---------------------------------------------------------------------------
def train_baseline(cell_type="LSTM", epochs=10, batch_size=16, lr=1e-3):
# --- Replace this with real data, e.g. loaded from CodeXGLUE defect-detection ---
samples = [
("def add(a, b):\n return a + b", 0),
("def add(a, b)\n return a + b", 1), # missing colon
("for i in range(10):\n print(i)", 0),
("for i in range(10)\n print(i)", 1), # missing colon
("if x == 1:\n print('one')", 0),
("if x = 1:\n print('one')", 1), # assignment vs equality
] * 20 # repeat for a runnable toy example; use a real dataset for real results
tokenized = [simple_tokenize(c) for c, _ in samples]
vocab = Vocab(tokenized)
dataset = CodeBugDataset(samples, vocab)
loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
model = RecurrentBugClassifier(vocab_size=len(vocab), cell_type=cell_type)
optimizer = torch.optim.Adam(model.parameters(), lr=lr)
criterion = nn.BCEWithLogitsLoss()
model.train()
for epoch in range(epochs):
total_loss = 0.0
for x, y in loader:
optimizer.zero_grad()
logits = model(x)
loss = criterion(logits, y)
loss.backward()
optimizer.step()
total_loss += loss.item()
print(f"[{cell_type}] Epoch {epoch+1}/{epochs} - loss: {total_loss/len(loader):.4f}")
return model, vocab
if __name__ == "__main__":
# Train and compare all three cell types
for cell_type in ["RNN", "GRU", "LSTM"]:
print(f"\n=== Training {cell_type} baseline ===")
train_baseline(cell_type=cell_type, epochs=5)