mllogs 0.2.0__tar.gz → 0.3.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: mllogs
3
- Version: 0.2.0
3
+ Version: 0.3.0
4
4
  Summary: Local experiment logging for machine learning
5
5
  Author: Finn Walsh
6
6
  License-Expression: MIT
@@ -45,15 +45,15 @@ client.list_runs()
45
45
  Retrieve or delete a saved run:
46
46
 
47
47
  ```python
48
- client.get_run(id)
49
- client.delete_run(id)
48
+ client.get_run(run.id)
49
+ client.delete_run(run.id)
50
50
  ```
51
51
 
52
52
  ## Features
53
53
 
54
54
  - Start and end experiment runs
55
55
  - Log parameters, metrics, and tags
56
- - Automatically persist completed runs to local storage
56
+ - Persist completed runs to a local SQLite database
57
57
  - Load, list, and delete saved runs
58
58
 
59
59
  ## Development
@@ -32,15 +32,15 @@ client.list_runs()
32
32
  Retrieve or delete a saved run:
33
33
 
34
34
  ```python
35
- client.get_run(id)
36
- client.delete_run(id)
35
+ client.get_run(run.id)
36
+ client.delete_run(run.id)
37
37
  ```
38
38
 
39
39
  ## Features
40
40
 
41
41
  - Start and end experiment runs
42
42
  - Log parameters, metrics, and tags
43
- - Automatically persist completed runs to local storage
43
+ - Persist completed runs to a local SQLite database
44
44
  - Load, list, and delete saved runs
45
45
 
46
46
  ## Development
@@ -7,7 +7,7 @@ where = ["src"]
7
7
 
8
8
  [project]
9
9
  name = "mllogs"
10
- version = "0.2.0"
10
+ version = "0.3.0"
11
11
  description = "Local experiment logging for machine learning"
12
12
  readme = "README.md"
13
13
  requires-python = ">=3.11"
@@ -0,0 +1,3 @@
1
+ from .client import MLLogsClient
2
+
3
+ __all__ = ["MLLogsClient"]
@@ -3,13 +3,20 @@ from datetime import datetime, UTC
3
3
  from uuid import uuid4
4
4
 
5
5
  from .run import Run, RunStatus
6
- from .storage import LocalFileStore
6
+ from .storage import Storage
7
+
8
+ from .storage.db import SQLiteStore
7
9
 
8
10
 
9
11
  class MLLogsClient:
10
- def __init__(self, file_store: LocalFileStore | None = None):
12
+ def __init__(self, storage: Storage | None = None):
11
13
  self._active_run: Run | None = None
12
- self._file_store = file_store if file_store is not None else LocalFileStore()
14
+
15
+ if storage is None:
16
+ db = SQLiteStore()
17
+ self._storage = Storage(db=db)
18
+ else:
19
+ self._storage = storage
13
20
 
14
21
 
15
22
  def _require_active_run(self) -> Run:
@@ -19,6 +26,10 @@ class MLLogsClient:
19
26
  return self._active_run
20
27
 
21
28
 
29
+ # ------------------
30
+ # --- Active Run ---
31
+ # ------------------
32
+
22
33
  def start_run(
23
34
  self,
24
35
  name: str | None = None,
@@ -47,22 +58,6 @@ class MLLogsClient:
47
58
  status = RunStatus.RUNNING,
48
59
  started_at=started_at,
49
60
  )
50
- # END
51
-
52
-
53
- def log_param(self, key: str, value: Any) -> None:
54
- run = self._require_active_run()
55
- run.params[key] = value
56
-
57
-
58
- def log_metric(self, key: str, value: float) -> None:
59
- run = self._require_active_run()
60
- run.metrics[key] = value
61
-
62
-
63
- def set_tag(self, key: str, value: str) -> None:
64
- run = self._require_active_run()
65
- run.tags[key] = value
66
61
 
67
62
 
68
63
  def end_run(self) -> Run:
@@ -77,21 +72,40 @@ class MLLogsClient:
77
72
  run.ended_at = datetime.now(UTC)
78
73
  run.status = RunStatus.COMPLETE
79
74
 
80
- self._file_store.save_run(run)
75
+ self._storage.save_run(run)
81
76
 
82
77
  # clear run
83
78
  self._active_run = None
84
79
 
85
80
  return run
