abidlabs/trackio-123
0
1import json2import os3import sqlite34 5from huggingface_hub import CommitScheduler6 7try:8 from trackio.dummy_commit_scheduler import DummyCommitScheduler9 from trackio.utils import RESERVED_KEYS, TRACKIO_DIR10except: # noqa: E72211 from dummy_commit_scheduler import DummyCommitScheduler12 from utils import RESERVED_KEYS, TRACKIO_DIR13 14 15class SQLiteStorage:16 def __init__(17 self, project: str, name: str, config: dict, dataset_id: str | None = None18 ):19 self.project = project20 self.name = name21 self.config = config22 self.db_path = os.path.join(TRACKIO_DIR, "trackio.db")23 self.dataset_id = dataset_id24 self.scheduler = self._get_scheduler()25 26 os.makedirs(TRACKIO_DIR, exist_ok=True)27 28 self._init_db()29 self._save_config()30 31 def _get_scheduler(self):32 hf_token = os.environ.get(33 "HF_TOKEN"34 ) # Get the token from the environment variable on Spaces35 dataset_id = self.dataset_id or os.environ.get("TRACKIO_DATASET_ID")36 if dataset_id is None:37 scheduler = DummyCommitScheduler()38 else:39 scheduler = CommitScheduler(40 repo_id=dataset_id,41 repo_type="dataset",42 folder_path=TRACKIO_DIR,43 private=True,44 squash_history=True,45 token=hf_token,46 )47 return scheduler48 49 def _init_db(self):50 """Initialize the SQLite database with required tables."""51 with self.scheduler.lock:52 with sqlite3.connect(self.db_path) as conn:53 cursor = conn.cursor()54 55 cursor.execute("""56 CREATE TABLE IF NOT EXISTS metrics (57 id INTEGER PRIMARY KEY AUTOINCREMENT,58 timestamp DATETIME DEFAULT CURRENT_TIMESTAMP,59 project_name TEXT NOT NULL,60 run_name TEXT NOT NULL,61 metrics TEXT NOT NULL62 )63 """)64 65 cursor.execute("""66 CREATE TABLE IF NOT EXISTS configs (67 project_name TEXT NOT NULL,68 run_name TEXT NOT NULL,69 config TEXT NOT NULL,70 created_at DATETIME DEFAULT CURRENT_TIMESTAMP,71 PRIMARY KEY (project_name, run_name)72 )73 """)74 75 conn.commit()76 77 def _save_config(self):78 """Save the run configuration to the database."""79 with self.scheduler.lock:80 with sqlite3.connect(self.db_path) as conn:81 cursor = conn.cursor()82 cursor.execute(83 "INSERT OR REPLACE INTO configs (project_name, run_name, config) VALUES (?, ?, ?)",84 (self.project, self.name, json.dumps(self.config)),85 )86 conn.commit()87 88 def log(self, metrics: dict):89 """Log metrics to the database."""90 for k in metrics.keys():91 if k in RESERVED_KEYS or k.startswith("__"):92 raise ValueError(93 f"Please do not use this reserved key as a metric: {k}"94 )95 96 with self.scheduler.lock:97 with sqlite3.connect(self.db_path) as conn:98 cursor = conn.cursor()99 cursor.execute(100 """101 INSERT INTO metrics 102 (project_name, run_name, metrics)103 VALUES (?, ?, ?)104 """,105 (self.project, self.name, json.dumps(metrics)),106 )107 conn.commit()108 109 def get_metrics(self, project: str, run: str) -> list[dict]:110 """Retrieve metrics for a specific run."""111 with sqlite3.connect(self.db_path) as conn:112 cursor = conn.cursor()113 cursor.execute(114 """115 SELECT timestamp, metrics116 FROM metrics117 WHERE project_name = ? AND run_name = ?118 ORDER BY timestamp119 """,120 (project, run),121 )122 rows = cursor.fetchall()123 124 results = []125 for row in rows:126 timestamp, metrics_json = row127 metrics = json.loads(metrics_json)128 metrics["timestamp"] = timestamp129 results.append(metrics)130 131 return results132 133 def get_projects(self) -> list[str]:134 """Get list of all projects."""135 with sqlite3.connect(self.db_path) as conn:136 cursor = conn.cursor()137 cursor.execute("SELECT DISTINCT project_name FROM metrics")138 return [row[0] for row in cursor.fetchall()]139 140 def get_runs(self, project: str) -> list[str]:141 """Get list of all runs for a project."""142 with sqlite3.connect(self.db_path) as conn:143 cursor = conn.cursor()144 cursor.execute(145 "SELECT DISTINCT run_name FROM metrics WHERE project_name = ?",146 (project,),147 )148 return [row[0] for row in cursor.fetchall()]149 150 def finish(self):151 """Cleanup when run is finished."""152 pass153 