CoolFace
Apppublic

abidlabs/trackio-123

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
sqlite_storage.py153 linesDownload Raw Back to root
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