Team Ai
Apppublic

ujjwalpardeshi/pytorch-training-debugger

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
code_templates.py249 linesDownload Raw Back to ml_training_debugger
1"""PyTorch code snippet templates for Task 6 code-level debugging.2 3Each template is a real, syntactically valid Python/PyTorch training script4with one injected bug.5"""6 7from __future__ import annotations8 9import ast10import io11import tokenize12from typing import Optional13 14import torch  # noqa: F40115 16# Bug variant templates: (buggy_code, correct_line_num, correct_replacement)17_TEMPLATES: dict[str, tuple[str, int, str]] = {18    "eval_mode": (19        """\20import torch21import torch.nn as nn22 23model = SimpleCNN()24model.eval()25optimizer = torch.optim.Adam(model.parameters(), lr=0.001)26criterion = nn.CrossEntropyLoss()27 28for epoch in range(100):29    for batch_x, batch_y in train_loader:30        optimizer.zero_grad()31        output = model(batch_x)32        loss = criterion(output, batch_y)33        loss.backward()34        optimizer.step()""",35        5,36        "model.train()",37    ),38    "detach_loss": (39        """\40import torch41import torch.nn as nn42 43model = SimpleCNN()44model.train()45optimizer = torch.optim.Adam(model.parameters(), lr=0.001)46criterion = nn.CrossEntropyLoss()47 48for epoch in range(100):49    for batch_x, batch_y in train_loader:50        optimizer.zero_grad()51        output = model(batch_x)52        loss = criterion(output, batch_y).detach()53        loss.backward()54        optimizer.step()""",55        14,56        "        loss = criterion(output, batch_y)",57    ),58    "zero_grad_missing": (59        """\60import torch61import torch.nn as nn62 63model = SimpleCNN()64model.train()65optimizer = torch.optim.Adam(model.parameters(), lr=0.001)66criterion = nn.CrossEntropyLoss()67 68for epoch in range(100):69    for batch_x, batch_y in train_loader:70        output = model(batch_x)71        loss = criterion(output, batch_y)72        loss.backward()73        optimizer.step()""",74        11,75        "        optimizer.zero_grad()",76    ),77    "inplace_relu": (78        """\79import torch80import torch.nn as nn81import torch.nn.functional as F82 83model = SimpleCNN()84model.train()85optimizer = torch.optim.Adam(model.parameters(), lr=0.001)86criterion = nn.CrossEntropyLoss()87 88for epoch in range(100):89    for batch_x, batch_y in train_loader:90        optimizer.zero_grad()91        output = model(batch_x)92        output = F.relu(output, inplace=True)93        loss = criterion(output, batch_y)94        loss.backward()95        optimizer.step()""",96        15,97        "        output = F.relu(output)",98    ),99}100 101# Semantic equivalence patterns per bug variant102_SEMANTIC_PATTERNS: dict[str, list[tuple[str, str]]] = {103    "eval_mode": [104        # (must_contain, must_not_contain)105        ("model.train()", "model.eval()"),106    ],107    "detach_loss": [108        ("criterion(", ".detach()"),109    ],110    "zero_grad_missing": [111        ("zero_grad()", ""),  # just needs zero_grad present112    ],113    "inplace_relu": [114        ("F.relu(", "inplace=True"),115    ],116}117 118 119def generate_code_snippet(bug_type: str, seed: int = 42) -> dict:120    """Generate a code snippet with the specified bug.121 122    Returns dict with keys: code, filename, line_count, imports, hint.123    """124    if bug_type not in _TEMPLATES:125        raise ValueError(f"Unknown bug_type: {bug_type}")126 127    code, _line, _replacement = _TEMPLATES[bug_type]128    lines = code.strip().split("\n")129    imports = [130        line for line in lines if line.startswith("import ") or line.startswith("from ")131    ]132 133    hint: Optional[str] = None134    if bug_type == "eval_mode":135        hint = "Check the model mode before the training loop."136    elif bug_type == "detach_loss":137        hint = "Examine how the loss is computed and used."138 139    return {140        "code": code,141        "filename": "train.py",142        "line_count": len(lines),143        "imports": imports,144        "hint": hint,145    }146 147 148def _normalize_code(s: str) -> str:149    """Strip whitespace and inline comments for comparison."""150    s = s.strip()151    # Remove inline comments152    result_lines: list[str] = []153    for line in s.split("\n"):154        # Remove trailing comment but preserve strings155        stripped = line.rstrip()156        result_lines.append(stripped)157    return "\n".join(result_lines)158 159 160def _tokenize_compare(original: str, replacement: str) -> bool:161    """Compare token streams ignoring whitespace and comments."""162 163    def get_tokens(code: str) -> list[tuple[int, str]]:164        try:165            tokens = list(tokenize.generate_tokens(io.StringIO(code).readline))166            # Filter out COMMENT, NL, NEWLINE, INDENT, DEDENT, ENCODING, ENDMARKER167            skip = {168                tokenize.COMMENT,169                tokenize.NL,170                tokenize.NEWLINE,171                tokenize.INDENT,172                tokenize.DEDENT,173                tokenize.ENCODING,174                tokenize.ENDMARKER,175            }176            return [(t.type, t.string) for t in tokens if t.type not in skip]177        except tokenize.TokenError:178            return []179 180    return get_tokens(original) == get_tokens(replacement)181 182 183def validate_fix(bug_type: str, line: int, replacement: str) -> bool:184    """Validate a code fix submission.185 186    Multi-strategy pipeline per spec Section 22:187    1. Normalize whitespace + strip comments188    2. Token-stream comparison189    3. Semantic equivalence patterns190    4. AST fallback191    """192    if bug_type not in _TEMPLATES:193        return False194 195    code, correct_line, correct_replacement = _TEMPLATES[bug_type]196    lines = code.strip().split("\n")197 198    # Check line number is valid199    if line < 1 or line > len(lines):200        return False201 202    # For zero_grad_missing, the fix is inserting a line, not replacing203    if bug_type == "zero_grad_missing":204        # Accept if the replacement contains zero_grad205        normalized = _normalize_code(replacement)206        if "zero_grad" in normalized:207            return True208        return False209 210    # Strategy 1: Normalize and compare211    norm_replacement = _normalize_code(replacement)212    norm_correct = _normalize_code(correct_replacement)213    if norm_replacement == norm_correct:214        return True215 216    # Strategy 2: Token-stream comparison217    if _tokenize_compare(correct_replacement, replacement):218        return True219 220    # Strategy 3: Semantic equivalence patterns221    patterns = _SEMANTIC_PATTERNS.get(bug_type, [])222    for must_contain, must_not_contain in patterns:223        if must_contain and must_contain in norm_replacement:224            if not must_not_contain or must_not_contain not in norm_replacement:225                return True226 227    # Strategy 4: AST fallback — verify buggy pattern absent228    try:229        # Replace the line in the full code and parse230        new_lines = lines.copy()231        new_lines[line - 1] = replacement.rstrip()232        new_code = "\n".join(new_lines)233        tree = ast.parse(new_code)234 235        # Check that the buggy pattern is absent236        ast.dump(tree)  # Validates AST is well-formed237        if bug_type == "eval_mode" and "eval" not in replacement.lower():238            if "train" in replacement.lower():239                return True240        if bug_type == "detach_loss" and "detach" not in replacement.lower():241            return True242        if bug_type == "inplace_relu" and "inplace" not in replacement.lower():243            if "relu" in replacement.lower():244                return True245    except SyntaxError:246        pass247 248    return False249