ujjwalpardeshi/pytorch-training-debugger
2
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 