File size: 3,040 Bytes
775f8ae
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
import ast
from pathlib import Path
import unittest


class DeferredCompressionCommitTests(unittest.TestCase):
    @classmethod
    def setUpClass(cls):
        validator_path = Path(__file__).parents[1] / "neurons" / "validator.py"
        validator_tree = ast.parse(validator_path.read_text(encoding="utf-8"))
        validator_class = next(
            node
            for node in validator_tree.body
            if isinstance(node, ast.ClassDef) and node.name == "Validator"
        )
        cls.methods = {
            node.name: node
            for node in validator_class.body
            if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
        }

    def test_compression_batches_are_scored_before_the_round_is_committed(self):
        method = self.methods["process_compression_miners"]
        scoring_loop = next(
            node
            for node in method.body
            if isinstance(node, ast.For)
            and any(
                isinstance(child, ast.Call)
                and isinstance(child.func, ast.Attribute)
                and child.func.attr == "score_compressions"
                for child in ast.walk(node)
            )
        )
        score_call = next(
            child
            for child in ast.walk(scoring_loop)
            if isinstance(child, ast.Call)
            and isinstance(child.func, ast.Attribute)
            and child.func.attr == "score_compressions"
        )
        commit_keyword = next(
            keyword
            for keyword in score_call.keywords
            if keyword.arg == "commit_scores"
        )
        self.assertIsInstance(commit_keyword.value, ast.Constant)
        self.assertIs(commit_keyword.value.value, False)

        loop_position = method.body.index(scoring_loop)
        later_calls = [
            child
            for statement in method.body[loop_position + 1 :]
            for child in ast.walk(statement)
            if isinstance(child, ast.Call)
            and isinstance(child.func, ast.Attribute)
        ]
        self.assertTrue(
            any(call.func.attr == "_commit_compression_scores" for call in later_calls)
        )

    def test_score_compressions_preserves_immediate_commit_by_default(self):
        method = self.methods["score_compressions"]
        commit_parameter = method.args.args[-1]
        self.assertEqual(commit_parameter.arg, "commit_scores")
        self.assertIs(method.args.defaults[-1].value, True)

        conditional_commit = next(
            node
            for node in ast.walk(method)
            if isinstance(node, ast.If)
            and isinstance(node.test, ast.Name)
            and node.test.id == "commit_scores"
        )
        self.assertTrue(
            any(
                isinstance(child, ast.Call)
                and isinstance(child.func, ast.Attribute)
                and child.func.attr == "_commit_compression_scores"
                for child in ast.walk(conditional_commit)
            )
        )


if __name__ == "__main__":
    unittest.main()