81
+
82
+
83
+ def log_param(self, key: str, value: Any) -> None:
84
+ run = self._require_active_run()
85
+ run.params[key] = value
86
+
87
+
88
+ def log_metric(self, key: str, value: float) -> None:
89
+ run = self._require_active_run()
90
+ run.metrics[key] = value
91
+
92
+
93
+ def set_tag(self, key: str, value: str) -> None:
94
+ run = self._require_active_run()
95
+ run.tags[key] = value
96
+
86
97
 
98
+ # ----------------------
99
+ # --- Persisted runs ---
100
+ # ----------------------
87
101
 
88
102
  def get_run(self, run_id: str | None = None) -> Run | None:
89
- return self._file_store.load_run(run_id=run_id)
103
+ return self._storage.load_run(run_id=run_id)
90
104
 
91
105
 
92
106
  def list_runs(self, limit: int | None = None) -> list[Run]:
93
- return self._file_store.list_runs(limit=limit)
107
+ return self._storage.list_runs(limit=limit)
94
108
 
95
109
 
96
110
  def delete_run(self, run_id: str) -> None:
97
- self._file_store.delete_run(run_id=run_id)
111
+ self._storage.delete_run(run_id=run_id)
@@ -0,0 +1,24 @@
1
+ from datetime import datetime
2
+ from dataclasses import dataclass, field
3
+ from enum import Enum
4
+
5
+
6
+ class RunStatus(Enum):
7
+ RUNNING = "running"
8
+ COMPLETE = "complete"
9
+ FAILED = "failed"
10
+
11
+
12
+ @dataclass
13
+ class Run:
14
+ id: str
15
+ started_at: datetime
16
+ status: RunStatus
17
+
18
+ name: str | None = None
19
+ run_type: str | None = None
20
+ ended_at: datetime | None = None
21
+
22
+ params: dict[str, str | int | float | bool] = field(default_factory=dict)
23
+ metrics: dict[str, float] = field(default_factory=dict)
24
+ tags: dict[str, str] = field(default_factory=dict)
@@ -0,0 +1,5 @@
1
+ from .storage import Storage
2
+
3
+ __all__ = [
4
+ "Storage",
5
+ ]
@@ -0,0 +1,7 @@
1
+ from .base import DBStore
2
+ from .sqlite import SQLiteStore
3
+
4
+ __all__ = [
5
+ "DBStore",
6
+ "SQLiteStore",
7
+ ]
@@ -0,0 +1,21 @@
1
+ from abc import ABC, abstractmethod
2
+
3
+ from mllogs.run import Run
4
+
5
+
6
+ class DBStore(ABC):
7
+ @abstractmethod
8
+ def save_run(self, run: Run) -> None:
9
+ ...
10
+
11
+ @abstractmethod
12
+ def load_run(self, run_id: str | None = None) -> Run | None:
13
+ ...
14
+
15
+ @abstractmethod
16
+ def list_runs(self, limit: int | None = None) -> list[Run]:
17
+ ...
18
+
19
+ @abstractmethod
20
+ def delete_run(self, run_id: str) -> None:
21
+ ...
@@ -0,0 +1,227 @@
1
+ from pathlib import Path
2
+ import sqlite3
3
+ import json
4
+ from datetime import datetime
5
+
6
+ from .base import DBStore
7
+ from mllogs.run import Run, RunStatus
8
+
9
+
10
+ class SQLiteStore(DBStore):
11
+ def __init__(self, db_path: str | Path = ".mllogs/mllogs.db") -> None:
12
+ self._db_path = Path(db_path)
13
+
14
+ self._db_path.parent.mkdir(parents=True, exist_ok=True)
15
+
16
+ self._connection = sqlite3.connect(self._db_path)
17
+ self._connection.row_factory = sqlite3.Row
18
+ self._connection.execute("PRAGMA foreign_keys = ON")
19
+
20
+ self._initialize_schema()
21
+
22
+
23
+ def _initialize_schema(self) -> None:
24
+ self._connection.executescript(
25
+ """
26
+ CREATE TABLE IF NOT EXISTS runs (
27
+ id TEXT PRIMARY KEY,
28
+ started_at TEXT NOT NULL,
29
+ status TEXT NOT NULL,
30
+ name TEXT,
31
+ run_type TEXT,
32
+ ended_at TEXT
33
+ );
34
+
35
+ CREATE TABLE IF NOT EXISTS params (
36
+ run_id TEXT NOT NULL,
37
+ key TEXT NOT NULL,
38
+ value TEXT NOT NULL,
39
+
40
+ PRIMARY KEY (run_id, key),
41
+ FOREIGN KEY (run_id)
42
+ REFERENCES runs(id)
43
+ ON DELETE CASCADE
44
+
45
+ );
46
+
47
+ CREATE TABLE IF NOT EXISTS metrics (
48
+ run_id TEXT NOT NULL,
49
+ key TEXT NOT NULL,
50
+ value REAL NOT NULL,
51
+
52
+ PRIMARY KEY (run_id, key),
53
+ FOREIGN KEY (run_id)
54
+ REFERENCES runs(id)
55
+ ON DELETE CASCADE
56
+ );
57
+
58
+ CREATE TABLE IF NOT EXISTS tags (
59
+ run_id TEXT NOT NULL,
60
+ key TEXT NOT NULL,
61
+ value TEXT NOT NULL,
62
+
63
+ PRIMARY KEY (run_id, key),
64
+ FOREIGN KEY (run_id)
65
+ REFERENCES runs(id)
66
+ ON DELETE CASCADE
67
+ );
68
+ """
69
+ )
70
+
71
+
72
+ def save_run(self, run: Run) -> None:
73
+ with self._connection:
74
+ self._connection.execute(
75
+ """
76
+ INSERT INTO runs (
77
+ id,
78
+ started_at,
79
+ status,
80
+ name,
81
+ run_type,
82
+ ended_at
83
+ )
84
+ VALUES (?, ?, ?, ?, ?, ?)
85
+ """,
86
+ (
87
+ run.id,
88
+ run.started_at.isoformat(),
89
+ run.status.value,
90
+ run.name,
91
+ run.run_type,
92
+ run.ended_at.isoformat() if run.ended_at else None,
93
+ ),
94
+ )
95
+
96
+ self._connection.executemany(
97
+ """
98
+ INSERT INTO params (run_id, key, value)
99
+ VALUES (?, ?, ?)
100
+ """,
101
+ [
102
+ (run.id, key, json.dumps(value))
103
+ for key, value in run.params.items()
104
+ ],
105
+ )
106
+
107
+ self._connection.executemany(
108
+ """
109
+ INSERT INTO metrics (run_id, key, value)
110
+ VALUES (?, ?, ?)
111
+ """,
112
+ [
113
+ (run.id, key, value)
114
+ for key, value in run.metrics.items()
115
+ ],
116
+ )
117
+
118
+ self._connection.executemany(
119
+ """
120
+ INSERT INTO tags (run_id, key, value)
121
+ VALUES (?, ?, ?)
122
+ """,
123
+ [
124
+ (run.id, key, str(value))
125
+ for key, value in run.tags.items()
126
+ ],
127
+ )
128
+
129
+
130
+ def load_run(self, run_id: str | None = None) -> Run | None:
131
+ if run_id is None:
132
+ run_row = self._connection.execute(
133
+ """
134
+ SELECT *
135
+ FROM runs
136
+ ORDER BY started_at DESC
137
+ LIMIT 1
138
+ """
139
+ ).fetchone()
140
+ else:
141
+ run_row = self._connection.execute(
142
+ "SELECT * FROM runs WHERE id = ?",
143
+ (run_id,),
144
+ ).fetchone()
145
+
146
+ if run_row is None:
147
+ return None
148
+
149
+ run_id = run_row["id"]
150
+
151
+ param_rows = self._connection.execute(
152
+ "SELECT key, value FROM params WHERE run_id = ?",
153
+ (run_id,),
154
+ ).fetchall()
155
+
156
+ metric_rows = self._connection.execute(
157
+ "SELECT key, value FROM metrics WHERE run_id = ?",
158
+ (run_id,),
159
+ ).fetchall()
160
+
161
+ tag_rows = self._connection.execute(
162
+ "SELECT key, value FROM tags WHERE run_id = ?",
163
+ (run_id,),
164
+ ).fetchall()
165
+
166
+ return Run(
167
+ id=run_row["id"],
168
+ started_at=datetime.fromisoformat(run_row["started_at"]),
169
+ status=RunStatus(run_row["status"]),
170
+ name=run_row["name"],
171
+ run_type=run_row["run_type"],
172
+ ended_at=(
173
+ datetime.fromisoformat(run_row["ended_at"])
174
+ if run_row["ended_at"]
175
+ else None
176
+ ),
177
+ params={
178
+ row["key"]: json.loads(row["value"])
179
+ for row in param_rows
180
+ },
181
+ metrics={
182
+ row["key"]: row["value"]
183
+ for row in metric_rows
184
+ },
185
+ tags={
186
+ row["key"]: row["value"]
187
+ for row in tag_rows
188
+ },
189
+ )
190
+
191
+
192
+ def list_runs(self, limit: int | None = None) -> list[Run]:
193
+ if limit is not None and limit <= 0:
194
+ raise ValueError("limit must be greater than 0")
195
+
196
+ query = """
197
+ SELECT id
198
+ FROM runs
199
+ ORDER BY started_at DESC
200
+ """
201
+
202
+ params = ()
203
+
204
+ if limit is not None:
205
+ query += " LIMIT ?"
206
+ params = (limit, )
207
+
208
+ rows = self._connection.execute(
209
+ query,
210
+ params,
211
+ ).fetchall()
212
+
213
+ return [
214
+ self.load_run(row["id"])
215
+ for row in rows
216
+ ]
217
+
218
+
219
+ def delete_run(self, run_id: str) -> None:
220
+ with self._connection:
221
+ cursor = self._connection.execute(
222
+ "DELETE FROM runs WHERE id = ?",
223
+ (run_id,),
224
+ )
225
+
226
+ if cursor.rowcount == 0:
227
+ raise KeyError(f"Run not found: {run_id}")
@@ -0,0 +1,23 @@
1
+ from .db import DBStore
2
+
3
+ from mllogs.run import Run
4
+
5
+
6
+ class Storage:
7
+ def __init__(
8
+ self,
9
+ db: DBStore,
10
+ ) -> None:
11
+ self._db = db
12
+
13
+ def save_run(self, run: Run) -> None:
14
+ self._db.save_run(run)
15
+
16
+ def load_run(self, run_id: str | None = None) -> Run | None:
17
+ return self._db.load_run(run_id)
18
+
19
+ def list_runs(self, limit: int | None = None) -> list[Run]:
20
+ return self._db.list_runs(limit)
21
+
22
+ def delete_run(self, run_id: str) -> None:
23
+ self._db.delete_run(run_id)
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: mllogs
3
- Version: 0.2.0
3
+ Version: 0.3.0
4
4
  Summary: Local experiment logging for machine learning
