delsj/function-gemma
0
1"""The main page for the Trackio UI."""2 3import os4import re5import secrets6import shutil7from dataclasses import dataclass8from typing import Any9 10import gradio as gr11import numpy as np12import pandas as pd13 14try:15 import trackio.utils as utils16 from trackio.media import (17 TrackioAudio,18 TrackioImage,19 TrackioVideo,20 get_project_media_path,21 )22 from trackio.sqlite_storage import SQLiteStorage23 from trackio.typehints import LogEntry, UploadEntry24 from trackio.ui import fns25 from trackio.ui.components.colored_checkbox import ColoredCheckboxGroup26 from trackio.ui.files import files_page27 from trackio.ui.helpers.run_selection import RunSelection28 from trackio.ui.media_page import media_page29 from trackio.ui.run_detail import run_detail_page30 from trackio.ui.runs import run_page31except ImportError:32 import utils33 from media import (34 TrackioAudio,35 TrackioImage,36 TrackioVideo,37 get_project_media_path,38 )39 from sqlite_storage import SQLiteStorage40 from typehints import LogEntry, UploadEntry41 from ui import fns42 from ui.components.colored_checkbox import ColoredCheckboxGroup43 from ui.files import files_page44 from ui.helpers.run_selection import RunSelection45 from ui.media_page import media_page46 from ui.run_detail import run_detail_page47 from ui.runs import run_page48 49 50INSTRUCTIONS_SPACES = """51## Start logging with Trackio 🤗52 53To start logging to this Trackio dashboard, first make sure you have the Trackio library installed. You can do this by running:54 55```bash56pip install trackio57```58 59Then, start logging to this Trackio dashboard by passing in the `space_id` to `trackio.init()`:60 61```python62import trackio63trackio.init(project="my-project", space_id="{}")64```65 66Then call `trackio.log()` to log metrics.67 68```python69for i in range(10):70 trackio.log({{"loss": 1/(i+1)}})71```72 73Finally, call `trackio.finish()` to finish the run.74 75```python76trackio.finish()77```78"""79 80INSTRUCTIONS_LOCAL = """81## Start logging with Trackio 🤗82 83You can create a new project by calling `trackio.init()`:84 85```python86import trackio87trackio.init(project="my-project")88 ```89 90Then call `trackio.log()` to log metrics.91 92```python93for i in range(10):94 trackio.log({"loss": 1/(i+1)})95```96 97Finally, call `trackio.finish()` to finish the run.98 99```python100trackio.finish()101```102 103Read the [Trackio documentation](https://huggingface.co/docs/trackio/en/index) for more examples.104"""105 106 107def get_runs(project) -> list[str]:108 if not project:109 return []110 return SQLiteStorage.get_runs(project)111 112 113def upload_db_to_space(114 project: str, uploaded_db: gr.FileData, hf_token: str | None115) -> None:116 """117 Uploads the database of a local Trackio project to a Hugging Face Space.118 """119 fns.check_hf_token_has_write_access(hf_token)120 db_project_path = SQLiteStorage.get_project_db_path(project)121 os.makedirs(os.path.dirname(db_project_path), exist_ok=True)122 shutil.copy(uploaded_db["path"], db_project_path)123 124 125def get_available_metrics(project: str, runs: list[str]) -> list[str]:126 """Get all available metrics across all runs for x-axis selection."""127 if not project or not runs:128 return ["step", "time"]129 130 all_metrics = set()131 for run in runs:132 metrics = SQLiteStorage.get_logs(project, run)133 if metrics:134 df = pd.DataFrame(metrics)135 numeric_cols = df.select_dtypes(include="number").columns136 numeric_cols = [c for c in numeric_cols if c not in utils.RESERVED_KEYS]137 all_metrics.update(numeric_cols)138 139 all_metrics.add("step")140 all_metrics.add("time")141 142 sorted_metrics = utils.sort_metrics_by_prefix(list(all_metrics))143 144 result = ["step", "time"]145 for metric in sorted_metrics:146 if metric not in result:147 result.append(metric)148 149 return result150 151 152@dataclass153class MediaData:154 caption: str | None155 file_path: str156 type: str157 158 159def extract_media(logs: list[dict]) -> dict[str, list[MediaData]]:160 media_by_key: dict[str, list[MediaData]] = {}161 logs = sorted(logs, key=lambda x: x.get("step", 0))162 for log in logs:163 for key, value in log.items():164 if isinstance(value, dict):165 type = value.get("_type")166 if (167 type == TrackioImage.TYPE168 or type == TrackioVideo.TYPE169 or type == TrackioAudio.TYPE170 ):171 if key not in media_by_key:172 media_by_key[key] = []173 try:174 media_data = MediaData(175 file_path=utils.MEDIA_DIR / value.get("file_path"),176 type=type,177 caption=value.get("caption"),178 )179 media_by_key[key].append(media_data)180 except Exception as e:181 print(f"Media currently unavailable: {key}: {e}")182 return media_by_key183 184 185def load_run_data(186 project: str | None,187 run: str | None,188 smoothing_granularity: int = 0,189 x_axis: str = "step",190 log_scale_x: bool = False,191 log_scale_y: bool = False,192) -> tuple[pd.DataFrame, dict]:193 if not project or not run:194 return None, None195 196 logs = SQLiteStorage.get_logs(project, run)197 if not logs:198 return None, None199 200 media = extract_media(logs)201 df = pd.DataFrame(logs)202 203 if "step" not in df.columns:204 df["step"] = range(len(df))205 206 if x_axis == "time" and "timestamp" in df.columns:207 df["timestamp"] = pd.to_datetime(df["timestamp"])208 first_timestamp = df["timestamp"].min()209 df["time"] = (df["timestamp"] - first_timestamp).dt.total_seconds()210 x_column = "time"211 elif x_axis == "step":212 x_column = "step"213 else:214 x_column = x_axis215 216 if log_scale_x and x_column in df.columns:217 x_vals = df[x_column]218 if (x_vals <= 0).any():219 df[x_column] = np.log10(np.maximum(x_vals, 0) + 1)220 else:221 df[x_column] = np.log10(x_vals)222 223 if log_scale_y:224 numeric_cols = df.select_dtypes(include="number").columns225 y_cols = [226 c for c in numeric_cols if c not in utils.RESERVED_KEYS and c != x_column227 ]228 for y_col in y_cols:229 if y_col in df.columns:230 y_vals = df[y_col]231 if (y_vals <= 0).any():232 df[y_col] = np.log10(np.maximum(y_vals, 0) + 1)233 else:234 df[y_col] = np.log10(y_vals)235 236 if smoothing_granularity > 0:237 numeric_cols = df.select_dtypes(include="number").columns238 numeric_cols = [c for c in numeric_cols if c not in utils.RESERVED_KEYS]239 240 df_original = df.copy()241 df_original["run"] = run242 df_original["data_type"] = "original"243 244 df_smoothed = df.copy()245 window_size = max(3, min(smoothing_granularity, len(df)))246 df_smoothed[numeric_cols] = (247 df_smoothed[numeric_cols]248 .rolling(window=window_size, center=True, min_periods=1)249 .mean()250 )251 df_smoothed["run"] = f"{run}_smoothed"252 df_smoothed["data_type"] = "smoothed"253 254 combined_df = pd.concat([df_original, df_smoothed], ignore_index=True)255 combined_df["x_axis"] = x_column256 return combined_df, media257 else:258 df["run"] = run259 df["data_type"] = "original"260 df["x_axis"] = x_column261 return df, media262 263 264def refresh_runs(265 project: str | None,266 filter_text: str | None,267 selection: RunSelection,268 selected_runs_from_url: list[str] | None = None,269):270 if project is None:271 runs: list[str] = []272 else:273 runs = get_runs(project)274 if filter_text:275 runs = [r for r in runs if filter_text in r]276 277 preferred = None278 if selected_runs_from_url:279 preferred = [r for r in runs if r in selected_runs_from_url]280 281 did_change = selection.update_choices(runs, preferred)282 return (283 fns.run_checkbox_update(selection) if did_change else gr.skip(),284 gr.Textbox(label=f"Runs ({len(runs)})"),285 selection,286 )287 288 289def generate_embed(project: str, metrics: str, selection: RunSelection) -> str:290 return utils.generate_embed_code(project, metrics, selection.selected)291 292 293def update_x_axis_choices(project, selection):294 """Update x-axis dropdown choices based on available metrics."""295 runs = selection.selected296 available_metrics = get_available_metrics(project, runs)297 return gr.Dropdown(298 label="X-axis",299 choices=available_metrics,300 value="step",301 )302 303 304def toggle_timer(cb_value):305 if cb_value:306 return gr.Timer(active=True)307 else:308 return gr.Timer(active=False)309 310 311def bulk_upload_media(uploads: list[UploadEntry], hf_token: str | None) -> None:312 """313 Uploads media files to a Trackio dashboard. Each entry in the list is a tuple of the project, run, and media file to be uploaded.314 Also handles uplaoding project-level files to the project's files directory (if the run and step are not provided).315 """316 fns.check_hf_token_has_write_access(hf_token)317 for upload in uploads:318 media_path = get_project_media_path(319 project=upload["project"],320 run=upload["run"],321 step=upload["step"],322 relative_path=upload["relative_path"],323 )324 shutil.copy(upload["uploaded_file"]["path"], media_path)325 326 327def log(328 project: str,329 run: str,330 metrics: dict[str, Any],331 step: int | None,332 hf_token: str | None,333) -> None:334 """335 Note: this method is not used in the latest versions of Trackio (replaced by bulk_log) but336 is kept for backwards compatibility for users who are connecting to a newer version of337 a Trackio Spaces dashboard with an older version of Trackio installed locally.338 """339 fns.check_hf_token_has_write_access(hf_token)340 SQLiteStorage.log(project=project, run=run, metrics=metrics, step=step)341 342 343def bulk_log(344 logs: list[LogEntry],345 hf_token: str | None,346) -> None:347 """348 Logs a list of metrics to a Trackio dashboard. Each entry in the list is a dictionary of the project, run, a dictionary of metrics, and optionally, a step and config.349 """350 fns.check_hf_token_has_write_access(hf_token)351 352 logs_by_run = {}353 for log_entry in logs:354 key = (log_entry["project"], log_entry["run"])355 if key not in logs_by_run:356 logs_by_run[key] = {"metrics": [], "steps": [], "config": None}357 logs_by_run[key]["metrics"].append(log_entry["metrics"])358 logs_by_run[key]["steps"].append(log_entry.get("step"))359 if log_entry.get("config") and logs_by_run[key]["config"] is None:360 logs_by_run[key]["config"] = log_entry["config"]361 362 for (project, run), data in logs_by_run.items():363 SQLiteStorage.bulk_log(364 project=project,365 run=run,366 metrics_list=data["metrics"],367 steps=data["steps"],368 config=data["config"],369 )370 371 372def get_metric_values(373 project: str,374 run: str,375 metric_name: str,376) -> list[dict]:377 """378 Get all values for a specific metric in a project/run.379 Returns a list of dictionaries with timestamp, step, and value.380 """381 return SQLiteStorage.get_metric_values(project, run, metric_name)382 383 384def get_runs_for_project(385 project: str,386) -> list[str]:387 """388 Get all runs for a given project.389 Returns a list of run names.390 """391 return SQLiteStorage.get_runs(project)392 393 394def get_metrics_for_run(395 project: str,396 run: str,397) -> list[str]:398 """399 Get all metrics for a given project and run.400 Returns a list of metric names.401 """402 return SQLiteStorage.get_all_metrics_for_run(project, run)403 404 405def filter_metrics_by_regex(metrics: list[str], filter_pattern: str) -> list[str]:406 """407 Filter metrics using regex pattern.408 409 Args:410 metrics: List of metric names to filter411 filter_pattern: Regex pattern to match against metric names412 413 Returns:414 List of metric names that match the pattern415 """416 if not filter_pattern.strip():417 return metrics418 419 try:420 pattern = re.compile(filter_pattern, re.IGNORECASE)421 return [metric for metric in metrics if pattern.search(metric)]422 except re.error:423 return [424 metric for metric in metrics if filter_pattern.lower() in metric.lower()425 ]426 427 428def get_all_projects() -> list[str]:429 """430 Get all project names.431 Returns a list of project names.432 """433 return SQLiteStorage.get_projects()434 435 436def get_project_summary(project: str) -> dict:437 """438 Get a summary of a project including number of runs and recent activity.439 440 Args:441 project: Project name442 443 Returns:444 Dictionary with project summary information445 """446 runs = SQLiteStorage.get_runs(project)447 if not runs:448 return {"project": project, "num_runs": 0, "runs": [], "last_activity": None}449 450 last_steps = SQLiteStorage.get_max_steps_for_runs(project)451 452 return {453 "project": project,454 "num_runs": len(runs),455 "runs": runs,456 "last_activity": max(last_steps.values()) if last_steps else None,457 }458 459 460def get_run_summary(project: str, run: str) -> dict:461 """462 Get a summary of a specific run including metrics and configuration.463 464 Args:465 project: Project name466 run: Run name467 468 Returns:469 Dictionary with run summary information470 """471 logs = SQLiteStorage.get_logs(project, run)472 metrics = SQLiteStorage.get_all_metrics_for_run(project, run)473 474 if not logs:475 return {476 "project": project,477 "run": run,478 "num_logs": 0,479 "metrics": [],480 "config": None,481 "last_step": None,482 }483 484 df = pd.DataFrame(logs)485 config = logs[0].get("config") if logs else None486 last_step = df["step"].max() if "step" in df.columns else len(logs) - 1487 488 return {489 "project": project,490 "run": run,491 "num_logs": len(logs),492 "metrics": metrics,493 "config": config,494 "last_step": last_step,495 }496 497 498def configure(request: gr.Request):499 sidebar_param = request.query_params.get("sidebar")500 match sidebar_param:501 case "collapsed":502 sidebar = gr.Sidebar(open=False, visible=True)503 case "hidden":504 sidebar = gr.Sidebar(open=False, visible=False)505 case _:506 sidebar = gr.Sidebar(open=True, visible=True)507 508 metrics_param = request.query_params.get("metrics", "")509 runs_param = request.query_params.get("runs", "")510 selected_runs = runs_param.split(",") if runs_param else []511 navbar_param = request.query_params.get("navbar")512 x_min_param = request.query_params.get("xmin")513 x_max_param = request.query_params.get("xmax")514 x_min = float(x_min_param) if x_min_param is not None else None515 x_max = float(x_max_param) if x_max_param is not None else None516 smoothing_param = request.query_params.get("smoothing")517 smoothing_value = int(smoothing_param) if smoothing_param is not None else 10518 519 match navbar_param:520 case "hidden":521 navbar = gr.Navbar(visible=False)522 case _:523 navbar = gr.Navbar(visible=True)524 525 return (526 [],527 sidebar,528 metrics_param,529 selected_runs,530 navbar,531 [x_min, x_max],532 smoothing_value,533 )534 535 536CSS = """537.logo-light { display: block; } 538.logo-dark { display: none; }539.dark .logo-light { display: none; }540.dark .logo-dark { display: block; }541.dark .caption-label { color: white; }542 543.info-container {544 position: relative;545 display: inline;546}547.info-checkbox {548 position: absolute;549 opacity: 0;550 pointer-events: none;551}552.info-icon {553 border-bottom: 1px dotted;554 cursor: pointer;555 user-select: none;556 color: var(--color-accent);557}558.info-expandable {559 display: none;560 opacity: 0;561 transition: opacity 0.2s ease-in-out;562}563.info-checkbox:checked ~ .info-expandable {564 display: inline;565 opacity: 1;566}567.info-icon:hover { opacity: 0.8; }568.accent-link { font-weight: bold; }569 570.media-gallery .fixed-height { min-height: 275px; }571.media-group, .media-group > div { background: none; }572.media-group .tabs { padding: 0.5em; }573.media-tab { max-height: 500px; overflow-y: scroll; }574.media-audio-accordion > button { 575 border-bottom-width: 1px;576 padding-bottom: 3px;577}578.media-audio-item {579 border-width: 1px !important;580 border-radius: 0.5em;581}582.media-audio-row {583 gap: 0.25em;584 margin-bottom: 0.25em;585}586 587.tab-like-container {588 visibility: hidden;589}590"""591 592HEAD = """593<script>594function setCookie(name, value, days) {595 var expires = "";596 if (days) {597 var date = new Date();598 date.setTime(date.getTime() + (days * 24 * 60 * 60 * 1000));599 expires = "; expires=" + date.toUTCString();600 }601 document.cookie = name + "=" + (value || "") + expires + "; path=/; SameSite=Lax";602}603 604function getCookie(name) {605 var nameEQ = name + "=";606 var ca = document.cookie.split(';');607 for(var i=0;i < ca.length;i++) {608 var c = ca[i];609 while (c.charAt(0)==' ') c = c.substring(1,c.length);610 if (c.indexOf(nameEQ) == 0) return c.substring(nameEQ.length,c.length);611 }612 return null;613}614 615(function() {616 const urlParams = new URLSearchParams(window.location.search);617 const writeToken = urlParams.get('write_token');618 const footerParam = urlParams.get('footer');619 620 if (writeToken) {621 setCookie('trackio_write_token', writeToken, 7);622 623 // Only remove write_token from URL if not in iframe624 // In iframes, keep it in URL as cookies may be blocked625 const inIframe = window.self !== window.top;626 if (!inIframe) {627 urlParams.delete('write_token');628 const newUrl = window.location.pathname + 629 (urlParams.toString() ? '?' + urlParams.toString() : '') + 630 window.location.hash;631 window.history.replaceState({}, document.title, newUrl);632 }633 }634 635 if (footerParam === 'false') {636 const style = document.createElement('style');637 style.textContent = 'footer { display: none !important; }';638 document.head.appendChild(style);639 }640})();641</script>642"""643 644 645gr.set_static_paths(paths=[utils.MEDIA_DIR])646 647with gr.Blocks(title="Trackio Dashboard") as demo:648 with gr.Sidebar(open=False) as sidebar:649 logo_urls = utils.get_logo_urls()650 logo = gr.Markdown(651 f"""652 <img src='{logo_urls["light"]}' width='80%' class='logo-light'>653 <img src='{logo_urls["dark"]}' width='80%' class='logo-dark'> 654 """655 )656 project_dd = gr.Dropdown(label="Project", allow_custom_value=True)657 658 embed_code = gr.Code(659 label="Embed this view",660 max_lines=2,661 lines=2,662 language="html",663 visible=bool(os.environ.get("SPACE_HOST")),664 )665 with gr.Group():666 run_tb = gr.Textbox(label="Runs", placeholder="Type to filter...")667 run_group_by_dd = gr.Dropdown(label="Group by...", choices=[], value=None)668 grouped_runs_panel = gr.Group(visible=False)669 run_cb = ColoredCheckboxGroup(choices=[], colors=[], label="Runs")670 671 gr.HTML("<hr>")672 realtime_cb = gr.Checkbox(label="Refresh metrics realtime", value=True)673 smoothing_slider = gr.Slider(674 label="Smoothing Factor",675 minimum=0,676 maximum=20,677 value=10,678 step=1,679 info="0 = no smoothing",680 )681 x_axis_dd = gr.Dropdown(682 label="X-axis",683 choices=["step", "time"],684 value="step",685 )686 log_scale_x_cb = gr.Checkbox(label="Log scale X-axis", value=False)687 log_scale_y_cb = gr.Checkbox(label="Log scale Y-axis", value=False)688 metric_filter_tb = gr.Textbox(689 label="Metric Filter (regex)",690 placeholder="e.g., loss|ndcg@10|gpu",691 value="",692 info="Filter metrics using regex patterns. Leave empty to show all metrics.",693 )694 695 navbar = gr.Navbar(696 value=[697 ("Metrics", ""),698 ("Media & Tables", "/media"),699 ("Runs", "/runs"),700 ("Files", "/files"),701 ],702 main_page_name=False,703 )704 timer = gr.Timer(value=1)705 metrics_subset = gr.State([])706 selected_runs_from_url = gr.State([])707 run_selection_state = gr.State(RunSelection())708 x_lim = gr.State(None)709 710 gr.on(711 [demo.load],712 fn=configure,713 outputs=[714 metrics_subset,715 sidebar,716 metric_filter_tb,717 selected_runs_from_url,718 navbar,719 x_lim,720 smoothing_slider,721 ],722 queue=False,723 api_visibility="private",724 )725 gr.on(726 [demo.load],727 fn=fns.get_projects,728 outputs=project_dd,729 show_progress="hidden",730 queue=False,731 api_visibility="private",732 )733 gr.on(734 [timer.tick],735 fn=refresh_runs,736 inputs=[project_dd, run_tb, run_selection_state, selected_runs_from_url],737 outputs=[run_cb, run_tb, run_selection_state],738 show_progress="hidden",739 api_visibility="private",740 )741 gr.on(742 [timer.tick],743 fn=lambda: gr.Dropdown(info=fns.get_project_info()),744 outputs=[project_dd],745 show_progress="hidden",746 api_visibility="private",747 )748 gr.on(749 [demo.load, project_dd.change],750 fn=refresh_runs,751 inputs=[project_dd, run_tb, run_selection_state, selected_runs_from_url],752 outputs=[run_cb, run_tb, run_selection_state],753 show_progress="hidden",754 queue=False,755 api_visibility="private",756 ).then(757 fn=update_x_axis_choices,758 inputs=[project_dd, run_selection_state],759 outputs=x_axis_dd,760 show_progress="hidden",761 queue=False,762 api_visibility="private",763 ).then(764 fn=generate_embed,765 inputs=[project_dd, metric_filter_tb, run_selection_state],766 outputs=[embed_code],767 show_progress="hidden",768 api_visibility="private",769 queue=False,770 ).then(771 fns.update_navbar_value,772 inputs=[project_dd],773 outputs=[navbar],774 show_progress="hidden",775 api_visibility="private",776 queue=False,777 ).then(778 fn=fns.get_group_by_fields,779 inputs=[project_dd],780 outputs=[run_group_by_dd],781 show_progress="hidden",782 api_visibility="private",783 queue=False,784 )785 786 gr.on(787 [run_cb.input],788 fn=update_x_axis_choices,789 inputs=[project_dd, run_selection_state],790 outputs=x_axis_dd,791 show_progress="hidden",792 queue=False,793 api_visibility="private",794 )795 gr.on(796 [metric_filter_tb.change, run_cb.change],797 fn=generate_embed,798 inputs=[project_dd, metric_filter_tb, run_selection_state],799 outputs=embed_code,800 show_progress="hidden",801 api_visibility="private",802 queue=False,803 )804 805 def toggle_group_view(group_by_dd):806 return (807 gr.CheckboxGroup(visible=not bool(group_by_dd)),808 gr.Group(visible=bool(group_by_dd)),809 )810 811 gr.on(812 [run_group_by_dd.change],813 fn=toggle_group_view,814 inputs=[run_group_by_dd],815 outputs=[run_cb, grouped_runs_panel],816 show_progress="hidden",817 api_visibility="private",818 queue=False,819 )820 821 realtime_cb.change(822 fn=toggle_timer,823 inputs=realtime_cb,824 outputs=timer,825 api_visibility="private",826 queue=False,827 )828 run_cb.input(829 fn=fns.handle_run_checkbox_change,830 inputs=[run_cb, run_selection_state],831 outputs=run_selection_state,832 api_visibility="private",833 queue=False,834 ).then(835 fn=generate_embed,836 inputs=[project_dd, metric_filter_tb, run_selection_state],837 outputs=embed_code,838 show_progress="hidden",839 api_visibility="private",840 queue=False,841 )842 run_tb.input(843 fn=refresh_runs,844 inputs=[project_dd, run_tb, run_selection_state],845 outputs=[run_cb, run_tb, run_selection_state],846 api_visibility="private",847 queue=False,848 show_progress="hidden",849 )850 851 gr.api(852 fn=upload_db_to_space,853 api_name="upload_db_to_space",854 )855 gr.api(856 fn=bulk_upload_media,857 api_name="bulk_upload_media",858 )859 gr.api(860 fn=log,861 api_name="log",862 )863 gr.api(864 fn=bulk_log,865 api_name="bulk_log",866 )867 gr.api(868 fn=get_metric_values,869 api_name="get_metric_values",870 )871 gr.api(872 fn=get_runs_for_project,873 api_name="get_runs_for_project",874 )875 gr.api(876 fn=get_metrics_for_run,877 api_name="get_metrics_for_run",878 )879 gr.api(880 fn=get_all_projects,881 api_name="get_all_projects",882 )883 gr.api(884 fn=get_project_summary,885 api_name="get_project_summary",886 )887 gr.api(888 fn=get_run_summary,889 api_name="get_run_summary",890 )891 892 last_steps = gr.State({})893 894 def update_x_lim(select_data: gr.SelectData):895 return select_data.index896 897 def update_last_steps(project):898 """Check the last step for each run to detect when new data is available."""899 if not project:900 return {}901 return SQLiteStorage.get_max_steps_for_runs(project)902 903 timer.tick(904 fn=update_last_steps,905 inputs=[project_dd],906 outputs=last_steps,907 show_progress="hidden",908 api_visibility="private",909 )910 911 @gr.render(912 triggers=[913 demo.load,914 run_cb.change,915 last_steps.change,916 smoothing_slider.change,917 x_lim.change,918 x_axis_dd.change,919 log_scale_x_cb.change,920 log_scale_y_cb.change,921 metric_filter_tb.change,922 ],923 inputs=[924 project_dd,925 run_cb,926 smoothing_slider,927 metrics_subset,928 x_lim,929 x_axis_dd,930 log_scale_x_cb,931 log_scale_y_cb,932 metric_filter_tb,933 run_selection_state,934 ],935 show_progress="hidden",936 queue=False,937 )938 def update_dashboard(939 project,940 runs,941 smoothing_granularity,942 metrics_subset,943 x_lim_value,944 x_axis,945 log_scale_x,946 log_scale_y,947 metric_filter,948 selection,949 ):950 dfs = []951 original_runs = runs.copy()952 953 for run in runs:954 df, _ = load_run_data(955 project, run, smoothing_granularity, x_axis, log_scale_x, log_scale_y956 )957 if df is not None:958 dfs.append(df)959 960 if dfs:961 if smoothing_granularity > 0:962 original_dfs = []963 smoothed_dfs = []964 for df in dfs:965 original_data = df[df["data_type"] == "original"]966 smoothed_data = df[df["data_type"] == "smoothed"]967 if not original_data.empty:968 original_dfs.append(original_data)969 if not smoothed_data.empty:970 smoothed_dfs.append(smoothed_data)971 972 all_dfs = original_dfs + smoothed_dfs973 master_df = (974 pd.concat(all_dfs, ignore_index=True) if all_dfs else pd.DataFrame()975 )976 977 else:978 master_df = pd.concat(dfs, ignore_index=True)979 else:980 master_df = pd.DataFrame()981 982 if master_df.empty:983 if not SQLiteStorage.get_projects():984 if space_id := utils.get_space():985 gr.Markdown(INSTRUCTIONS_SPACES.format(space_id))986 else:987 gr.Markdown(INSTRUCTIONS_LOCAL)988 else:989 gr.Markdown("*Waiting for runs to appear...*")990 return991 992 x_column = "step"993 if dfs and not dfs[0].empty and "x_axis" in dfs[0].columns:994 x_column = dfs[0]["x_axis"].iloc[0]995 996 numeric_cols = master_df.select_dtypes(include="number").columns997 numeric_cols = [c for c in numeric_cols if c not in utils.RESERVED_KEYS]998 if x_column and x_column in numeric_cols:999 numeric_cols.remove(x_column)1000 1001 if metrics_subset:1002 numeric_cols = [c for c in numeric_cols if c in metrics_subset]1003 1004 if metric_filter and metric_filter.strip():1005 numeric_cols = filter_metrics_by_regex(list(numeric_cols), metric_filter)1006 1007 ordered_groups, nested_metric_groups = utils.order_metrics_by_plot_preference(1008 list(numeric_cols)1009 )1010 all_runs = selection.choices if selection else original_runs1011 color_map = utils.get_color_mapping(all_runs, smoothing_granularity > 0)1012 1013 metric_idx = 01014 for group_name in ordered_groups:1015 group_data = nested_metric_groups[group_name]1016 1017 total_plot_count = sum(1018 11019 for m in group_data["direct_metrics"]1020 if not master_df.dropna(subset=[m]).empty1021 ) + sum(1022 sum(1 for m in metrics if not master_df.dropna(subset=[m]).empty)1023 for metrics in group_data["subgroups"].values()1024 )1025 group_label = (1026 f"{group_name} ({total_plot_count})"1027 if total_plot_count > 01028 else group_name1029 )1030 1031 with gr.Accordion(1032 label=group_label,1033 open=True,1034 key=f"accordion-{group_name}",1035 preserved_by_key=["value", "open"],1036 ):1037 if group_data["direct_metrics"]:1038 with gr.Draggable(1039 key=f"row-{group_name}-direct", orientation="row"1040 ):1041 for metric_name in group_data["direct_metrics"]:1042 metric_df = master_df.dropna(subset=[metric_name])1043 color = "run" if "run" in metric_df.columns else None1044 downsampled_df, updated_x_lim = utils.downsample(1045 metric_df,1046 x_column,1047 metric_name,1048 color,1049 x_lim_value,1050 )1051 if not metric_df.empty:1052 plot = gr.LinePlot(1053 downsampled_df,1054 x=x_column,1055 y=metric_name,1056 y_title=metric_name.split("/")[-1],1057 color=color,1058 color_map=color_map,1059 colors_in_legend=original_runs,1060 title=metric_name,1061 key=f"plot-{metric_idx}",1062 preserved_by_key=None,1063 buttons=["fullscreen", "export"],1064 x_lim=updated_x_lim,1065 min_width=400,1066 )1067 plot.select(1068 update_x_lim,1069 outputs=x_lim,1070 key=f"select-{metric_idx}",1071 )1072 plot.double_click(1073 lambda: None,1074 outputs=x_lim,1075 key=f"double-{metric_idx}",1076 )1077 metric_idx += 11078 1079 if group_data["subgroups"]:1080 for subgroup_name in sorted(group_data["subgroups"].keys()):1081 subgroup_metrics = group_data["subgroups"][subgroup_name]1082 1083 subgroup_plot_count = sum(1084 11085 for m in subgroup_metrics1086 if not master_df.dropna(subset=[m]).empty1087 )1088 subgroup_label = (1089 f"{subgroup_name} ({subgroup_plot_count})"1090 if subgroup_plot_count > 01091 else subgroup_name1092 )1093 1094 with gr.Accordion(1095 label=subgroup_label,1096 open=True,1097 key=f"accordion-{group_name}-{subgroup_name}",1098 preserved_by_key=["value", "open"],1099 ):1100 with gr.Draggable(1101 key=f"row-{group_name}-{subgroup_name}",1102 orientation="row",1103 ):1104 for metric_name in subgroup_metrics:1105 metric_df = master_df.dropna(subset=[metric_name])1106 color = (1107 "run" if "run" in metric_df.columns else None1108 )1109 downsampled_df, updated_x_lim = utils.downsample(1110 metric_df,1111 x_column,1112 metric_name,1113 color,1114 x_lim_value,1115 )1116 if not metric_df.empty:1117 plot = gr.LinePlot(1118 downsampled_df,1119 x=x_column,1120 y=metric_name,1121 y_title=metric_name.split("/")[-1],1122 color=color,1123 color_map=color_map,1124 colors_in_legend=original_runs,1125 title=metric_name,1126 key=f"plot-{metric_idx}",1127 preserved_by_key=None,1128 buttons=["fullscreen", "export"],1129 x_lim=updated_x_lim,1130 min_width=400,1131 )1132 plot.select(1133 update_x_lim,1134 outputs=x_lim,1135 key=f"select-{metric_idx}",1136 )1137 plot.double_click(1138 lambda: None,1139 outputs=x_lim,1140 key=f"double-{metric_idx}",1141 )1142 metric_idx += 11143 1144 with grouped_runs_panel:1145 1146 @gr.render(1147 triggers=[1148 demo.load,1149 project_dd.change,1150 run_group_by_dd.change,1151 run_tb.input,1152 run_selection_state.change,1153 last_steps.change,1154 ],1155 inputs=[project_dd, run_group_by_dd, run_tb, run_selection_state],1156 show_progress="hidden",1157 queue=False,1158 )1159 def render_grouped_runs(project, group_key, filter_text, selection):1160 if not group_key:1161 return1162 selection = selection or RunSelection()1163 groups = fns.group_runs_by_config(project, group_key, filter_text)1164 1165 for label, runs in groups.items():1166 ordered_current = utils.ordered_subset(runs, selection.selected)1167 1168 with gr.Group():1169 show_group_cb = gr.Checkbox(1170 label="Show/Hide",1171 value=bool(ordered_current),1172 key=f"show-cb-{group_key}-{label}",1173 preserved_by_key=["value"],1174 )1175 1176 with gr.Accordion(1177 f"{label} ({len(runs)})",1178 open=False,1179 key=f"accordion-{group_key}-{label}",1180 preserved_by_key=["open"],1181 ):1182 color_palette = utils.get_color_palette()1183 choice_indices = {1184 run: i for i, run in enumerate(selection.choices)1185 }1186 colors = [1187 color_palette[1188 choice_indices.get(run, 0) % len(color_palette)1189 ]1190 for run in runs1191 ]1192 group_cb = ColoredCheckboxGroup(1193 choices=runs,1194 value=ordered_current,1195 colors=colors,1196 label=f"Runs ({len(runs)})",1197 key=f"group-cb-{group_key}-{label}",1198 preserved_by_key=None,1199 )1200 