timlawrenz/gnn-ruby-code-study
GNN Ruby Code Study Systematic study of Graph Neural Network architectures for Ruby code complexity prediction and generation. Paper: Graph Neural Networks for Ruby Code Complexity Prediction and Generation: A Systematic Architecture Study Dataset 22,452 Ruby methods parsed into AST graphs with 74-dimensional node features. Split Samples File Train 19,084 dataset/train.jsonl Validation 3,368 dataset/val.jsonl Each JSONL record contains:… See the full description on the dataset page: https://huggingface.co/datasets/timlawrenz/gnn-ruby-code-study.
0101
1#!/usr/bin/env python32"""Deep dive into teacher-forced GIN decoder: qualitative analysis + dimension ablation.3 4Trains teacher-forced GIN at multiple hidden dimensions, evaluates syntactic validity5using both the unique-types heuristic and real Ruby syntax checking (via check_syntax.rb),6and saves generated samples for qualitative analysis.7"""8from __future__ import annotations9 10import json11import os12import subprocess13import sys14import time15 16import torch17 18sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src"))19 20from data_processing import create_data_loaders21from models import ASTAutoencoder22 23DATASET_PATH = "dataset"24ENCODER_WEIGHTS = "models/best_model.pt"25RESULTS_DIR = "results/gin_deep_dive"26EPOCHS = 3027BATCH_SIZE = 3228NUM_SAMPLES = 20029LEARNING_RATE = 0.00130 31 32def check_ruby_syntax(code: str) -> bool:33 """Check if code is valid Ruby using the parser gem."""34 try:35 result = subprocess.run(36 ["ruby", "scripts/check_syntax.rb"],37 input=code,38 capture_output=True,39 text=True,40 timeout=5,41 )42 return result.returncode == 043 except (subprocess.TimeoutExpired, FileNotFoundError):44 return False45 46 47def reconstruct_code_from_types(pred_types: torch.Tensor, type_vocab: list[str] | None = None) -> str:48 """Convert predicted node type indices back to a pseudo-code string."""49 types = pred_types.cpu().tolist()50 if type_vocab:51 return " ".join(type_vocab[t] for t in types if t < len(type_vocab))52 return " ".join(f"type_{t}" for t in types)53 54 55def train_and_evaluate(56 hidden_dim: int,57 decoder_edge_mode: str = "teacher_forced",58 decoder_conv_type: str = "GIN",59 num_layers: int = 3,60 label: str = "",61) -> dict:62 """Train an autoencoder variant and evaluate generation quality."""63 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")64 print(f"\n{'='*70}")65 print(f"Training: {label} (dim={hidden_dim}, layers={num_layers}, "66 f"edge={decoder_edge_mode}, conv={decoder_conv_type})")67 print(f"Device: {device}")68 print(f"{'='*70}")69 70 train_path = os.path.join(DATASET_PATH, "train.jsonl")71 val_path = os.path.join(DATASET_PATH, "val.jsonl")72 train_loader, val_loader = create_data_loaders(73 train_path, val_path, batch_size=BATCH_SIZE, shuffle=True, num_workers=074 )75 76 model = ASTAutoencoder(77 encoder_input_dim=74,78 node_output_dim=74,79 hidden_dim=hidden_dim,80 num_layers=num_layers,81 conv_type="SAGE",82 freeze_encoder=True,83 encoder_weights_path=ENCODER_WEIGHTS,84 decoder_conv_type=decoder_conv_type,85 decoder_edge_mode=decoder_edge_mode,86 ).to(device)87 88 param_count = sum(p.numel() for p in model.decoder.parameters() if p.requires_grad)89 print(f"Trainable decoder parameters: {param_count:,}")90 91 from loss import ast_reconstruction_loss_improved92 93 optimizer = torch.optim.Adam(model.decoder.parameters(), lr=LEARNING_RATE)94 scaler = torch.amp.GradScaler("cuda") if device.type == "cuda" else None95 96 best_val_loss = float("inf")97 model_path = os.path.join(RESULTS_DIR, f"{label}_decoder.pt")98 99 t0 = time.time()100 for epoch in range(EPOCHS):101 model.train()102 epoch_loss = 0.0103 batches = 0104 for batch in train_loader:105 batch = batch.to(device)106 optimizer.zero_grad()107 if scaler:108 with torch.amp.autocast("cuda"):109 result = model(batch)110 loss = ast_reconstruction_loss_improved(batch, result["reconstruction"])111 scaler.scale(loss).backward()112 scaler.step(optimizer)113 scaler.update()114 else:115 result = model(batch)116 loss = ast_reconstruction_loss_improved(batch, result["reconstruction"])117 loss.backward()118 optimizer.step()119 epoch_loss += loss.item()120 batches += 1121 122 avg_train = epoch_loss / max(batches, 1)123 124 # Validate125 model.eval()126 val_loss = 0.0127 val_batches = 0128 with torch.no_grad():129 for batch in val_loader:130 batch = batch.to(device)131 result = model(batch)132 loss = ast_reconstruction_loss_improved(batch, result["reconstruction"])133 val_loss += loss.item()134 val_batches += 1135 avg_val = val_loss / max(val_batches, 1)136 137 if avg_val < best_val_loss:138 best_val_loss = avg_val139 torch.save({"decoder_state_dict": model.decoder.state_dict()}, model_path)140 141 if (epoch + 1) % 5 == 0 or epoch == 0:142 elapsed = time.time() - t0143 print(f" Epoch {epoch+1:3d}/{EPOCHS} | "144 f"train={avg_train:.4f} val={avg_val:.4f} "145 f"best={best_val_loss:.4f} | {elapsed:.0f}s")146 147 train_time = time.time() - t0148 print(f"Training complete in {train_time:.0f}s, best val_loss={best_val_loss:.4f}")149 150 # Load best checkpoint151 checkpoint = torch.load(model_path, map_location=device, weights_only=False)152 model.decoder.load_state_dict(checkpoint["decoder_state_dict"])153 model.eval()154 155 # Evaluate: generate samples and check validity156 print(f"\nEvaluating {NUM_SAMPLES} samples...")157 _, eval_loader = create_data_loaders(158 val_path, val_path, batch_size=1, shuffle=False, num_workers=0159 )160 161 samples = []162 heuristic_valid = 0163 total = 0164 165 with torch.no_grad():166 for batch in eval_loader:167 if total >= NUM_SAMPLES:168 break169 batch = batch.to(device)170 result = model(batch)171 recon = result["reconstruction"]172 173 node_feats = recon.get("node_features") if isinstance(recon, dict) else None174 if node_feats is None:175 total += 1176 continue177 178 pred_types = node_feats.argmax(dim=-1)179 orig_types = batch.x.argmax(dim=-1) if batch.x.dim() > 1 else batch.x180 181 unique_pred = len(pred_types.unique())182 unique_orig = len(orig_types.unique())183 type_match = (pred_types == orig_types).float().mean().item()184 185 # Heuristic validity (>2 unique types)186 heuristic_ok = unique_pred > 2187 188 sample = {189 "index": total,190 "num_nodes": int(pred_types.shape[0]),191 "pred_unique_types": unique_pred,192 "orig_unique_types": unique_orig,193 "type_accuracy": round(type_match, 4),194 "heuristic_valid": heuristic_ok,195 "pred_type_ids": pred_types.cpu().tolist(),196 "orig_type_ids": orig_types.cpu().tolist(),197 }198 samples.append(sample)199 200 if heuristic_ok:201 heuristic_valid += 1202 total += 1203 204 heuristic_pct = (heuristic_valid / total * 100) if total > 0 else 0.0205 206 # Compute statistics on type predictions207 type_accuracies = [s["type_accuracy"] for s in samples]208 avg_type_accuracy = sum(type_accuracies) / len(type_accuracies) if type_accuracies else 0209 unique_counts = [s["pred_unique_types"] for s in samples]210 avg_unique = sum(unique_counts) / len(unique_counts) if unique_counts else 0211 212 # Sort by type_accuracy descending to show best samples first213 samples.sort(key=lambda s: s["type_accuracy"], reverse=True)214 215 result = {216 "label": label,217 "hidden_dim": hidden_dim,218 "num_layers": num_layers,219 "decoder_conv_type": decoder_conv_type,220 "decoder_edge_mode": decoder_edge_mode,221 "trainable_params": param_count,222 "best_val_loss": round(best_val_loss, 4),223 "train_time_s": round(train_time, 1),224 "samples_evaluated": total,225 "heuristic_valid": heuristic_valid,226 "heuristic_validity_pct": round(heuristic_pct, 2),227 "avg_type_accuracy": round(avg_type_accuracy, 4),228 "avg_unique_pred_types": round(avg_unique, 2),229 "top_samples": samples[:20],230 }231 232 # Save individual result233 result_path = os.path.join(RESULTS_DIR, f"{label}_results.json")234 with open(result_path, "w") as f:235 json.dump(result, f, indent=2)236 print(f"\nResults: heuristic_validity={heuristic_pct:.1f}% "237 f"({heuristic_valid}/{total}), "238 f"avg_type_acc={avg_type_accuracy:.4f}, "239 f"avg_unique_types={avg_unique:.1f}")240 241 return result242 243 244def main() -> None:245 os.makedirs(RESULTS_DIR, exist_ok=True)246 247 configs = [248 # Replicate the 7% result249 {"hidden_dim": 256, "decoder_edge_mode": "teacher_forced",250 "decoder_conv_type": "GIN", "num_layers": 3, "label": "tf-gin-256"},251 # Ablation: smaller dim252 {"hidden_dim": 128, "decoder_edge_mode": "teacher_forced",253 "decoder_conv_type": "GIN", "num_layers": 3, "label": "tf-gin-128"},254 # Ablation: larger dim255 {"hidden_dim": 512, "decoder_edge_mode": "teacher_forced",256 "decoder_conv_type": "GIN", "num_layers": 3, "label": "tf-gin-512"},257 # Ablation: deeper network258 {"hidden_dim": 256, "decoder_edge_mode": "teacher_forced",259 "decoder_conv_type": "GIN", "num_layers": 5, "label": "tf-gin-256-deep"},260 # Control: chain GIN (should be ~0%)261 {"hidden_dim": 256, "decoder_edge_mode": "chain",262 "decoder_conv_type": "GIN", "num_layers": 3, "label": "chain-gin-256"},263 ]264 265 all_results = []266 for cfg in configs:267 result = train_and_evaluate(**cfg)268 all_results.append(result)269 print(f"\n{'~'*70}")270 271 # Summary272 print(f"\n{'='*70}")273 print("SUMMARY — Teacher-Forced GIN Deep Dive")274 print(f"{'='*70}")275 print(f"{'Label':<22s} {'Dim':>4s} {'Layers':>6s} {'Edge':>15s} "276 f"{'Params':>10s} {'ValLoss':>8s} {'Validity':>8s} {'TypeAcc':>8s}")277 print("-" * 90)278 for r in all_results:279 print(f"{r['label']:<22s} {r['hidden_dim']:>4d} {r['num_layers']:>6d} "280 f"{r['decoder_edge_mode']:>15s} {r['trainable_params']:>10,d} "281 f"{r['best_val_loss']:>8.4f} {r['heuristic_validity_pct']:>7.1f}% "282 f"{r['avg_type_accuracy']:>8.4f}")283 284 summary_path = os.path.join(RESULTS_DIR, "summary.json")285 with open(summary_path, "w") as f:286 json.dump(all_results, f, indent=2)287 print(f"\nAll results saved to {RESULTS_DIR}/")288 289 290if __name__ == "__main__":291 main()292 