5
5
  Author: Finn Walsh
6
6
  License-Expression: MIT
@@ -45,15 +45,15 @@ client.list_runs()
45
45
  Retrieve or delete a saved run:
46
46
 
47
47
  ```python
48
- client.get_run(id)
49
- client.delete_run(id)
48
+ client.get_run(run.id)
49
+ client.delete_run(run.id)
50
50
  ```
51
51
 
52
52
  ## Features
53
53
 
54
54
  - Start and end experiment runs
55
55
  - Log parameters, metrics, and tags
56
- - Automatically persist completed runs to local storage
56
+ - Persist completed runs to a local SQLite database
57
57
  - Load, list, and delete saved runs
58
58
 
59
59
  ## Development
@@ -4,12 +4,15 @@ pyproject.toml
4
4
  src/mllogs/__init__.py
5
5
  src/mllogs/client.py
6
6
  src/mllogs/run.py
7
- src/mllogs/storage.py
8
7
  src/mllogs.egg-info/PKG-INFO
9
8
  src/mllogs.egg-info/SOURCES.txt
10
9
  src/mllogs.egg-info/dependency_links.txt
11
10
  src/mllogs.egg-info/requires.txt
12
11
  src/mllogs.egg-info/top_level.txt
12
+ src/mllogs/storage/__init__.py
13
+ src/mllogs/storage/storage.py
14
+ src/mllogs/storage/db/__init__.py
15
+ src/mllogs/storage/db/base.py
16
+ src/mllogs/storage/db/sqlite.py
13
17
  tests/test_client.py
