|
| 1 | +"""Unit tests for the experiments metadata model and store (Layer 2, PR 1). |
| 2 | +
|
| 3 | +Exercises :mod:`visdom.experiments` against a real ``JSONStore`` over a |
| 4 | +temporary directory, so no running visdom server is needed. These cover the |
| 5 | +model schemas (validation and round-trip (de)serialisation) and the |
| 6 | +``ExperimentStore`` CRUD/list operations, including that experiment metadata |
| 7 | +persists to disk and coexists with ordinary environment window data. |
| 8 | +""" |
| 9 | + |
| 10 | +import tempfile |
| 11 | +import unittest |
| 12 | + |
| 13 | +from visdom.data_model import JSONStore |
| 14 | +from visdom.experiments import ( |
| 15 | + Experiment, |
| 16 | + ExperimentStore, |
| 17 | + Metric, |
| 18 | + Param, |
| 19 | + Tag, |
| 20 | + STATUS_FAILED, |
| 21 | + STATUS_FINISHED, |
| 22 | + STATUS_RUNNING, |
| 23 | +) |
| 24 | + |
| 25 | + |
| 26 | +class TestModels(unittest.TestCase): |
| 27 | + """The Experiment/Param/Metric/Tag schemas validate and round-trip.""" |
| 28 | + |
| 29 | + def test_param_infers_dtype(self): |
| 30 | + """Param.dtype is inferred from the value when not given explicitly.""" |
| 31 | + self.assertEqual(Param("lr", 0.01).dtype, "float") |
| 32 | + self.assertEqual(Param("epochs", 10).dtype, "int") |
| 33 | + self.assertEqual(Param("name", "resnet").dtype, "str") |
| 34 | + # bool must not be mislabelled as int (bool subclasses int). |
| 35 | + self.assertEqual(Param("amp", True).dtype, "bool") |
| 36 | + |
| 37 | + def test_param_respects_explicit_dtype(self): |
| 38 | + """An explicit dtype is kept rather than being overwritten.""" |
| 39 | + self.assertEqual(Param("k", "1", dtype="int").dtype, "int") |
| 40 | + |
| 41 | + def test_experiment_defaults(self): |
| 42 | + """A new experiment defaults name to env_id and starts running.""" |
| 43 | + exp = Experiment(env_id="main") |
| 44 | + self.assertEqual(exp.name, "main") |
| 45 | + self.assertEqual(exp.status, STATUS_RUNNING) |
| 46 | + self.assertIsNone(exp.finished_at) |
| 47 | + |
| 48 | + def test_experiment_rejects_bad_status(self): |
| 49 | + """Constructing with an unknown status raises ValueError.""" |
| 50 | + with self.assertRaises(ValueError): |
| 51 | + Experiment(env_id="main", status="bogus") |
| 52 | + |
| 53 | + def test_set_param_replaces_by_key(self): |
| 54 | + """set_param overwrites an existing param of the same name in place.""" |
| 55 | + exp = Experiment(env_id="main") |
| 56 | + exp.set_param("lr", 0.1) |
| 57 | + exp.set_param("lr", 0.01) |
| 58 | + self.assertEqual(len(exp.params), 1) |
| 59 | + self.assertEqual(exp.get_param("lr").value, 0.01) |
| 60 | + |
| 61 | + def test_set_tag_replaces_by_key(self): |
| 62 | + """set_tag overwrites an existing tag of the same name in place.""" |
| 63 | + exp = Experiment(env_id="main") |
| 64 | + exp.set_tag("dataset", "mnist") |
| 65 | + exp.set_tag("dataset", "cifar") |
| 66 | + self.assertEqual(len(exp.tags), 1) |
| 67 | + self.assertEqual(exp.tags[0].value, "cifar") |
| 68 | + |
| 69 | + def test_add_metric_appends_time_series(self): |
| 70 | + """add_metric appends (never replaces); latest_metric reads the last.""" |
| 71 | + exp = Experiment(env_id="main") |
| 72 | + exp.add_metric("acc", 0.5, step=1) |
| 73 | + exp.add_metric("acc", 0.9, step=2) |
| 74 | + self.assertEqual(len(exp.metrics), 2) |
| 75 | + self.assertEqual(exp.latest_metric("acc").value, 0.9) |
| 76 | + self.assertIsNone(exp.latest_metric("missing")) |
| 77 | + |
| 78 | + def test_finish_sets_terminal_state(self): |
| 79 | + """finish stamps finished_at and only accepts terminal statuses.""" |
| 80 | + exp = Experiment(env_id="main") |
| 81 | + exp.finish() |
| 82 | + self.assertEqual(exp.status, STATUS_FINISHED) |
| 83 | + self.assertIsNotNone(exp.finished_at) |
| 84 | + exp2 = Experiment(env_id="two") |
| 85 | + exp2.finish(STATUS_FAILED) |
| 86 | + self.assertEqual(exp2.status, STATUS_FAILED) |
| 87 | + with self.assertRaises(ValueError): |
| 88 | + Experiment(env_id="three").finish(STATUS_RUNNING) |
| 89 | + |
| 90 | + def test_round_trip_serialisation(self): |
| 91 | + """to_dict/from_dict preserve every field of a populated experiment.""" |
| 92 | + exp = Experiment(env_id="main", description="run 1") |
| 93 | + exp.set_param("lr", 0.01) |
| 94 | + exp.set_tag("dataset", "mnist") |
| 95 | + exp.add_metric("acc", 0.9, step=3) |
| 96 | + exp.finish() |
| 97 | + rebuilt = Experiment.from_dict(exp.to_dict()) |
| 98 | + self.assertEqual(rebuilt.to_dict(), exp.to_dict()) |
| 99 | + self.assertIsInstance(rebuilt.params[0], Param) |
| 100 | + self.assertIsInstance(rebuilt.metrics[0], Metric) |
| 101 | + self.assertIsInstance(rebuilt.tags[0], Tag) |
| 102 | + |
| 103 | + |
| 104 | +class TestExperimentStore(unittest.TestCase): |
| 105 | + """ExperimentStore CRUD/list operations over a real JSONStore backend.""" |
| 106 | + |
| 107 | + def setUp(self): |
| 108 | + """Give each test a fresh temp env_path and a store over it.""" |
| 109 | + self._tmp = tempfile.TemporaryDirectory() |
| 110 | + self.env_path = self._tmp.name |
| 111 | + self.backend = JSONStore(self.env_path) |
| 112 | + self.store = ExperimentStore(self.backend) |
| 113 | + |
| 114 | + def tearDown(self): |
| 115 | + self._tmp.cleanup() |
| 116 | + |
| 117 | + def test_log_experiment_persists_to_disk(self): |
| 118 | + """A logged experiment is readable by a brand-new store over the dir.""" |
| 119 | + self.store.log_experiment( |
| 120 | + "main", params={"lr": 0.01, "epochs": 10}, tags={"dataset": "mnist"} |
| 121 | + ) |
| 122 | + # A fresh store instance proves the data survives on disk, not just in |
| 123 | + # the object we wrote through. |
| 124 | + reopened = ExperimentStore(JSONStore(self.env_path)) |
| 125 | + exp = reopened.get_experiment("main") |
| 126 | + self.assertIsNotNone(exp) |
| 127 | + self.assertEqual(exp.get_param("lr").value, 0.01) |
| 128 | + self.assertEqual(exp.get_param("epochs").dtype, "int") |
| 129 | + self.assertEqual(exp.tags[0].value, "mnist") |
| 130 | + |
| 131 | + def test_get_missing_experiment_returns_none(self): |
| 132 | + """Envs with no experiment blob yield None (feature is opt-in).""" |
| 133 | + self.assertIsNone(self.store.get_experiment("never_logged")) |
| 134 | + |
| 135 | + def test_log_experiment_updates_in_place(self): |
| 136 | + """Re-logging merges params and keeps prior metrics rather than resetting.""" |
| 137 | + self.store.log_experiment("main", params={"lr": 0.1}) |
| 138 | + self.store.log_metric("main", "acc", 0.7) |
| 139 | + exp = self.store.log_experiment( |
| 140 | + "main", description="updated", params={"lr": 0.01, "wd": 0.0} |
| 141 | + ) |
| 142 | + self.assertEqual(exp.description, "updated") |
| 143 | + self.assertEqual(exp.get_param("lr").value, 0.01) |
| 144 | + self.assertEqual(exp.get_param("wd").value, 0.0) |
| 145 | + # The metric logged before the update is still present. |
| 146 | + self.assertEqual(len(exp.metrics), 1) |
| 147 | + |
| 148 | + def test_log_metric_auto_creates_experiment(self): |
| 149 | + """Logging a metric for an env with no experiment creates one.""" |
| 150 | + self.store.log_metric("main", "loss", 1.5, step=0) |
| 151 | + exp = self.store.get_experiment("main") |
| 152 | + self.assertIsNotNone(exp) |
| 153 | + self.assertEqual(exp.latest_metric("loss").value, 1.5) |
| 154 | + |
| 155 | + def test_finish_experiment(self): |
| 156 | + """finish_experiment persists a terminal status.""" |
| 157 | + self.store.log_experiment("main") |
| 158 | + self.store.finish_experiment("main", STATUS_FAILED) |
| 159 | + self.assertEqual(self.store.get_experiment("main").status, STATUS_FAILED) |
| 160 | + |
| 161 | + def test_finish_missing_experiment_raises(self): |
| 162 | + """Finishing an env that never logged an experiment raises KeyError.""" |
| 163 | + with self.assertRaises(KeyError): |
| 164 | + self.store.finish_experiment("nope") |
| 165 | + |
| 166 | + def test_list_experiments(self): |
| 167 | + """list_experiments returns only envs that actually have a blob.""" |
| 168 | + self.store.log_experiment("a") |
| 169 | + self.store.log_experiment("b") |
| 170 | + # An ordinary env with window data but no experiment must be skipped. |
| 171 | + self.backend.save_env("plain", {"jsons": {"w": {"id": "w"}}, "reload": {}}) |
| 172 | + listed = sorted(exp.env_id for exp in self.store.list_experiments()) |
| 173 | + self.assertEqual(listed, ["a", "b"]) |
| 174 | + |
| 175 | + def test_delete_experiment_keeps_env(self): |
| 176 | + """delete_experiment drops the blob but leaves the environment intact.""" |
| 177 | + self.backend.save_env("main", {"jsons": {"w": {"id": "w"}}, "reload": {}}) |
| 178 | + self.store.log_experiment("main", params={"lr": 0.01}) |
| 179 | + self.assertTrue(self.store.delete_experiment("main")) |
| 180 | + self.assertIsNone(self.store.get_experiment("main")) |
| 181 | + # The window data the env had before is untouched. |
| 182 | + env = self.backend.load_env("main") |
| 183 | + self.assertIn("w", env["jsons"]) |
| 184 | + self.assertNotIn("experiment", env) |
| 185 | + |
| 186 | + def test_delete_missing_experiment_returns_false(self): |
| 187 | + """Deleting from an env with no experiment reports False.""" |
| 188 | + self.assertFalse(self.store.delete_experiment("nope")) |
| 189 | + |
| 190 | + def test_experiment_coexists_with_window_data(self): |
| 191 | + """Logging onto an env with windows preserves those windows on disk.""" |
| 192 | + self.backend.save_env( |
| 193 | + "main", {"jsons": {"win_0": {"id": "win_0"}}, "reload": {"foo": 1}} |
| 194 | + ) |
| 195 | + self.store.log_experiment("main", params={"lr": 0.01}) |
| 196 | + env = self.backend.load_env("main") |
| 197 | + self.assertIn("win_0", env["jsons"]) |
| 198 | + self.assertEqual(env["reload"], {"foo": 1}) |
| 199 | + self.assertEqual(env["experiment"]["params"][0]["value"], 0.01) |
| 200 | + |
| 201 | + |
| 202 | +class TestExperimentStoreNoPersistence(unittest.TestCase): |
| 203 | + """With persistence disabled the store degrades gracefully (no crashes).""" |
| 204 | + |
| 205 | + def setUp(self): |
| 206 | + self.store = ExperimentStore(JSONStore(None)) |
| 207 | + |
| 208 | + def test_log_returns_experiment_but_read_is_empty(self): |
| 209 | + """log_* still returns a valid object though nothing is persisted.""" |
| 210 | + exp = self.store.log_experiment("main", params={"lr": 0.01}) |
| 211 | + self.assertIsInstance(exp, Experiment) |
| 212 | + # Nothing is stored, so a subsequent read finds no experiment. |
| 213 | + self.assertIsNone(self.store.get_experiment("main")) |
| 214 | + self.assertEqual(self.store.list_experiments(), []) |
| 215 | + |
| 216 | + |
| 217 | +if __name__ == "__main__": |
| 218 | + unittest.main() |
0 commit comments