WalisonCruz/function-gemma
0
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 