ujjwalpardeshi/pytorch-training-debugger
2
1"""Test code bug generation and fix validation."""2 3from __future__ import annotations4 5import pytest6 7from ml_training_debugger.code_templates import generate_code_snippet, validate_fix8 9 10class TestGenerateCodeSnippet:11 def test_eval_mode(self):12 snippet = generate_code_snippet("eval_mode")13 assert "model.eval()" in snippet["code"]14 assert snippet["filename"] == "train.py"15 assert snippet["line_count"] > 016 assert len(snippet["imports"]) > 017 18 def test_detach_loss(self):19 snippet = generate_code_snippet("detach_loss")20 assert ".detach()" in snippet["code"]21 22 def test_zero_grad_missing(self):23 snippet = generate_code_snippet("zero_grad_missing")24 assert "zero_grad" not in snippet["code"]25 26 def test_inplace_relu(self):27 snippet = generate_code_snippet("inplace_relu")28 assert "inplace=True" in snippet["code"]29 30 def test_unknown_bug_raises(self):31 with pytest.raises(ValueError):32 generate_code_snippet("nonexistent_bug")33 34 35class TestValidateFix:36 def test_eval_mode_correct_fix(self):37 assert validate_fix("eval_mode", 5, "model.train()")38 39 def test_eval_mode_with_whitespace(self):40 assert validate_fix("eval_mode", 5, " model.train() ")41 42 def test_eval_mode_wrong_fix(self):43 assert not validate_fix("eval_mode", 5, "pass")44 45 def test_detach_loss_correct_fix(self):46 assert validate_fix(47 "detach_loss", 14, " loss = criterion(output, batch_y)"48 )49 50 def test_detach_loss_with_trailing_spaces(self):51 assert validate_fix(52 "detach_loss", 14, " loss = criterion(output, batch_y) "53 )54 55 def test_zero_grad_correct_fix(self):56 assert validate_fix("zero_grad_missing", 11, " optimizer.zero_grad()")57 58 def test_inplace_relu_correct_fix(self):59 assert validate_fix("inplace_relu", 15, " output = F.relu(output)")60 61 def test_wrong_line_number(self):62 assert not validate_fix("eval_mode", 999, "model.train()")63 64 def test_unknown_bug_type(self):65 assert not validate_fix("nonexistent", 1, "pass")66 