File size: 5,155 Bytes
62e0b07
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
import copy
import importlib.util
import math
from pathlib import Path
import unittest

spec = importlib.util.spec_from_file_location("plot_results", Path(__file__).parents[1] / "plot_results.py")
plot = importlib.util.module_from_spec(spec)
spec.loader.exec_module(plot)


def fixture():
    summary = {"episodes": [{"condition": c, "task": f"task-{i}", "has_valid_grade": True,
                             "clock_valid": True, "invalid_grade_count": 0, "status": "completed"}
                            for c in plot.CONDITIONS for i in range(5)]}
    rows = []
    for axis in plot.AXES:
        for condition in plot.CONDITIONS:
            scores = [0, .2, .7, .5] if condition == "single" else [0, .7, .8, .6]
            for time, score in zip([0, 30, 60, 120], scores):
                rows.append(dict(axis=axis, time_seconds=time, condition=condition,
                                 mean_hidden_test_fraction=score, tasks_total=5,
                                 tasks_with_valid_grade=5, episodes_configured=5, clock_valid=True))
    return summary, rows


class PlotTests(unittest.TestCase):
    def test_pending_does_not_fabricate_zero_curves(self):
        figure = plot.build_plotly(None, [])
        self.assertFalse(figure["layout"]["meta"]["has_results"])
        self.assertEqual(len(figure["data"]), 3)
        self.assertTrue(all(trace["y"] == [None] for trace in figure["data"]))
        self.assertIsNone(figure["layout"]["meta"]["threshold_comparison"])

    def test_regression_and_zero_time_are_not_smoothed_or_shifted(self):
        summary, rows = fixture()
        data = plot.figure_data(summary, rows)
        self.assertEqual([r["score"] for r in data["groups"]["single"]["rows"]], [.2, .7, .5])
        self.assertEqual([r["seconds"] for r in data["groups"]["single"]["rows"]], [30, 60, 120])
        figure = plot.build_plotly(summary, rows)
        self.assertEqual(figure["data"][0]["line"]["shape"], "hv")

    def test_speedup_requires_complete_coverage_and_observed_crossing(self):
        summary, rows = fixture()
        comparison = plot.threshold_comparison(plot.figure_data(summary, rows))
        self.assertEqual(comparison["ratio"], 2)
        self.assertEqual(comparison["single"], 60)
        annotation = next(a for a in plot.build_plotly(summary, rows)["layout"]["annotations"] if a.get("axref") == "x")
        self.assertEqual(annotation["x"], math.log10(60))
        self.assertEqual(annotation["ax"], math.log10(30))
        summary["episodes"][0]["has_valid_grade"] = False
        self.assertIsNone(plot.threshold_comparison(plot.figure_data(summary, rows)))

    def test_unreached_threshold_is_not_extrapolated(self):
        summary, rows = fixture()
        for row in rows:
            row["mean_hidden_test_fraction"] *= .5
        self.assertIsNone(plot.threshold_comparison(plot.figure_data(summary, rows)))

    def test_bad_clock_suppresses_only_normalized_condition(self):
        summary, rows = fixture()
        summary["episodes"][0]["clock_valid"] = False
        normalized = plot.figure_data(summary, rows)
        wall = plot.figure_data(summary, rows, "wall_seconds")
        self.assertFalse(normalized["groups"]["single"]["available"])
        self.assertIn("unavailable", normalized["status"])
        self.assertIn("clock unavailable", plot.legend_suffix(normalized["groups"]["single"]))
        self.assertTrue(wall["groups"]["single"]["available"])
        for row in rows:
            row["clock_valid"] = "False"
        self.assertFalse(plot.figure_data(summary, rows)["available"])

    def test_missing_results_are_provisional_and_cannot_claim_speedup(self):
        summary, rows = fixture()
        summary["episodes"][0].update(status="failed", invalid_grade_count=1)
        data = plot.figure_data(summary, rows)
        self.assertIn("Provisional", data["status"])
        self.assertIsNone(plot.threshold_comparison(data))
        self.assertIn("zero before any valid grade", data["detail"])

    def test_wrong_denominator_or_nonfinite_data_is_rejected(self):
        summary, rows = fixture()
        for key, value in [("tasks_total", 4), ("mean_hidden_test_fraction", float("nan")),
                           ("time_seconds", -1), ("tasks_with_valid_grade", 6)]:
            invalid = copy.deepcopy(rows)
            invalid[0][key] = value
            with self.subTest(key=key), self.assertRaises(ValueError):
                plot.figure_data(summary, invalid)

    def test_plotly_schema_accepts_complete_and_pending_figures(self):
        import plotly.graph_objects as go
        summary, rows = fixture()
        for axis in plot.AXES:
            go.Figure(plot.build_plotly(summary, rows, axis))
            go.Figure(plot.build_plotly(None, [], axis))

    def test_single_positive_checkpoint_is_visible(self):
        summary, rows = fixture()
        rows = [r for r in rows if r["time_seconds"] == 30]
        figure = plot.build_plotly(summary, rows)
        self.assertEqual(figure["data"][0]["marker"]["size"], [5])
        self.assertEqual(figure["data"][0]["mode"], "lines+markers")


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