Team Ai
Apppublic

ujjwalpardeshi/pytorch-training-debugger

sourceHugging Faceupdated 6mo agoView on Hugging Face
2likes
test_code_templates.py66 linesDownload Raw Back to tests
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