Team Ai
Apppublic

pankajagarwal/function-gemma

sourceHugging Faceupdated 10mo agoView on Hugging Face
0likes
run.py242 linesDownload Raw Back to root
1import threading2import time3import warnings4from datetime import datetime, timezone5 6import huggingface_hub7from gradio_client import Client, handle_file8 9from trackio import utils10from trackio.histogram import Histogram11from trackio.media import TrackioMedia12from trackio.sqlite_storage import SQLiteStorage13from trackio.table import Table14from trackio.typehints import LogEntry, UploadEntry15from trackio.utils import _get_default_namespace16 17BATCH_SEND_INTERVAL = 0.518 19 20class Run:21    def __init__(22        self,23        url: str,24        project: str,25        client: Client | None,26        name: str | None = None,27        group: str | None = None,28        config: dict | None = None,29        space_id: str | None = None,30    ):31        self.url = url32        self.project = project33        self._client_lock = threading.Lock()34        self._client_thread = None35        self._client = client36        self._space_id = space_id37        self.name = name or utils.generate_readable_name(38            SQLiteStorage.get_runs(project), space_id39        )40        self.group = group41        self.config = utils.to_json_safe(config or {})42 43        if isinstance(self.config, dict):44            for key in self.config:45                if key.startswith("_"):46                    raise ValueError(47                        f"Config key '{key}' is reserved (keys starting with '_' are reserved for internal use)"48                    )49 50        self.config["_Username"] = self._get_username()51        self.config["_Created"] = datetime.now(timezone.utc).isoformat()52        self.config["_Group"] = self.group53 54        self._queued_logs: list[LogEntry] = []55        self._queued_uploads: list[UploadEntry] = []56        self._stop_flag = threading.Event()57        self._config_logged = False58 59        self._client_thread = threading.Thread(target=self._init_client_background)60        self._client_thread.daemon = True61        self._client_thread.start()62 63    def _get_username(self) -> str | None:64        """Get the current HuggingFace username if logged in, otherwise None."""65        try:66            return _get_default_namespace()67        except Exception:68            return None69 70    def _batch_sender(self):71        """Send batched logs every BATCH_SEND_INTERVAL."""72        while not self._stop_flag.is_set() or len(self._queued_logs) > 0:73            if not self._stop_flag.is_set():74                time.sleep(BATCH_SEND_INTERVAL)75 76            with self._client_lock:77                if self._client is None:78                    return79                if self._queued_logs:80                    logs_to_send = self._queued_logs.copy()81                    self._queued_logs.clear()82                    self._client.predict(83                        api_name="/bulk_log",84                        logs=logs_to_send,85                        hf_token=huggingface_hub.utils.get_token(),86                    )87                if self._queued_uploads:88                    uploads_to_send = self._queued_uploads.copy()89                    self._queued_uploads.clear()90                    self._client.predict(91                        api_name="/bulk_upload_media",92                        uploads=uploads_to_send,93                        hf_token=huggingface_hub.utils.get_token(),94                    )95 96    def _init_client_background(self):97        if self._client is None:98            fib = utils.fibo()99            for sleep_coefficient in fib:100                try:101                    client = Client(self.url, verbose=False)102 103                    with self._client_lock:104                        self._client = client105                    break106                except Exception:107                    pass108                if sleep_coefficient is not None:109                    time.sleep(0.1 * sleep_coefficient)110 111        self._batch_sender()112 113    def _queue_upload(114        self,115        file_path,116        step: int | None,117        relative_path: str | None = None,118        use_run_name: bool = True,119    ):120        """121        Queues a media file for upload to a Space.122 123        Args:124            file_path:125                The path to the file to upload.126            step (`int` or `None`, *optional*):127                The step number associated with this upload.128            relative_path (`str` or `None`, *optional*):129                The relative path within the project's files directory. Used when130                uploading files via `trackio.save()`.131            use_run_name (`bool`, *optional*):132                Whether to use the run name for the uploaded file. This is set to133                `False` when uploading files via `trackio.save()`.134        """135        upload_entry: UploadEntry = {136            "project": self.project,137            "run": self.name if use_run_name else None,138            "step": step,139            "relative_path": relative_path,140            "uploaded_file": handle_file(file_path),141        }142        with self._client_lock:143            self._queued_uploads.append(upload_entry)144 145    def _process_media(self, value: TrackioMedia, step: int | None) -> dict:146        """147        Serialize media in metrics and upload to space if needed.148        """149        value._save(self.project, self.name, step)150        if self._space_id:151            self._queue_upload(value._get_absolute_file_path(), step)152        return value._to_dict()153 154    def _scan_and_queue_media_uploads(self, table_dict: dict, step: int | None):155        """156        Scan a serialized table for media objects and queue them for upload to space.157        """158        if not self._space_id:159            return160 161        table_data = table_dict.get("_value", [])162        for row in table_data:163            for value in row.values():164                if isinstance(value, dict) and value.get("_type") in [165                    "trackio.image",166                    "trackio.video",167                    "trackio.audio",168                ]:169                    file_path = value.get("file_path")170                    if file_path:171                        from trackio.utils import MEDIA_DIR172 173                        absolute_path = MEDIA_DIR / file_path174                        self._queue_upload(absolute_path, step)175                elif isinstance(value, list):176                    for item in value:177                        if isinstance(item, dict) and item.get("_type") in [178                            "trackio.image",179                            "trackio.video",180                            "trackio.audio",181                        ]:182                            file_path = item.get("file_path")183                            if file_path:184                                from trackio.utils import MEDIA_DIR185 186                                absolute_path = MEDIA_DIR / file_path187                                self._queue_upload(absolute_path, step)188 189    def log(self, metrics: dict, step: int | None = None):190        renamed_keys = []191        new_metrics = {}192 193        for k, v in metrics.items():194            if k in utils.RESERVED_KEYS or k.startswith("__"):195                new_key = f"__{k}"196                renamed_keys.append(k)197                new_metrics[new_key] = v198            else:199                new_metrics[k] = v200 201        if renamed_keys:202            warnings.warn(f"Reserved keys renamed: {renamed_keys} → '__{{key}}'")203 204        metrics = new_metrics205        for key, value in metrics.items():206            if isinstance(value, Table):207                metrics[key] = value._to_dict(208                    project=self.project, run=self.name, step=step209                )210                self._scan_and_queue_media_uploads(metrics[key], step)211            elif isinstance(value, Histogram):212                metrics[key] = value._to_dict()213            elif isinstance(value, TrackioMedia):214                metrics[key] = self._process_media(value, step)215        metrics = utils.serialize_values(metrics)216 217        config_to_log = None218        if not self._config_logged and self.config:219            config_to_log = utils.to_json_safe(self.config)220            self._config_logged = True221 222        log_entry: LogEntry = {223            "project": self.project,224            "run": self.name,225            "metrics": metrics,226            "step": step,227            "config": config_to_log,228        }229 230        with self._client_lock:231            self._queued_logs.append(log_entry)232 233    def finish(self):234        """Cleanup when run is finished."""235        self._stop_flag.set()236 237        time.sleep(2 * BATCH_SEND_INTERVAL)238 239        if self._client_thread is not None:240            print("* Run finished. Uploading logs to Trackio (please wait...)")241            self._client_thread.join()242