14
- tests/test_run.py
15
- tests/test_storage.py
18
+ tests/test_package.py
@@ -1,13 +1,24 @@
1
- from pathlib import Path
2
1
  import pytest
3
2
 
4
3
  from mllogs.run import RunStatus
5
4
  from mllogs.client import MLLogsClient
6
- from mllogs.storage import LocalFileStore
5
+ from mllogs.storage import Storage
6
+ from mllogs.storage.db import SQLiteStore
7
7
 
8
+ # ====================
9
+ # ----- Fixtures -----
10
+ # ====================
8
11
 
9
- def test_start_run() -> None:
10
- client = MLLogsClient()
12
+
13
+ @pytest.fixture
14
+ def client(tmp_path):
15
+ storage=Storage(
16
+ db=SQLiteStore(tmp_path / "mllogs.db")
17
+ )
18
+ return MLLogsClient(storage=storage)
19
+
20
+
21
+ def test_start_run(client):
11
22
  client.start_run(
12
23
  name="run1",
13
24
  run_type="training",
@@ -34,14 +45,11 @@ def test_start_run() -> None:
34
45
  assert run.ended_at is None
35
46
 
36
47
 
37
- def test_end_run(tmp_path: Path) -> None:
38
- file_store = LocalFileStore(root_dir=tmp_path)
39
- client = MLLogsClient(file_store=file_store)
40
-
48
+ def test_end_run(client):
41
49
  client.start_run()
42
50
  run = client.end_run()
43
51
 
44
- # assert run is ended
52
+ # assert active run is cleared
45
53
  assert client._active_run is None
46
54
 
47
55
  # assert ended run is complete
@@ -49,12 +57,10 @@ def test_end_run(tmp_path: Path) -> None:
49
57
  assert run.ended_at is not None
50
58
 
51
59
  # assert ended run is persisted to storage
52
- assert (tmp_path / "runs" / f"{run.id}.json").is_file()
53
-
60
+ assert client.get_run(run.id) == run
54
61
 
55
- def test_mutate_params() -> None:
56
- client = MLLogsClient()
57
62
 
63
+ def test_log_run_data(client):
58
64
  client.start_run()
59
65
 
60
66
  # log param
@@ -70,9 +76,7 @@ def test_mutate_params() -> None:
70
76
  assert client._active_run.tags["ml_model"] == "ridge"
71
77
 
72
78
 
73
- def test_ops_requiring_active_run() -> None:
74
- client = MLLogsClient()
75
-
79
+ def test_ops_requiring_active_run(client):
76
80
  with pytest.raises(RuntimeError):
77
81
  client.log_param("alpha", 0.1)
78
82
 
@@ -83,5 +87,4 @@ def test_ops_requiring_active_run() -> None:
83
87
  client.set_tag("model", "ridge")
84
88
 
85
89
  with pytest.raises(RuntimeError):
86
- client.end_run()
87
-
90
+ client.end_run()
@@ -0,0 +1,4 @@
1
+ from mllogs import MLLogsClient
2
+
3
+ def test_public_import() -> None:
4
+ assert MLLogsClient is not None
File without changes
@@ -1,88 +0,0 @@
1
- from datetime import datetime
2
- from dataclasses import dataclass, field
3
- from typing import Any
4
- from enum import Enum
5
-
6
-
7
- class RunStatus(Enum):
8
- RUNNING = "running"
9
- COMPLETE = "complete"
10
- FAILED = "failed"
11
-
12
-
13
- @dataclass
14
- class Artifact:
15
- name: str
16
- uri: str
17
- artifact_type: str
18
-
19
- def to_dict(self) -> dict[str, str]:
20
- return {
21
- "name": self.name,
22
- "uri": self.uri,
23
- "artifact_type": self.artifact_type,
24
- }
25
-
26
-
27
- @classmethod
28
- def from_dict(cls, d: dict[str, str]) -> "Artifact":
29
- return cls(
30
- name = d["name"],
31
- uri = d["uri"],
32
- artifact_type = d["artifact_type"],
33
- )
34
-
35
-
36
- @dataclass
37
- class Run:
38
- id: str
39
- started_at: datetime
40
- status: RunStatus
41
-
42
- name: str | None = None
43
- run_type: str | None = None
44
- ended_at: datetime | None = None
45
-
46
- params: dict[str, Any] = field(default_factory=dict)
47
- metrics: dict[str, float] = field(default_factory=dict)
48
- tags: dict[str, str] = field(default_factory=dict)
49
- artifacts: list[Artifact] = field(default_factory=list)
50
-
51
-
52
- def to_dict(self) -> dict[str, Any]:
53
- return {
54
- "id": self.id,
55
- "started_at": self.started_at.isoformat(),
56
- "status": self.status.value,
57
- "name": self.name,
58
- "run_type": self.run_type,
59
- "ended_at": (
60
- self.ended_at.isoformat()
61
- if self.ended_at is not None
62
- else None
63
- ),
64
- "params": self.params,
65
- "metrics": self.metrics,
66
- "tags": self.tags,
67
- "artifacts": [a.to_dict() for a in self.artifacts],
68
- }
69
-
70
-
71
- @classmethod
72
- def from_dict(cls, d: dict[str, Any]) -> "Run":
73
- return cls(
74
- id = d["id"],
75
- started_at = datetime.fromisoformat(d["started_at"]),
76
- status = RunStatus(d["status"]),
77
- name = d["name"],
78
- run_type = d["run_type"],
79
- ended_at = (
80
- datetime.fromisoformat(d["ended_at"])
81
- if d["ended_at"] is not None
82
- else None
83
- ),
84
- params = d["params"],
85
- metrics = d["metrics"],
86
- tags = d["tags"],
87
- artifacts = [Artifact.from_dict(a) for a in d["artifacts"]],
88
- )
@@ -1,86 +0,0 @@
1
- from pathlib import Path
2
- import json
3
-
4
- from .run import Run
5
-
6
-
7
- class LocalFileStore:
8
- def __init__(self, root_dir: str | Path = ".mllogs") -> None:
9
- """
10
- Initializes LocalFileStore and creates runs directory if it
11
- does not already exist.
12
- """
13
- self._root_dir = Path(root_dir)
14
- self._runs_dir = self._root_dir / "runs"
15
-
16
- self._runs_dir.mkdir(parents=True, exist_ok=True)
17
-
18
- def save_run(self, run: Run) -> None:
19
- """
20
- Writes Run to storage as JSON file.
21
- """
22
- path = self._runs_dir / f"{run.id}.json"
23
-
24
- with path.open("w") as f:
25
- json.dump(run.to_dict(), f, indent=4)
26
-
27
-
28
- def load_run(self, run_id: str | None = None) -> Run | None:
29
- """
30
- Reads JSON run file from storage and returns a Run.
31
-
32
- Args:
33
- - run_id: returns latest run if no argument is passed
34
-
35
- Returns:
36
- - Run if exists else None
37
- """
38
- if run_id is None:
39
- runs = self.list_runs(limit=1)
40
-
41
- if not runs:
42
- return None
43
-
44
- return runs[0]
45
-
46
- path = self._runs_dir / f"{run_id}.json"
47
-
48
- with path.open("r") as f:
49
- run_dict = json.load(f)
50
-
51
- return Run.from_dict(run_dict)
52
-
53
-
54
- def list_runs(self, limit: int | None = None) -> list[Run]:
55
- """
56
- Returns the most recent runs.
57
-
58
- If limit is None, returns all runs.
59
- """
60
- if limit is not None and limit <= 0:
61
- raise ValueError("limit must be greater than 0")
62
-
63
- paths = sorted(
64
- self._runs_dir.glob("*.json"),
65
- reverse=True, # latest first
66
- )
67
-
68
- if limit is not None:
69
- paths = paths[:limit]
70
-
71
- runs = []
72
-
73
- for path in paths:
74
- run = self.load_run(path.stem)
75
- runs.append(run)
76
-
77
- return runs
78
-
79
-
80
- def delete_run(self, run_id: str) -> None:
81
- """
82
- Deletes a run from storage by run ID.
83
- """
84
- path = self._runs_dir / f"{run_id}.json"
85
-
86
- path.unlink()
@@ -1,69 +0,0 @@
1
- from datetime import datetime, UTC
2
-
3
- from mllogs.run import Run, RunStatus
4
-
5
-
6
- def test_run_defaults():
7
- started_at = datetime(2000, 1, 1, 0, 0, tzinfo=UTC)
8
- run = Run(
9
- id="run-123",
10
- status=RunStatus.RUNNING,
11
- started_at=started_at,
12
- )
13
-
14
- assert run.name is None
15
- assert run.run_type is None
16
- assert run.ended_at is None
17
- assert run.params == {}
18
- assert run.metrics == {}
19
- assert run.tags == {}
20
- assert run.artifacts == []
21
-
22
-
23
- def test_to_dict():
24
- started_at = datetime(2000, 1, 1, 0, 0, tzinfo=UTC)
25
- run = Run(
26
- id="run-123",
27
- status=RunStatus.RUNNING,
28
- started_at=started_at,
29
- )
30
-
31
- d = run.to_dict()
32
-
33
- assert d["id"] == "run-123"
34
- assert d["status"] == "running"
35
- assert d["started_at"] == "2000-01-01T00:00:00+00:00"
36
-
37
-
38
- def test_from_dict():
39
- d = {
40
- "id": "run-123",
41
- "status": "running",
42
- "started_at": "2000-01-01T00:00:00+00:00",
43
- "name": None,
44
- "run_type": None,
45
- "ended_at": None,
46
- "params": {},
47
- "metrics": {},
48
- "tags": {},
49
- "artifacts": [],
50
- }
51
-
52
- run = Run.from_dict(d)
53
-
54
- assert run.started_at == datetime(2000, 1, 1, 0, 0, tzinfo=UTC)
55
- assert run.status == RunStatus.RUNNING
56
-
57
-
58
- def test_serialization_round_trip():
59
- started_at = datetime(2000, 1, 1, 0, 0, tzinfo=UTC)
60
-
61
- before = Run(
62
- id="run-123",
63
- status=RunStatus.RUNNING,
64
- started_at=started_at,
65
- )
66
-
67
- after = Run.from_dict(before.to_dict())
68
-
69
- assert before == after
@@ -1,89 +0,0 @@
1
- from pathlib import Path
2
- from datetime import datetime, UTC
3
- import pytest
4
-
5
- from mllogs.run import Run, RunStatus
6
- from mllogs.storage import LocalFileStore
7
-
8
-
9
- def test_persistence_round_trip(tmp_path: Path) -> None:
10
- storage = LocalFileStore(root_dir=tmp_path)
11
-
12
- run = Run(
13
- id="run123",
14
- status=RunStatus.RUNNING,
15
- started_at=datetime.now(UTC),
16
- )
17
- run_path = storage._runs_dir / f"{run.id}.json"
18
-
19
- # save
20
- storage.save_run(run)
21
- assert run_path.is_file()
22
-
23
- # load
24
- run_after = storage.load_run(run.id)
25
- assert run == run_after
26
-
27
- # delete
28
- storage.delete_run(run.id)
29
- assert not run_path.exists()
30
-
31
-
32
- def test_list_runs(tmp_path: Path) -> None:
33
- storage = LocalFileStore(root_dir=tmp_path)
34
-
35
- for run_id in ["001", "002", "003"]:
36
- storage.save_run(Run(
37
- id=run_id,
38
- status=RunStatus.RUNNING,
39
- started_at=datetime.now(UTC),
40
- ))
41
-
42
- runs = storage.list_runs()
43
-
44
- assert [run.id for run in runs] == ["003", "002", "001"]
45
-
46
-
47
- def test_list_runs_with_limit(tmp_path: Path) -> None:
48
- storage = LocalFileStore(root_dir=tmp_path)
49
-
50
- for run_id in ["001", "002", "003"]:
51
- storage.save_run(Run(
52
- id=run_id,
53
- status=RunStatus.RUNNING,
54
- started_at=datetime.now(UTC),
55
- ))
56
-
57
- runs = storage.list_runs(limit=2)
58
-
59
- assert [run.id for run in runs] == ["003", "002"]
60
-
61
-
62
- def test_runs_with_invalid_limit(tmp_path: Path) -> None:
63
- storage = LocalFileStore(root_dir=tmp_path)
64
-
65
- with pytest.raises(ValueError):
66
- storage.list_runs(limit=0)
67
-
68
-
69
- def test_load_latest_run(tmp_path: Path) -> None:
70
- storage = LocalFileStore(root_dir=tmp_path)
71
-
72
- run1 = Run(
73
- id="001",
74
- status=RunStatus.RUNNING,
75
- started_at=datetime.now(UTC),
76
- )
77
-
78
- run2 = Run(
79
- id="002",
80
- status=RunStatus.RUNNING,
81
- started_at=datetime.now(UTC),
82
- )
83
-
84
- storage.save_run(run1)
85
- storage.save_run(run2)
86
-
87
- latest_run = storage.load_run()
88
-
89
- assert latest_run == run2
File without changes
File without changes