Team Ai
Datasetpublic

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.

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes101downloads
qualitative_eval.py259 linesDownload Raw Back to scripts
1#!/usr/bin/env python32"""Qualitative analysis: reconstruct Ruby code from teacher-forced GIN and check syntax.3 4Uses the full pipeline: model → predict types + parents → build AST tree → Ruby pretty-print → syntax check.5"""6from __future__ import annotations7 8import json9import os10import subprocess11import sys12 13import torch14 15sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src"))16 17from data_processing import ASTGraphConverter, ASTNodeEncoder, create_data_loaders18from models import ASTAutoencoder19from torch_geometric.data import Data20 21RESULTS_DIR = "results/gin_deep_dive"22DATASET_PATH = "dataset"23ENCODER_WEIGHTS = "models/best_model.pt"24NUM_SAMPLES = 20025 26 27def run_ruby_script(script_path: str, stdin_data: str) -> str:28    """Run a Ruby script with stdin and return stdout."""29    try:30        result = subprocess.run(31            ["ruby", script_path],32            input=stdin_data,33            capture_output=True,34            text=True,35            timeout=10,36        )37        return result.stdout.strip()38    except Exception:39        return ""40 41 42def check_ruby_syntax(code: str) -> bool:43    """Check if code is valid Ruby."""44    if not code or not code.strip():45        return False46    try:47        result = subprocess.run(48            ["ruby", "scripts/check_syntax.rb"],49            input=code,50            capture_output=True,51            text=True,52            timeout=5,53        )54        return result.returncode == 055    except Exception:56        return False57 58 59def evaluate_model_full(model_path: str, label: str, hidden_dim: int,60                        num_layers: int, decoder_conv_type: str,61                        decoder_edge_mode: str) -> dict:62    """Full evaluation: model → AST → Ruby code → syntax check."""63    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")64 65    converter = ASTGraphConverter()66 67    model = ASTAutoencoder(68        encoder_input_dim=74,69        node_output_dim=74,70        hidden_dim=hidden_dim,71        num_layers=num_layers,72        conv_type="SAGE",73        freeze_encoder=True,74        encoder_weights_path=ENCODER_WEIGHTS,75        decoder_conv_type=decoder_conv_type,76        decoder_edge_mode=decoder_edge_mode,77    ).to(device)78 79    checkpoint = torch.load(model_path, map_location=device, weights_only=False)80    model.decoder.load_state_dict(checkpoint["decoder_state_dict"])81    model.eval()82 83    val_path = os.path.join(DATASET_PATH, "val.jsonl")84    _, val_loader = create_data_loaders(85        val_path, val_path, batch_size=1, shuffle=False, num_workers=086    )87 88    # Also load raw source for comparison89    raw_sources = []90    with open(val_path) as f:91        for line in f:92            d = json.loads(line)93            raw_sources.append(d["raw_source"])94 95    results = []96    syntax_valid = 097    heuristic_valid = 098    total = 099 100    print(f"\nEvaluating {label} on {NUM_SAMPLES} samples...")101 102    with torch.no_grad():103        for batch in val_loader:104            if total >= NUM_SAMPLES:105                break106 107            batch = batch.to(device)108            result = model(batch)109            recon = result["reconstruction"]110 111            node_feats = recon.get("node_features") if isinstance(recon, dict) else None112            parent_logits = recon.get("parent_logits") if isinstance(recon, dict) else None113 114            if node_feats is None:115                total += 1116                continue117 118            pred_types = node_feats.argmax(dim=-1)119            orig_types = batch.x.argmax(dim=-1) if batch.x.dim() > 1 else batch.x120 121            unique_pred = len(pred_types.unique())122            type_match = (pred_types == orig_types).float().mean().item()123            heuristic_ok = unique_pred > 2124 125            # Build reconstructed AST tree from types + parents126            reconstructed_code = ""127            if parent_logits is not None:128                parent_preds = parent_logits.squeeze(0).argmax(dim=-1)129                node_map = {}130                for i, type_idx in enumerate(pred_types.cpu().numpy()):131                    nt = converter.node_encoder.node_types[type_idx] \132                        if type_idx < len(converter.node_encoder.node_types) else "unknown"133                    node_map[i] = {"type": nt, "children": []}134 135                root_nodes = []136                for i, parent_idx in enumerate(parent_preds.cpu().numpy()):137                    pi = int(parent_idx)138                    if i == pi or pi >= len(node_map):139                        root_nodes.append(node_map[i])140                    elif pi in node_map:141                        node_map[pi]["children"].append(node_map[i])142 143                if root_nodes:144                    ast_json = json.dumps(root_nodes)145                    reconstructed_code = run_ruby_script(146                        "scripts/pretty_print_ast.rb", ast_json147                    )148 149            is_valid = check_ruby_syntax(reconstructed_code)150 151            sample = {152                "index": total,153                "num_nodes": int(pred_types.shape[0]),154                "type_accuracy": round(type_match, 4),155                "heuristic_valid": heuristic_ok,156                "syntax_valid": is_valid,157                "original_code": raw_sources[total] if total < len(raw_sources) else "",158                "reconstructed_code": reconstructed_code,159                "pred_types": [converter.node_encoder.node_types[t]160                               if t < len(converter.node_encoder.node_types) else "unknown"161                               for t in pred_types.cpu().tolist()],162            }163            results.append(sample)164 165            if heuristic_ok:166                heuristic_valid += 1167            if is_valid:168                syntax_valid += 1169            total += 1170 171            if total % 50 == 0:172                print(f"  {total}/{NUM_SAMPLES} done, "173                      f"syntax_valid={syntax_valid}, heuristic_valid={heuristic_valid}")174 175    heuristic_pct = (heuristic_valid / total * 100) if total > 0 else 0176    syntax_pct = (syntax_valid / total * 100) if total > 0 else 0177 178    print(f"\n{label}: heuristic={heuristic_pct:.1f}%, syntax={syntax_pct:.1f}% "179          f"({syntax_valid}/{total})")180 181    # Show examples of valid reconstructions182    valid_examples = [r for r in results if r["syntax_valid"]]183    print(f"\n--- Valid reconstructions ({len(valid_examples)} total) ---")184    for r in valid_examples[:5]:185        print(f"\nSample #{r['index']} ({r['num_nodes']} nodes, type_acc={r['type_accuracy']}):")186        print(f"  ORIGINAL:      {r['original_code'][:80].strip()}")187        print(f"  RECONSTRUCTED: {r['reconstructed_code'][:80].strip()}")188 189    # Show best type-accuracy samples that failed syntax190    failed_high_acc = sorted(191        [r for r in results if not r["syntax_valid"] and r["type_accuracy"] > 0.9],192        key=lambda x: -x["type_accuracy"],193    )194    if failed_high_acc:195        print(f"\n--- High accuracy but invalid syntax ({len(failed_high_acc)} total) ---")196        for r in failed_high_acc[:3]:197            print(f"\nSample #{r['index']} (type_acc={r['type_accuracy']}):")198            print(f"  ORIGINAL:      {r['original_code'][:80].strip()}")199            print(f"  RECONSTRUCTED: {r['reconstructed_code'][:80].strip()}")200            print(f"  PRED TYPES: {r['pred_types'][:12]}")201 202    output = {203        "label": label,204        "total": total,205        "heuristic_valid": heuristic_valid,206        "heuristic_pct": round(heuristic_pct, 2),207        "syntax_valid": syntax_valid,208        "syntax_pct": round(syntax_pct, 2),209        "valid_examples": valid_examples[:10],210        "failed_high_acc": [211            {k: v for k, v in r.items() if k != "pred_types"}212            for r in failed_high_acc[:10]213        ],214    }215 216    out_path = os.path.join(RESULTS_DIR, f"{label}_qualitative.json")217    with open(out_path, "w") as f:218        json.dump(output, f, indent=2)219    print(f"\nSaved to {out_path}")220    return output221 222 223def main() -> None:224    os.makedirs(RESULTS_DIR, exist_ok=True)225 226    configs = [227        {228            "model_path": os.path.join(RESULTS_DIR, "tf-gin-256-deep_decoder.pt"),229            "label": "tf-gin-256-deep",230            "hidden_dim": 256, "num_layers": 5,231            "decoder_conv_type": "GIN", "decoder_edge_mode": "teacher_forced",232        },233        {234            "model_path": os.path.join(RESULTS_DIR, "chain-gin-256_decoder.pt"),235            "label": "chain-gin-256",236            "hidden_dim": 256, "num_layers": 3,237            "decoder_conv_type": "GIN", "decoder_edge_mode": "chain",238        },239    ]240 241    all_results = []242    for cfg in configs:243        if not os.path.exists(cfg["model_path"]):244            print(f"Skipping {cfg['label']}: no model at {cfg['model_path']}")245            continue246        result = evaluate_model_full(**cfg)247        all_results.append(result)248 249    print("\n" + "=" * 60)250    print("FULL RECONSTRUCTION SUMMARY")251    print("=" * 60)252    for r in all_results:253        print(f"  {r['label']:25s}: heuristic={r['heuristic_pct']:5.1f}%  "254              f"syntax={r['syntax_pct']:5.1f}% ({r['syntax_valid']}/{r['total']})")255 256 257if __name__ == "__main__":258    main()259