Spaces:
Running
Running
| from html import escape | |
| from pathlib import Path | |
| import base64 | |
| import random | |
| import gradio as gr | |
| import pandas as pd | |
| import plotly.graph_objects as go | |
| from model_display import display_model_name | |
| _LOGO_PATH = Path(__file__).resolve().parent / "pruna-logo.png" | |
| _LOGO_DATA_URI = ( | |
| "data:image/png;base64," | |
| + base64.b64encode(_LOGO_PATH.read_bytes()).decode("ascii") | |
| if _LOGO_PATH.exists() | |
| else "https://playground.pruna.ai/logo.svg" | |
| ) | |
| MAX_COMPARE_MODELS = 4 | |
| DEFAULT_COMPARE_PROMPTS = 3 | |
| MAX_COMPARE_PROMPTS = 8 | |
| MAX_PARETO_METRICS = 8 | |
| _PARETO_SLOT_COUNT = 1 + MAX_PARETO_METRICS * 8 | |
| _PARETO_PRICE_COLUMN = "Price / Image (USD)" | |
| _PARETO_TIME_COLUMN = "Min Generation Time (s)" | |
| TAB_LEADERBOARDS = "leaderboards" | |
| TAB_PARETO = "pareto" | |
| TAB_SAMPLES = "samples" | |
| TAB_ABOUT = "about" | |
| _MODEL_CHOICES_CACHE = {} | |
| _VIEW_EVENTS = { | |
| "show_progress": "hidden", | |
| "trigger_mode": "always_last", | |
| "concurrency_id": "workspace-views", | |
| "concurrency_limit": 1, | |
| } | |
| ABOUT_OVERVIEW_CONTENT = """ | |
| # About P-Bench | |
| P-Bench compares **text-to-image models**, including optimized or accelerated | |
| endpoints, on **quality, speed, and price**. Each view is a **dataset** scored | |
| with a **metric**, written as `Dataset | Metric`. There is no single score | |
| across P-Bench. | |
| ## How to read it | |
| 1. Pick a **dataset** and a **metric**. | |
| 2. **Leaderboards**: ranked by that metric. Price and generation time sit in | |
| the same table when the source publishes them. | |
| 3. **Pareto plots**: mark models that are not beaten on both higher score | |
| and lower price (or time). Only datasets with price or generation time | |
| can open this tab (not Arena AI). | |
| 4. **Samples**: the same prompts, side by side. Only for datasets we | |
| generated (Qwen Image Dataset and OneIG Alignment Dataset). | |
| ## How a score is made | |
| On **Qwen Image Dataset** and **OneIG Alignment Dataset**: | |
| 1. Each endpoint is given the same prompt suite. | |
| 2. It generates one image per prompt when the run succeeds. Not every model | |
| has every prompt or every metric. | |
| 3. Quality is scored automatically (OneIG alignment, P-Judger) and, where | |
| available, by human preference (Datapoint Elo, Rapidata Elo). | |
| 4. Price per image and generation time are joined from the evaluation table. | |
| **Artificial Analysis** and **Arena AI** are external leaderboards. We import | |
| their published Elos (and Artificial Analysis price). We do not run their | |
| prompt suites, so samples are not shown. | |
| ## Current datasets | |
| ### Qwen Image Dataset | |
| 100 prompts from the 1,000-prompt Qwen Image Bench set, sampled for coverage | |
| across its fine-grained (L3) categories. Metrics include Datapoint Elo, | |
| Rapidata Elo, and P-Judger. Samples are available. | |
| ### OneIG Alignment Dataset | |
| Prompt-image **alignment** on anime / stylization (100), human / portrait | |
| (100), and general object prompts (99). This is the alignment slice of | |
| OneIG, not the full suite. Alignment Overall is the mean of the category | |
| scores that exist for that row. Also includes Datapoint category Elo, | |
| Rapidata Elo, and P-Judger. Samples are available. | |
| ### Artificial Analysis Dataset | |
| External text-to-image Elo and price per image from Artificial Analysis. | |
| Their prompt set is private, so samples are not shown. Pareto plots use | |
| price vs score only. | |
| ### Arena AI Dataset | |
| External text-to-image Elo (overall and category) from Arena AI. Their | |
| prompt set is private, so samples are not shown. Price and generation time | |
| are not in this export, so Pareto plots are unavailable. | |
| """ | |
| ABOUT_DETAILS_CONTENT = """ | |
| # Metrics & limits | |
| ## Metrics | |
| - **Alignment**: how well the image matches the prompt (OneIG category scores). | |
| - **P-Judger**: Pruna's automatic judger. A private scoring model that rates | |
| each (prompt, image) pair in one pass and returns an absolute quality score, | |
| not a pairwise Elo. It is not a text-to-image model on the leaderboard. Use | |
| it alongside Datapoint and Rapidata Elo, not instead of them. | |
| - **Datapoint Elo**: human-preference Elo from Datapoint pairwise comparisons. | |
| - **Rapidata Elo**: human-preference Elo from Rapidata pairwise comparisons. | |
| Rapidata rejects prompts over 400 characters, so this Elo is on a subset | |
| of each suite (see Setup). Rapidata is not a dataset. | |
| - **Artificial Analysis Elo**: Elo published by Artificial Analysis on their | |
| own dataset. | |
| - **Arena Elo**: Elo published by Arena AI on their own dataset, plus | |
| category Elos (branding, 3D, cartoon/anime, photorealistic, art, portraits, | |
| text rendering). | |
| - **Generation time**: median and minimum generation time in seconds, as | |
| reported in the evaluation table. This is not a p95, and we do not state | |
| warm vs cold or concurrent load. Not available for Arena AI. | |
| - **Price**: USD per image in the evaluation table. We do not state list | |
| price vs amount paid, or whether failed generations are included. Not | |
| available for Arena AI. | |
| Scores from different datasets or metrics are **not interchangeable**. A high | |
| OneIG alignment score is not the same quantity as a Datapoint Elo. Compare | |
| models *within* a Dataset | Metric view. | |
| ## Setup | |
| - **Evaluation window:** July–August 2026. | |
| - **Update policy:** numbers come from evaluation snapshots in the tables, | |
| not a live API poll. | |
| - **Prompt counts:** OneIG Alignment uses 100 anime, 100 human, and 99 object | |
| prompts (299 total). Qwen Image Dataset uses 100 prompts sampled from the | |
| 1,000-prompt pool for roughly even coverage of its fine-grained (L3) | |
| categories. Artificial Analysis and Arena AI use their own private prompt | |
| sets. | |
| - **Generation (Qwen and OneIG):** one image per prompt per endpoint when | |
| the run exists. Default resolution is 1024×1024. Exceptions: FLUX 1.1 Pro | |
| Ultra at 2K, FLUX 2 Flex at 1008×1008, and any endpoint labeled 2K. The | |
| seed is derived from the prompt, so every model gets the same seed for the | |
| same prompt. Steps, CFG, prompt rewrite, and safety filters follow each | |
| endpoint's default. This does not describe Artificial Analysis or Arena AI. | |
| - **Datapoint (Qwen and OneIG):** every model pair is compared on every | |
| prompt, with 10 votes per battle. | |
| - **Rapidata (Qwen and OneIG):** prompts longer than 400 characters are | |
| dropped, leaving 212 OneIG prompts and 85 Qwen Image Dataset prompts. 4 | |
| votes per pair; about 26,000 votes on OneIG and 35,000 on Qwen Image | |
| Dataset. | |
| ## Limits | |
| - Empty cells mean that track was not run or not reported for that model. | |
| - Artificial Analysis and Arena AI samples, prompts, and (for Arena) price | |
| or latency are not available to P-Bench. | |
| - Rapidata Elo is not on the full prompt suite, so it is not directly | |
| comparable to Datapoint Elo even on the same dataset. | |
| - Elo ratings can shift when the comparison pool changes: treat them as | |
| relative rankings for the snapshot, not absolute constants. | |
| - Close scores can be a tie in practice; the table does not show confidence | |
| intervals. | |
| - Optimized or accelerated endpoints (when labeled) may differ from the base | |
| model publisher's default serving stack. | |
| - APIs and prices drift; a rank is for the snapshot, not a guarantee. | |
| Built by [Pruna AI](https://www.pruna.ai/). If you have questions, reach out | |
| to us on [Discord](https://discord.gg/JFQmtFKCjd). | |
| """ | |
| FOOTER_CONTENT = """ | |
| <div class="page-footer"> | |
| <h3>Join the Pruna AI community</h3> | |
| <div class="community-footer-links"> | |
| <a rel="nofollow" href="https://x.com/PrunaAI" target="_blank">X</a> | |
| <a rel="nofollow" href="https://www.linkedin.com/company/pruna-ai" target="_blank">LinkedIn</a> | |
| <a rel="nofollow" href="https://www.instagram.com/pruna.ai/" target="_blank">Instagram</a> | |
| <a rel="nofollow" href="https://discord.gg/JFQmtFKCjd" target="_blank">Discord</a> | |
| <a rel="nofollow" href="https://github.com/PrunaAI/pruna" target="_blank">GitHub</a> | |
| <a rel="nofollow" href="https://www.pruna.ai/" target="_blank">pruna.ai</a> | |
| </div> | |
| </div> | |
| """ | |
| CITATION_CONTENT = """ | |
| ```bibtex | |
| @misc{P-Bench, | |
| title={P-Bench: A Leaderboard for Text-to-Image Models}, | |
| author={PrunaAI}, | |
| year={2026}, | |
| howpublished={\\url{https://huggingface.co/spaces/PrunaAI/P-Bench}} | |
| } | |
| ``` | |
| """ | |
| def render_header(): | |
| gr.HTML( | |
| f""" | |
| <header class="app-header"> | |
| <div class="app-header-bar"> | |
| <div class="app-header-brand"> | |
| <img class="app-header-logo" src="{_LOGO_DATA_URI}" alt="" /> | |
| <h1>P-Bench</h1> | |
| </div> | |
| <button type="button" class="theme-toggle" data-mode="dark" aria-label="Switch to light mode" title="Switch to light mode"> | |
| <svg class="theme-icon-sun" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" aria-hidden="true"> | |
| <circle cx="12" cy="12" r="4"></circle> | |
| <path d="M12 2v2M12 20v2M4.93 4.93l1.41 1.41M17.66 17.66l1.41 1.41M2 12h2M20 12h2M4.93 19.07l1.41-1.41M17.66 6.34l1.41-1.41"></path> | |
| </svg> | |
| <svg class="theme-icon-moon" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" aria-hidden="true"> | |
| <path d="M21 14.5A8.5 8.5 0 1 1 9.5 3 7 7 0 0 0 21 14.5z"></path> | |
| </svg> | |
| </button> | |
| </div> | |
| <p class="app-header-tagline">Compare text-to-image models on quality, speed, and price</p> | |
| </header> | |
| """, | |
| padding=False, | |
| elem_classes="app-header-host", | |
| ) | |
| def _item(items, item_id): | |
| for item in items: | |
| if item["id"] == item_id: | |
| return item | |
| return items[0] if items else None | |
| def _dataset_choices(datasets, *, require_samples=False, require_pareto=False): | |
| return [ | |
| (dataset["name"], dataset["id"]) | |
| for dataset in datasets | |
| if (not require_samples or dataset.get("samples")) | |
| and (not require_pareto or _dataset_has_pareto(datasets, dataset["id"])) | |
| ] | |
| def _dataset_has_samples(datasets, dataset_id): | |
| dataset = _item(datasets, dataset_id) | |
| return bool(dataset and dataset.get("samples")) | |
| def _dataset_has_pareto(datasets, dataset_id): | |
| dataset = _item(datasets, dataset_id) | |
| columns = getattr(dataset.get("data") if dataset else None, "columns", []) | |
| return _PARETO_PRICE_COLUMN in columns or _PARETO_TIME_COLUMN in columns | |
| def _dataset_dropdown_update(datasets, tab, dataset_id): | |
| """Limit the dataset list to what the current tab can show.""" | |
| return gr.update( | |
| choices=_dataset_choices( | |
| datasets, | |
| require_samples=tab == TAB_SAMPLES | |
| and _dataset_has_samples(datasets, dataset_id), | |
| require_pareto=tab == TAB_PARETO | |
| and _dataset_has_pareto(datasets, dataset_id), | |
| ), | |
| value=dataset_id, | |
| ) | |
| ALL_METRICS_ID = "__all__" | |
| def _metric_choices(datasets, metrics, dataset_id): | |
| dataset = _item(datasets, dataset_id) | |
| if not dataset: | |
| return [] | |
| allowed = set(dataset.get("metric_ids") or []) | |
| data = dataset.get("data") | |
| columns = getattr(data, "columns", []) | |
| return [ | |
| (_display_label(metric["column"]), metric["id"]) | |
| for metric in metrics | |
| if metric["id"] in allowed and metric["column"] in columns | |
| ] | |
| def _metric_dropdown_choices(datasets, metrics, dataset_id): | |
| choices = _metric_choices(datasets, metrics, dataset_id) | |
| if not choices: | |
| return [] | |
| return [("Select all", ALL_METRICS_ID)] + choices | |
| def _normalize_metric_ids(metric_id): | |
| if metric_id is None or metric_id == "": | |
| return [] | |
| if isinstance(metric_id, (list, tuple)): | |
| return [item for item in metric_id if item] | |
| return [metric_id] | |
| def _coerce_metric(datasets, metrics, dataset_id, metric_id): | |
| """Valid metric ids for this dataset. Empty means all metrics.""" | |
| wanted = _normalize_metric_ids(metric_id) | |
| valid_ids = [choice[1] for choice in _metric_choices(datasets, metrics, dataset_id)] | |
| valid = set(valid_ids) | |
| if ALL_METRICS_ID in wanted: | |
| return list(valid_ids) | |
| return [item for item in wanted if item in valid] | |
| def _metric_dropdown_value(metric_id): | |
| return [ | |
| item | |
| for item in _normalize_metric_ids(metric_id) | |
| if item != ALL_METRICS_ID | |
| ] | |
| def _model_choices(datasets, dataset_id): | |
| cached = _MODEL_CHOICES_CACHE.get(dataset_id) | |
| if cached is not None: | |
| return cached | |
| dataset = _item(datasets, dataset_id) | |
| data = dataset.get("data") if dataset else None | |
| if data is None or "Model" not in getattr(data, "columns", []): | |
| _MODEL_CHOICES_CACHE[dataset_id] = [] | |
| return [] | |
| models = data["Model"].dropna().astype(str).unique().tolist() | |
| # (label, value) so the UI shows the shared name but filters on the raw id. | |
| choices = sorted( | |
| ((display_model_name(model), model) for model in models), | |
| key=lambda item: item[0].casefold(), | |
| ) | |
| _MODEL_CHOICES_CACHE[dataset_id] = choices | |
| return choices | |
| def _model_choice_values(choices): | |
| return [choice[1] if isinstance(choice, tuple) else choice for choice in choices] | |
| def _metric_columns(datasets, metrics, dataset_id): | |
| dataset = _item(datasets, dataset_id) | |
| if not dataset: | |
| return [] | |
| data = dataset.get("data") | |
| columns = getattr(data, "columns", []) | |
| allowed = set(dataset.get("metric_ids") or []) | |
| return [ | |
| metric["column"] | |
| for metric in metrics | |
| if metric["id"] in allowed and metric["column"] in columns | |
| ] | |
| _LEADERBOARD_IDENTITY_COLUMNS = [ | |
| "Model", | |
| "Platform", | |
| "Endpoint Owner", | |
| "Optimized", | |
| ] | |
| _LEADERBOARD_META_COLUMNS = [ | |
| "Median Generation Time (s)", | |
| "Min Generation Time (s)", | |
| "Price / Image (USD)", | |
| "Evaluation Date (UTC)", | |
| "Date", | |
| ] | |
| def _columns_for_metric(dataset, score_columns): | |
| """When metrics are selected, show identity + those scores + time/price.""" | |
| available = list(getattr(dataset.get("data"), "columns", [])) or list( | |
| dataset.get("columns") or [] | |
| ) | |
| identity = [column for column in _LEADERBOARD_IDENTITY_COLUMNS if column in available] | |
| meta = [column for column in _LEADERBOARD_META_COLUMNS if column in available] | |
| scores = [column for column in (score_columns or []) if column in available] | |
| if scores: | |
| return [*identity, *scores, *meta] | |
| return [column for column in (dataset.get("columns") or available) if column != "URL"] | |
| def resolve_view(datasets, metrics, dataset_id, metric_id): | |
| dataset = _item(datasets, dataset_id) | |
| if not dataset: | |
| return None | |
| data_columns = getattr(dataset.get("data"), "columns", []) | |
| score_columns = [] | |
| for metric_key in _coerce_metric(datasets, metrics, dataset_id, metric_id): | |
| metric = _item(metrics, metric_key) | |
| if metric and metric["column"] in data_columns: | |
| score_columns.append(metric["column"]) | |
| if score_columns: | |
| columns = _columns_for_metric(dataset, score_columns) | |
| else: | |
| columns = [ | |
| column | |
| for column in (dataset.get("columns") or []) | |
| if column != "URL" | |
| ] | |
| score_columns = _metric_columns(datasets, metrics, dataset_id) | |
| return { | |
| "data": dataset["data"], | |
| "columns": columns, | |
| "score_column": score_columns[0] if score_columns else None, | |
| "score_columns": score_columns, | |
| "samples": dataset.get("samples"), | |
| "note": dataset.get("note"), | |
| } | |
| def _format_leaderboard_cell(column, value): | |
| if pd.isna(value) or value is None or value == "": | |
| return "-" | |
| label = str(column).lower() | |
| if label == "rank": | |
| return str(int(value)) | |
| if "price" in label: | |
| return _format_price(value) | |
| if "time" in label or "generation" in label: | |
| try: | |
| return f"{float(value):.2f}" | |
| except (TypeError, ValueError): | |
| return escape(str(value)) | |
| if label in {"model", "platform", "endpoint owner", "optimized"}: | |
| if label == "model": | |
| return escape(display_model_name(value)) | |
| return escape(str(value)) | |
| try: | |
| number = float(value) | |
| except (TypeError, ValueError): | |
| return escape(str(value)) | |
| if abs(number) >= 100: | |
| return f"{number:.1f}" | |
| return f"{number:.4f}".rstrip("0").rstrip(".") | |
| def _leaderboard_sort_value(column, value): | |
| """Raw value used by client-side column sorting.""" | |
| if pd.isna(value) or value is None or value == "": | |
| return "" | |
| label = str(column).lower() | |
| if label in {"model", "platform", "endpoint owner", "optimized", "date"}: | |
| return str(value).casefold() | |
| try: | |
| return f"{float(value):.10g}" | |
| except (TypeError, ValueError): | |
| return str(value).casefold() | |
| def _leaderboard_sort_type(column): | |
| label = str(column).lower() | |
| if label in {"model", "platform", "endpoint owner", "optimized", "date"}: | |
| return "text" | |
| return "number" | |
| def _leaderboard_cell_class(column): | |
| if column == "Rank": | |
| return "rank" | |
| if column == "Model": | |
| return "model-cell" | |
| if _leaderboard_sort_type(column) == "text": | |
| return "text-cell" | |
| return "metric-score" | |
| def _assign_leaderboard_ranks(data, overall_column): | |
| """Rank the full table by the selected metric. Filters keep these numbers.""" | |
| if data is None: | |
| return data | |
| ranked = data.copy() | |
| if "Rank" in ranked.columns: | |
| ranked = ranked.drop(columns=["Rank"]) | |
| if overall_column and overall_column in ranked.columns: | |
| ranked = ranked.sort_values( | |
| overall_column, ascending=False, na_position="last" | |
| ) | |
| ranked = ranked.reset_index(drop=True) | |
| ranked.insert(0, "Rank", ranked.index + 1) | |
| return ranked | |
| def _leaderboard_html(data, columns): | |
| leaderboard = _leaderboard_dataframe(data, columns) | |
| if leaderboard.empty: | |
| return ( | |
| '<div class="ranking-table-scroll">' | |
| '<div class="ranking-table empty-state">No models match the current filters.</div>' | |
| "</div>" | |
| ) | |
| header_cells = [] | |
| for index, column in enumerate(leaderboard.columns): | |
| sort_type = _leaderboard_sort_type(column) | |
| cell_class = _leaderboard_cell_class(column) | |
| header_cells.append( | |
| f'<th class="sortable-col {cell_class}" data-col="{index}" ' | |
| f'data-type="{sort_type}" title="Sort by {escape(str(column))}">' | |
| f"{escape(str(column))}</th>" | |
| ) | |
| body_rows = [] | |
| for _, row in leaderboard.iterrows(): | |
| cells = [] | |
| for column in leaderboard.columns: | |
| css = _leaderboard_cell_class(column) | |
| sort_value = escape(_leaderboard_sort_value(column, row[column]), quote=True) | |
| if column == "Model": | |
| cells.append( | |
| f'<td class="{css}" data-sort-value="{sort_value}">' | |
| f"<strong>{_format_leaderboard_cell(column, row[column])}</strong></td>" | |
| ) | |
| else: | |
| cells.append( | |
| f'<td class="{css}" data-sort-value="{sort_value}">' | |
| f"{_format_leaderboard_cell(column, row[column])}</td>" | |
| ) | |
| body_rows.append(f"<tr>{''.join(cells)}</tr>") | |
| return f""" | |
| <div class="ranking-table-scroll"> | |
| <div class="ranking-table"> | |
| <table> | |
| <thead><tr>{''.join(header_cells)}</tr></thead> | |
| <tbody>{''.join(body_rows)}</tbody> | |
| </table> | |
| </div> | |
| </div> | |
| """ | |
| def _filter_choices(data, column): | |
| if data is None or column not in data.columns: | |
| return [] | |
| return sorted(data[column].dropna().astype(str).unique().tolist()) | |
| def _filter_leaderboard(data, platform, owner, optimized, models=None): | |
| filtered = data.copy() | |
| if models: | |
| if "Model" in filtered.columns: | |
| filtered = filtered[filtered["Model"].astype(str).isin(models)] | |
| for column, values in [ | |
| ("Platform", platform), | |
| ("Endpoint Owner", owner), | |
| ("Optimized", optimized), | |
| ]: | |
| if values and column in filtered.columns: | |
| filtered = filtered[filtered[column].astype(str).isin(values)] | |
| return filtered | |
| def _leaderboard_dataframe(data, columns): | |
| ordered = [] | |
| if "Rank" in data.columns: | |
| ordered.append("Rank") | |
| for column in columns: | |
| if ( | |
| column in data.columns | |
| and column not in {"URL", "Rank"} | |
| and column not in ordered | |
| ): | |
| ordered.append(column) | |
| return data[ordered].rename(columns=_display_label) | |
| def _display_label(column): | |
| labels = { | |
| "OneIG Overall Score": "Overall", | |
| "OneIG (Anime Alignment)": "Anime", | |
| "OneIG (Human Alignment)": "Human", | |
| "OneIG (Object Alignment)": "Object", | |
| "OneIG Anime Elo": "Anime Elo (Datapoint)", | |
| "OneIG Human Elo": "Human Elo (Datapoint)", | |
| "OneIG Object Elo": "Object Elo (Datapoint)", | |
| "P-Judge Overall": "P-Judger (Pruna)", | |
| "Datapoint Elo": "Datapoint Elo", | |
| "Rapidata Elo": "Rapidata Elo", | |
| "Artificial Analysis Elo": "Artificial Analysis Elo", | |
| "Arena Elo": "Overall Elo", | |
| "Arena Branding / Commercial Elo": "Branding / Commercial", | |
| "Arena 3D Imaging Elo": "3D Imaging", | |
| "Arena Cartoon / Anime Elo": "Cartoon / Anime", | |
| "Arena Photorealistic Elo": "Photorealistic", | |
| "Arena Art Elo": "Art", | |
| "Arena Portraits Elo": "Portraits", | |
| "Arena Text Rendering Elo": "Text Rendering", | |
| "Raw Win Rate": "Raw win rate", | |
| "Median Generation Time (s)": "Median generation time", | |
| "Min Generation Time (s)": "Min generation time", | |
| "Price / Image (USD)": "Price per image", | |
| "Evaluation Date (UTC)": "Date", | |
| "Date": "Date", | |
| } | |
| return labels.get(column, column) | |
| def _format_price(value): | |
| return "-" if pd.isna(value) or value is None else f"${float(value):.3f}" | |
| def _pareto_frontier_mask(x_values, scores): | |
| """True for non-dominated points when maximizing score and minimizing x.""" | |
| n = len(x_values) | |
| mask = [True] * n | |
| for i in range(n): | |
| for j in range(n): | |
| if i == j: | |
| continue | |
| better_or_equal = x_values[j] <= x_values[i] and scores[j] >= scores[i] | |
| strictly_better = x_values[j] < x_values[i] or scores[j] > scores[i] | |
| if better_or_equal and strictly_better: | |
| mask[i] = False | |
| break | |
| return mask | |
| def _pareto_unavailable_html(message): | |
| return ( | |
| "<p class='pareto-note-copy'>" | |
| f"{escape(message)}" | |
| "</p>" | |
| ) | |
| def _pareto_note_update(message): | |
| if message: | |
| return gr.update( | |
| value=_pareto_unavailable_html(message), | |
| visible=True, | |
| ) | |
| return gr.update(value="", visible=False) | |
| def _pareto_plot_update(fig): | |
| if fig is not None: | |
| return gr.update(value=fig, visible=True) | |
| return gr.update(value=None, visible=False) | |
| def _skip_all(count): | |
| return tuple(gr.skip() for _ in range(count)) | |
| def _pareto_skip_updates(): | |
| return _skip_all(_PARETO_SLOT_COUNT) | |
| def _selection_key(dataset_id, metric_id, models): | |
| return ( | |
| dataset_id, | |
| tuple(_normalize_metric_ids(metric_id)), | |
| tuple(models or ()), | |
| ) | |
| def _applied_key(view_state): | |
| view_state = view_state or {} | |
| return _selection_key( | |
| view_state.get("dataset_id"), | |
| view_state.get("metric_id"), | |
| view_state.get("models"), | |
| ) | |
| def _build_pareto_figure( | |
| data, | |
| score_column, | |
| x_column, | |
| x_title, | |
| x_hover_prefix="", | |
| x_hover_suffix="", | |
| ): | |
| scatter = ( | |
| data[["Model", score_column, x_column]] | |
| .dropna() | |
| .copy() | |
| .reset_index(drop=True) | |
| ) | |
| if scatter.empty: | |
| return None | |
| x_values = scatter[x_column].astype(float).tolist() | |
| scores = scatter[score_column].astype(float).tolist() | |
| on_frontier = _pareto_frontier_mask(x_values, scores) | |
| dominated = scatter.loc[[not flag for flag in on_frontier]].copy() | |
| frontier = scatter.loc[on_frontier].sort_values(x_column).copy() | |
| if not dominated.empty: | |
| dominated["Model"] = dominated["Model"].map(display_model_name) | |
| if not frontier.empty: | |
| frontier["Model"] = frontier["Model"].map(display_model_name) | |
| hover = ( | |
| "<b>%{text}</b><br>" | |
| f"{escape(x_title)}: {x_hover_prefix}%{{x:.4f}}{x_hover_suffix}" | |
| "<br>Score: %{y:.4f}<extra></extra>" | |
| ) | |
| fig = go.Figure() | |
| if not dominated.empty: | |
| fig.add_trace( | |
| go.Scatter( | |
| x=dominated[x_column], | |
| y=dominated[score_column], | |
| mode="markers", | |
| name="Below frontier", | |
| text=dominated["Model"], | |
| hovertemplate=hover, | |
| marker={ | |
| "size": 9, | |
| "color": "#d8b4fe", | |
| "opacity": 0.8, | |
| "line": {"width": 0}, | |
| }, | |
| ) | |
| ) | |
| if not frontier.empty: | |
| fig.add_trace( | |
| go.Scatter( | |
| x=frontier[x_column], | |
| y=frontier[score_column], | |
| mode="lines+markers", | |
| name="On frontier", | |
| text=frontier["Model"], | |
| hovertemplate=hover, | |
| line={"color": "#69a45c", "width": 2.5}, | |
| marker={ | |
| "size": 12, | |
| "color": "#69a45c", | |
| "line": {"width": 1.5, "color": "#86c077"}, | |
| }, | |
| ) | |
| ) | |
| score_label = _display_label(score_column) | |
| fig.update_layout( | |
| title=None, | |
| xaxis_title=x_title, | |
| yaxis_title=score_label, | |
| autosize=True, | |
| height=420, | |
| margin={"l": 56, "r": 28, "t": 28, "b": 80}, | |
| legend={ | |
| "orientation": "h", | |
| "yanchor": "top", | |
| "y": -0.24, | |
| "xanchor": "center", | |
| "x": 0.5, | |
| "bgcolor": "rgba(0,0,0,0)", | |
| "font": {"color": "#d4d4d4", "size": 12}, | |
| }, | |
| plot_bgcolor="#1d1429", | |
| paper_bgcolor="#171021", | |
| font={"color": "#d4d4d4", "size": 13}, | |
| ) | |
| axis_font = {"color": "#fafafa", "size": 13} | |
| tick_font = {"color": "#a3a3a3", "size": 12} | |
| fig.update_xaxes( | |
| showgrid=True, | |
| gridcolor="rgba(74, 57, 98, 0.55)", | |
| zeroline=False, | |
| title_font=axis_font, | |
| tickfont=tick_font, | |
| color="#d4d4d4", | |
| ) | |
| fig.update_yaxes( | |
| showgrid=True, | |
| gridcolor="rgba(74, 57, 98, 0.55)", | |
| zeroline=False, | |
| title_font=axis_font, | |
| tickfont=tick_font, | |
| color="#d4d4d4", | |
| ) | |
| return fig | |
| def _pareto_axis(data, score_column, x_column, x_title, missing_message, empty_message, **hover): | |
| if x_column not in data.columns: | |
| return None, missing_message | |
| fig = _build_pareto_figure( | |
| data, | |
| score_column, | |
| x_column=x_column, | |
| x_title=x_title, | |
| **hover, | |
| ) | |
| if fig is None: | |
| return None, empty_message | |
| return fig, None | |
| def _pareto_pair(data, score_column): | |
| score_missing = "No score data is available for this metric." | |
| if data is None or not score_column or score_column not in data.columns: | |
| return None, score_missing, None, score_missing | |
| price_fig, price_message = _pareto_axis( | |
| data, | |
| score_column, | |
| _PARETO_PRICE_COLUMN, | |
| "Price per image (USD)", | |
| "Price per image isn't available for this dataset.", | |
| "No models have both a score and a price for this metric.", | |
| x_hover_prefix="$", | |
| ) | |
| time_fig, time_message = _pareto_axis( | |
| data, | |
| score_column, | |
| _PARETO_TIME_COLUMN, | |
| "Min generation time (s)", | |
| "Min generation time isn't available for this dataset.", | |
| "No models have both a score and a min generation time for this metric.", | |
| x_hover_suffix="s", | |
| ) | |
| return price_fig, price_message, time_fig, time_message | |
| def _pareto_dataset_message(data): | |
| has_price = data is not None and _PARETO_PRICE_COLUMN in data.columns | |
| has_time = data is not None and _PARETO_TIME_COLUMN in data.columns | |
| if has_price or has_time: | |
| return None | |
| return ( | |
| "Price per image and min generation time aren't available for " | |
| "this dataset, so these plots can't be drawn." | |
| ) | |
| def _pareto_slot_note(price_fig, price_message, time_fig, time_message, data): | |
| has_price = data is not None and _PARETO_PRICE_COLUMN in data.columns | |
| has_time = data is not None and _PARETO_TIME_COLUMN in data.columns | |
| notes = [] | |
| if has_price and not has_time: | |
| notes.append( | |
| "Min generation time isn't available for this dataset, so only " | |
| "price vs score is shown." | |
| ) | |
| elif has_time and not has_price: | |
| notes.append( | |
| "Price per image isn't available for this dataset, so only min " | |
| "generation time vs score is shown." | |
| ) | |
| if price_fig is None and has_price: | |
| notes.append(price_message) | |
| if time_fig is None and has_time: | |
| notes.append(time_message) | |
| if len(notes) == 2 and notes[0] == notes[1]: | |
| notes = notes[:1] | |
| return " ".join(notes) | |
| def _pareto_slot_updates(data, score_columns): | |
| """Updates for a fixed bank of Gradio Plot slots (visible/hidden).""" | |
| score_columns = [column for column in (score_columns or []) if column] | |
| has_price = data is not None and _PARETO_PRICE_COLUMN in data.columns | |
| has_time = data is not None and _PARETO_TIME_COLUMN in data.columns | |
| dataset_note = _pareto_dataset_message(data) | |
| updates = [_pareto_note_update(dataset_note)] | |
| hide_all_slots = not has_price and not has_time | |
| hidden_slot = ( | |
| gr.update(visible=False), | |
| "", | |
| _pareto_note_update(""), | |
| gr.update(visible=False), | |
| gr.update(visible=False), | |
| _pareto_plot_update(None), | |
| gr.update(visible=False), | |
| _pareto_plot_update(None), | |
| ) | |
| for index in range(MAX_PARETO_METRICS): | |
| if hide_all_slots or index >= len(score_columns): | |
| updates.extend(hidden_slot) | |
| continue | |
| score_column = score_columns[index] | |
| price_fig, price_message, time_fig, time_message = _pareto_pair( | |
| data, score_column | |
| ) | |
| show_price = price_fig is not None | |
| show_time = time_fig is not None | |
| updates.extend( | |
| [ | |
| gr.update(visible=True), | |
| f"#### {_display_label(score_column)}", | |
| _pareto_note_update( | |
| _pareto_slot_note( | |
| price_fig, | |
| price_message, | |
| time_fig, | |
| time_message, | |
| data, | |
| ) | |
| ), | |
| gr.update(visible=show_price or show_time), | |
| gr.update(visible=show_price), | |
| _pareto_plot_update(price_fig), | |
| gr.update(visible=show_time), | |
| _pareto_plot_update(time_fig), | |
| ] | |
| ) | |
| return updates | |
| def _samples_html(samples, selected_models, num_prompts, seed=0): | |
| if not samples: | |
| return _pareto_unavailable_html( | |
| "Samples aren't available for this dataset." | |
| ) | |
| images = samples.get("images", {}) | |
| models = [model for model in (selected_models or []) if model in images] | |
| if not models: | |
| models = (samples.get("models") or [])[:2] | |
| return _build_compare_samples_html(samples, models, num_prompts, seed) | |
| def _build_compare_samples_html(samples, selected_models, num_prompts, seed=0): | |
| selected_models = list(selected_models or [])[:MAX_COMPARE_MODELS] | |
| if not selected_models: | |
| return ( | |
| '<div class="compare-empty">' | |
| "Select at least one model to compare samples." | |
| "</div>" | |
| ) | |
| shared_prompt_ids = None | |
| for model in selected_models: | |
| model_prompt_ids = set(samples["images"][model]) | |
| shared_prompt_ids = ( | |
| model_prompt_ids | |
| if shared_prompt_ids is None | |
| else shared_prompt_ids & model_prompt_ids | |
| ) | |
| shared_prompt_ids = sorted(shared_prompt_ids or []) | |
| if not shared_prompt_ids: | |
| return ( | |
| '<div class="compare-empty">' | |
| "No shared prompts found for the selected models." | |
| "</div>" | |
| ) | |
| rng = random.Random(seed) | |
| prompt_pool = list(shared_prompt_ids) | |
| rng.shuffle(prompt_pool) | |
| chosen = prompt_pool[: max(1, min(int(num_prompts), len(prompt_pool)))] | |
| columns = len(selected_models) | |
| blocks = [] | |
| for index, prompt_id in enumerate(chosen, start=1): | |
| prompt_text = escape(samples["prompts"].get(prompt_id, "")) | |
| cells = [] | |
| for model in selected_models: | |
| image_url = escape(samples["images"][model][prompt_id], quote=True) | |
| cells.append( | |
| f""" | |
| <div class="compare-cell"> | |
| <div class="compare-model-label">{escape(display_model_name(model))}</div> | |
| <a href="{image_url}" target="_blank" rel="noopener noreferrer"> | |
| <img src="{image_url}" alt="{escape(display_model_name(model))} sample" loading="lazy" /> | |
| </a> | |
| </div> | |
| """ | |
| ) | |
| blocks.append( | |
| f""" | |
| <div class="compare-prompt-block"> | |
| <div class="compare-prompt-meta"> | |
| <span>Prompt {index}</span> | |
| <span>{escape(prompt_id)}</span> | |
| </div> | |
| <p class="compare-prompt-text">{prompt_text}</p> | |
| <div class="compare-row" style="grid-template-columns: repeat({columns}, minmax(0, 1fr));"> | |
| {''.join(cells)} | |
| </div> | |
| </div> | |
| """ | |
| ) | |
| return "\n".join(blocks) | |
| def _leaderboard_intro_markdown(note): | |
| parts = [ | |
| "Models are ranked by the selected metric, with price and generation " | |
| "time in the same table." | |
| ] | |
| extra = (note or "").strip() | |
| if extra: | |
| parts.append(extra) | |
| return "<p class='view-help'>" + " ".join(parts) + "</p>" | |
| def _filter_row(datasets, metrics, default_dataset_id, default_metric_id=None): | |
| metric_id = _coerce_metric( | |
| datasets, metrics, default_dataset_id, default_metric_id | |
| ) | |
| with gr.Row(elem_classes="view-filters"): | |
| dataset_dd = gr.Dropdown( | |
| choices=_dataset_choices(datasets), | |
| value=default_dataset_id, | |
| label="Dataset", | |
| type="value", | |
| filterable=False, | |
| scale=2, | |
| min_width=160, | |
| ) | |
| metric_dd = gr.Dropdown( | |
| choices=_metric_dropdown_choices(datasets, metrics, default_dataset_id), | |
| value=_metric_dropdown_value(metric_id), | |
| label="Metric", | |
| type="value", | |
| multiselect=True, | |
| allow_custom_value=False, | |
| filterable=True, | |
| scale=2, | |
| min_width=180, | |
| elem_classes="filter-chips", | |
| ) | |
| models_dd = gr.Dropdown( | |
| choices=_model_choices(datasets, default_dataset_id), | |
| value=[], | |
| multiselect=True, | |
| label="Models", | |
| type="value", | |
| allow_custom_value=False, | |
| filterable=True, | |
| scale=2, | |
| min_width=180, | |
| elem_classes="filter-chips", | |
| ) | |
| return dataset_dd, metric_dd, models_dd | |
| def render_image_workspace(datasets, metrics, default_dataset_id, default_metric_id): | |
| default_metric_id = _coerce_metric( | |
| datasets, metrics, default_dataset_id, default_metric_id | |
| ) | |
| initial = resolve_view(datasets, metrics, default_dataset_id, default_metric_id) | |
| initial_data = initial["data"] | |
| initial_columns = initial["columns"] | |
| initial_samples = initial.get("samples") | |
| with gr.Column(elem_classes="workspace-shell"): | |
| with gr.Column(elem_classes="workspace-filters") as filters_host: | |
| gr.Markdown( | |
| "<p class='filter-help'>" | |
| "These filters apply to Leaderboards, Pareto plots, and Samples. " | |
| "On Samples, only datasets we have generations for are listed. " | |
| "On Pareto plots, only datasets with price or generation time " | |
| "are listed. Search in Models, or leave it empty to include " | |
| "every model." | |
| "</p>", | |
| elem_classes="filter-help-host", | |
| ) | |
| dataset_dd, metric_dd, models_dd = _filter_row( | |
| datasets, metrics, default_dataset_id, None | |
| ) | |
| with gr.Tabs(elem_classes="main-tabs") as main_tabs: | |
| with gr.TabItem("Leaderboards", id=TAB_LEADERBOARDS) as lb_tab: | |
| lb_note = gr.Markdown( | |
| _leaderboard_intro_markdown(initial.get("note")), | |
| elem_classes="view-help-host", | |
| ) | |
| platform_choices = _filter_choices(initial_data, "Platform") | |
| owner_choices = _filter_choices(initial_data, "Endpoint Owner") | |
| optimized_choices = _filter_choices(initial_data, "Optimized") | |
| with gr.Row( | |
| elem_classes="leaderboard-controls", | |
| visible=bool( | |
| platform_choices or owner_choices or optimized_choices | |
| ), | |
| ) as lb_controls: | |
| platform = gr.Dropdown( | |
| choices=platform_choices, | |
| value=[], | |
| label="Providers", | |
| multiselect=True, | |
| allow_custom_value=False, | |
| filterable=False, | |
| scale=1, | |
| visible=bool(platform_choices), | |
| ) | |
| owner = gr.Dropdown( | |
| choices=owner_choices, | |
| value=[], | |
| label="Endpoint owners", | |
| multiselect=True, | |
| allow_custom_value=False, | |
| filterable=False, | |
| scale=1, | |
| visible=bool(owner_choices), | |
| ) | |
| optimized = gr.Dropdown( | |
| choices=optimized_choices, | |
| value=[], | |
| label="Optimized", | |
| multiselect=True, | |
| allow_custom_value=False, | |
| filterable=False, | |
| scale=1, | |
| visible=bool(optimized_choices), | |
| ) | |
| ranking = gr.HTML( | |
| _leaderboard_html( | |
| _assign_leaderboard_ranks( | |
| initial_data, initial.get("score_column") | |
| ), | |
| initial_columns, | |
| ), | |
| padding=False, | |
| elem_classes="ranking-table-host", | |
| ) | |
| with gr.TabItem( | |
| "Pareto Plots", | |
| id=TAB_PARETO, | |
| interactive=_dataset_has_pareto(datasets, default_dataset_id), | |
| ) as pp_tab: | |
| gr.Markdown( | |
| "<p class='view-help'>" | |
| "Score against price and generation time. Green points are on the " | |
| "frontier; lavender points sit below it. Hover a point to see " | |
| "which model it is." | |
| "</p>", | |
| elem_classes="view-help-host", | |
| ) | |
| pareto_dataset_note = gr.HTML( | |
| "", | |
| padding=False, | |
| visible=False, | |
| elem_classes="pareto-note", | |
| ) | |
| pareto_slots = [] | |
| for slot_index in range(MAX_PARETO_METRICS): | |
| with gr.Column( | |
| visible=False, | |
| elem_classes="pareto-metric-block", | |
| ) as slot_group: | |
| slot_title = gr.Markdown( | |
| "", | |
| elem_classes="pareto-metric-title", | |
| ) | |
| with gr.Row( | |
| equal_height=True, | |
| elem_classes="pareto-layout", | |
| ) as slot_layout: | |
| with gr.Column( | |
| scale=1, | |
| min_width=320, | |
| elem_classes="pareto-col", | |
| ) as slot_price_col: | |
| gr.Markdown( | |
| "#### Price vs score", | |
| elem_classes="pareto-subhead", | |
| ) | |
| slot_price = gr.Plot( | |
| value=None, | |
| show_label=False, | |
| elem_classes="pareto-plot", | |
| ) | |
| with gr.Column( | |
| scale=1, | |
| min_width=320, | |
| elem_classes="pareto-col", | |
| ) as slot_time_col: | |
| gr.Markdown( | |
| "#### Min generation time vs score", | |
| elem_classes="pareto-subhead", | |
| ) | |
| slot_time = gr.Plot( | |
| value=None, | |
| show_label=False, | |
| elem_classes="pareto-plot", | |
| ) | |
| slot_note = gr.HTML( | |
| "", | |
| padding=False, | |
| visible=False, | |
| elem_classes="pareto-note", | |
| ) | |
| pareto_slots.append( | |
| ( | |
| slot_group, | |
| slot_title, | |
| slot_note, | |
| slot_layout, | |
| slot_price_col, | |
| slot_price, | |
| slot_time_col, | |
| slot_time, | |
| ) | |
| ) | |
| with gr.TabItem( | |
| "Samples", | |
| id=TAB_SAMPLES, | |
| interactive=_dataset_has_samples(datasets, default_dataset_id), | |
| ) as sm_tab: | |
| with gr.Column(visible=bool(initial_samples)) as samples_panel: | |
| gr.Markdown( | |
| f"<p class='view-help'>" | |
| f"The same prompts, side by side. Select up to " | |
| f"<strong>{MAX_COMPARE_MODELS}</strong> models above, or leave " | |
| f"Models empty for two defaults." | |
| f"</p>", | |
| elem_classes="view-help-host", | |
| ) | |
| with gr.Row(equal_height=False, elem_classes="compare-controls"): | |
| prompt_count = gr.Slider( | |
| minimum=1, | |
| maximum=MAX_COMPARE_PROMPTS, | |
| value=DEFAULT_COMPARE_PROMPTS, | |
| step=1, | |
| label="Prompts to show", | |
| container=False, | |
| show_reset_button=False, | |
| scale=1, | |
| min_width=180, | |
| elem_classes="compare-prompt-count", | |
| ) | |
| shuffle_button = gr.Button( | |
| "Shuffle prompts", | |
| variant="primary", | |
| scale=0, | |
| min_width=140, | |
| elem_classes="compare-shuffle", | |
| ) | |
| gallery = gr.HTML( | |
| value=_samples_html( | |
| initial_samples, [], DEFAULT_COMPARE_PROMPTS, seed=0 | |
| ), | |
| elem_classes="compare-gallery", | |
| ) | |
| seed_state = gr.State(0) | |
| with gr.TabItem("About", id=TAB_ABOUT) as about_tab: | |
| render_about() | |
| def _synced_filters(dataset_id, metric_id, models, *, clear_metric=False): | |
| if clear_metric: | |
| metric_id = [] | |
| else: | |
| metric_id = _coerce_metric(datasets, metrics, dataset_id, metric_id) | |
| model_choices = _model_choices(datasets, dataset_id) | |
| model_values = set(_model_choice_values(model_choices)) | |
| models = [model for model in (models or []) if model in model_values] | |
| metric_choices = _metric_dropdown_choices(datasets, metrics, dataset_id) | |
| return ( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| gr.update( | |
| choices=metric_choices, | |
| value=_metric_dropdown_value(metric_id), | |
| ), | |
| gr.update(choices=model_choices, value=models), | |
| ) | |
| def _leaderboard_extras(data, platform_value, owner_value, optimized_value): | |
| platform_choices = _filter_choices(data, "Platform") | |
| owner_choices = _filter_choices(data, "Endpoint Owner") | |
| optimized_choices = _filter_choices(data, "Optimized") | |
| platform_value = [ | |
| value for value in (platform_value or []) if value in platform_choices | |
| ] | |
| owner_value = [ | |
| value for value in (owner_value or []) if value in owner_choices | |
| ] | |
| optimized_value = [ | |
| value for value in (optimized_value or []) if value in optimized_choices | |
| ] | |
| return ( | |
| gr.update( | |
| choices=platform_choices, | |
| value=platform_value, | |
| visible=bool(platform_choices), | |
| ), | |
| gr.update( | |
| choices=owner_choices, | |
| value=owner_value, | |
| visible=bool(owner_choices), | |
| ), | |
| gr.update( | |
| choices=optimized_choices, | |
| value=optimized_value, | |
| visible=bool(optimized_choices), | |
| ), | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| gr.update( | |
| visible=bool( | |
| platform_choices or owner_choices or optimized_choices | |
| ) | |
| ), | |
| ) | |
| def _content_flags(tab): | |
| return { | |
| "include_leaderboard": tab == TAB_LEADERBOARDS, | |
| "include_pareto": tab == TAB_PARETO, | |
| "include_samples": tab == TAB_SAMPLES, | |
| } | |
| def _commit_state( | |
| view_state, | |
| dataset_id, | |
| metric_id, | |
| models, | |
| tab, | |
| flags, | |
| extras=None, | |
| ): | |
| prev = dict(view_state or {}) | |
| extras = extras or {} | |
| return { | |
| "dataset_id": dataset_id, | |
| "metric_id": metric_id, | |
| "models": list(models or []), | |
| "current_tab": tab, | |
| "platform": list( | |
| extras.get("platform", prev.get("platform") or []) | |
| ), | |
| "owner": list(extras.get("owner", prev.get("owner") or [])), | |
| "optimized": list( | |
| extras.get("optimized", prev.get("optimized") or []) | |
| ), | |
| "stale": { | |
| TAB_LEADERBOARDS: not flags["include_leaderboard"], | |
| TAB_PARETO: not flags["include_pareto"], | |
| TAB_SAMPLES: not flags["include_samples"], | |
| }, | |
| } | |
| def _save_leaderboard_filters( | |
| view_state, platform_value, owner_value, optimized_value | |
| ): | |
| view_state["platform"] = list(platform_value or []) | |
| view_state["owner"] = list(owner_value or []) | |
| view_state["optimized"] = list(optimized_value or []) | |
| return view_state | |
| def _restore_leaderboard_filters( | |
| view_state, platform_value, owner_value, optimized_value | |
| ): | |
| stored_platform = (view_state or {}).get("platform") or [] | |
| stored_owner = (view_state or {}).get("owner") or [] | |
| stored_optimized = (view_state or {}).get("optimized") or [] | |
| view = resolve_view( | |
| datasets, | |
| metrics, | |
| (view_state or {}).get("dataset_id"), | |
| (view_state or {}).get("metric_id"), | |
| ) | |
| extras = _leaderboard_extras( | |
| view["data"] if view else None, | |
| stored_platform, | |
| stored_owner, | |
| stored_optimized, | |
| ) | |
| platform_update = extras[0] if list(platform_value or []) != extras[3] else gr.skip() | |
| owner_update = extras[1] if list(owner_value or []) != extras[4] else gr.skip() | |
| optimized_update = extras[2] if list(optimized_value or []) != extras[5] else gr.skip() | |
| return ( | |
| extras[6], | |
| platform_update, | |
| owner_update, | |
| optimized_update, | |
| extras[3], | |
| extras[4], | |
| extras[5], | |
| ) | |
| def _views( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| *, | |
| include_leaderboard=True, | |
| include_pareto=False, | |
| include_samples=False, | |
| ): | |
| view = resolve_view(datasets, metrics, dataset_id, metric_id) | |
| data = view["data"] | |
| if include_leaderboard: | |
| note = _leaderboard_intro_markdown(view.get("note")) | |
| ranking_html = _leaderboard_html( | |
| _filter_leaderboard( | |
| _assign_leaderboard_ranks(data, view["score_column"]), | |
| platform_value or [], | |
| owner_value or [], | |
| optimized_value or [], | |
| models=models, | |
| ), | |
| view["columns"], | |
| ) | |
| else: | |
| note = gr.skip() | |
| ranking_html = gr.skip() | |
| if include_pareto: | |
| pareto_data = _filter_leaderboard(data, [], [], [], models=models) | |
| pareto_updates = _pareto_slot_updates(pareto_data, view["score_columns"]) | |
| else: | |
| pareto_updates = _pareto_skip_updates() | |
| if include_samples: | |
| samples = view.get("samples") | |
| samples_html = _samples_html( | |
| samples, | |
| models, | |
| int(num_prompts or DEFAULT_COMPARE_PROMPTS), | |
| int(seed or 0), | |
| ) | |
| samples_visible = gr.update(visible=bool(samples)) | |
| else: | |
| samples_html = gr.skip() | |
| samples_visible = gr.skip() | |
| return ( | |
| note, | |
| ranking_html, | |
| *pareto_updates, | |
| samples_html, | |
| samples_visible, | |
| ) | |
| def _apply_filter_change( | |
| source, | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ): | |
| view_state = dict(view_state or {}) | |
| tab = view_state.get("current_tab") or TAB_LEADERBOARDS | |
| selected_raw = _normalize_metric_ids(metric_id) | |
| incoming_models = list(models or []) | |
| dataset_changed = source == "dataset" and dataset_id != view_state.get( | |
| "dataset_id" | |
| ) | |
| if source == "dataset": | |
| synced = _synced_filters( | |
| dataset_id, metric_id, models, clear_metric=dataset_changed | |
| ) | |
| dataset_id, metric_id, models = synced[:3] | |
| metric_update, models_update = synced[3], synced[4] | |
| else: | |
| metric_id = _coerce_metric(datasets, metrics, dataset_id, metric_id) | |
| if source == "models": | |
| allowed = set( | |
| _model_choice_values(_model_choices(datasets, dataset_id)) | |
| ) | |
| models = [model for model in incoming_models if model in allowed] | |
| models_update = ( | |
| gr.update(value=models) | |
| if models != incoming_models | |
| else gr.skip() | |
| ) | |
| else: | |
| models = incoming_models | |
| models_update = gr.skip() | |
| metric_update = ( | |
| gr.update( | |
| choices=_metric_dropdown_choices(datasets, metrics, dataset_id), | |
| value=_metric_dropdown_value(metric_id), | |
| ) | |
| if source == "metric" and ALL_METRICS_ID in selected_raw | |
| else gr.skip() | |
| ) | |
| unchanged = _applied_key(view_state) == _selection_key( | |
| dataset_id, metric_id, models | |
| ) | |
| if ( | |
| not dataset_changed | |
| and not (source == "metric" and ALL_METRICS_ID in selected_raw) | |
| and unchanged | |
| ): | |
| return None | |
| selected_tab = tab | |
| extras = ( | |
| list(platform_value or []), | |
| list(owner_value or []), | |
| list(optimized_value or []), | |
| ) | |
| extra_updates = None | |
| can_pareto = _dataset_has_pareto(datasets, dataset_id) | |
| can_samples = _dataset_has_samples(datasets, dataset_id) | |
| if source == "dataset": | |
| if tab == TAB_SAMPLES and not can_samples: | |
| selected_tab = TAB_LEADERBOARDS | |
| elif tab == TAB_PARETO and not can_pareto: | |
| selected_tab = TAB_LEADERBOARDS | |
| view = resolve_view(datasets, metrics, dataset_id, metric_id) | |
| extra_updates = _leaderboard_extras( | |
| view["data"] if view else None, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| ) | |
| extras = extra_updates[3:6] | |
| flags = _content_flags(selected_tab) | |
| extras_payload = ( | |
| { | |
| "platform": extras[0], | |
| "owner": extras[1], | |
| "optimized": extras[2], | |
| } | |
| if selected_tab == TAB_LEADERBOARDS | |
| else {} | |
| ) | |
| return { | |
| "dataset_id": dataset_id, | |
| "metric_update": metric_update, | |
| "models_update": models_update, | |
| "selected_tab": selected_tab, | |
| "tab": tab, | |
| "extra_updates": extra_updates, | |
| "can_pareto": can_pareto, | |
| "can_samples": can_samples, | |
| "views": _views( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| extras[0], | |
| extras[1], | |
| extras[2], | |
| num_prompts, | |
| seed, | |
| **flags, | |
| ), | |
| "state": _commit_state( | |
| view_state, | |
| dataset_id, | |
| metric_id, | |
| models, | |
| selected_tab, | |
| flags, | |
| extras=extras_payload, | |
| ), | |
| } | |
| def on_dataset( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ): | |
| result = _apply_filter_change( | |
| "dataset", | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ) | |
| if result is None: | |
| return _skip_all(len(dataset_outputs)) | |
| extras = result["extra_updates"] | |
| return ( | |
| _dataset_dropdown_update( | |
| datasets, result["selected_tab"], result["dataset_id"] | |
| ), | |
| result["metric_update"], | |
| result["models_update"], | |
| extras[6], | |
| extras[0], | |
| extras[1], | |
| extras[2], | |
| *result["views"], | |
| gr.update(interactive=result["can_pareto"]), | |
| gr.update(interactive=result["can_samples"]), | |
| gr.update(selected=result["selected_tab"]) | |
| if result["selected_tab"] != result["tab"] | |
| else gr.skip(), | |
| result["state"], | |
| ) | |
| def on_metric( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ): | |
| result = _apply_filter_change( | |
| "metric", | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ) | |
| if result is None: | |
| return _skip_all(len(metric_outputs)) | |
| return (result["metric_update"], *result["views"], result["state"]) | |
| def on_models( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ): | |
| result = _apply_filter_change( | |
| "models", | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ) | |
| if result is None: | |
| return _skip_all(len(models_outputs)) | |
| return (result["models_update"], *result["views"], result["state"]) | |
| def on_tab_select( | |
| tab, | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ): | |
| view_state = dict(view_state or {}) | |
| prev_tab = view_state.get("current_tab") or TAB_LEADERBOARDS | |
| requested_tab = tab | |
| if tab == TAB_SAMPLES and not _dataset_has_samples(datasets, dataset_id): | |
| tab = TAB_LEADERBOARDS | |
| elif tab == TAB_PARETO and not _dataset_has_pareto(datasets, dataset_id): | |
| tab = TAB_LEADERBOARDS | |
| if prev_tab == TAB_LEADERBOARDS: | |
| _save_leaderboard_filters( | |
| view_state, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| ) | |
| dataset_update = _dataset_dropdown_update(datasets, tab, dataset_id) | |
| metric_id = _coerce_metric(datasets, metrics, dataset_id, metric_id) | |
| models = list(models or []) | |
| view_state["current_tab"] = tab | |
| view_state["dataset_id"] = dataset_id | |
| view_state["metric_id"] = metric_id | |
| view_state["models"] = models | |
| show_filters = tab != TAB_ABOUT | |
| show_metric = tab in (TAB_LEADERBOARDS, TAB_PARETO) | |
| was_filters = prev_tab != TAB_ABOUT | |
| was_metric = prev_tab in (TAB_LEADERBOARDS, TAB_PARETO) | |
| filters_vis = ( | |
| gr.update(visible=show_filters) | |
| if show_filters != was_filters | |
| else gr.skip() | |
| ) | |
| metric_vis = ( | |
| gr.update(visible=show_metric) | |
| if show_metric != was_metric | |
| else gr.skip() | |
| ) | |
| if tab == TAB_LEADERBOARDS: | |
| restored = _restore_leaderboard_filters( | |
| view_state, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| ) | |
| platform_value = restored[4] | |
| owner_value = restored[5] | |
| optimized_value = restored[6] | |
| lb_filters = restored[:4] | |
| else: | |
| lb_filters = _skip_all(4) | |
| stale = dict(view_state.get("stale") or {}) | |
| chrome = ( | |
| filters_vis, | |
| dataset_update, | |
| metric_vis, | |
| gr.skip(), | |
| *lb_filters, | |
| ) | |
| tab_select = ( | |
| gr.update(selected=tab) | |
| if tab != requested_tab | |
| else gr.skip() | |
| ) | |
| if tab == TAB_ABOUT or not stale.get(tab, True): | |
| return ( | |
| *chrome, | |
| *_skip_all(len(view_outputs)), | |
| tab_select, | |
| view_state, | |
| ) | |
| flags = _content_flags(tab) | |
| views = _views( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| **flags, | |
| ) | |
| stale[tab] = False | |
| view_state["stale"] = stale | |
| return (*chrome, *views, tab_select, view_state) | |
| def on_leaderboard_filters( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| view_state, | |
| ): | |
| view_state = dict(view_state or {}) | |
| _save_leaderboard_filters( | |
| view_state, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| ) | |
| view = resolve_view(datasets, metrics, dataset_id, metric_id) | |
| return ( | |
| _leaderboard_html( | |
| _filter_leaderboard( | |
| _assign_leaderboard_ranks(view["data"], view["score_column"]), | |
| platform_value or [], | |
| owner_value or [], | |
| optimized_value or [], | |
| models=models, | |
| ), | |
| view["columns"], | |
| ), | |
| view_state, | |
| ) | |
| def on_samples_controls(dataset_id, models, num_prompts, seed): | |
| view = resolve_view(datasets, metrics, dataset_id, None) | |
| return _samples_html( | |
| view.get("samples") if view else None, | |
| models, | |
| int(num_prompts or DEFAULT_COMPARE_PROMPTS), | |
| int(seed or 0), | |
| ) | |
| def on_shuffle(dataset_id, models, num_prompts, seed): | |
| next_seed = int(seed or 0) + 1 | |
| view = resolve_view(datasets, metrics, dataset_id, None) | |
| return next_seed, _samples_html( | |
| view.get("samples") if view else None, | |
| models, | |
| int(num_prompts or DEFAULT_COMPARE_PROMPTS), | |
| next_seed, | |
| ) | |
| def _on_tab(tab): | |
| def handler( | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ): | |
| return on_tab_select( | |
| tab, | |
| dataset_id, | |
| metric_id, | |
| models, | |
| platform_value, | |
| owner_value, | |
| optimized_value, | |
| num_prompts, | |
| seed, | |
| view_state, | |
| ) | |
| handler.__name__ = f"on_tab_{tab}" | |
| return handler | |
| view_state = gr.State( | |
| { | |
| "dataset_id": default_dataset_id, | |
| "metric_id": None, | |
| "models": [], | |
| "current_tab": TAB_LEADERBOARDS, | |
| "platform": [], | |
| "owner": [], | |
| "optimized": [], | |
| "stale": { | |
| TAB_LEADERBOARDS: False, | |
| TAB_PARETO: True, | |
| TAB_SAMPLES: False, | |
| }, | |
| } | |
| ) | |
| pareto_outputs = [ | |
| pareto_dataset_note, | |
| *[ | |
| component | |
| for slot_group, slot_title, slot_note, slot_layout, slot_price_col, slot_price, slot_time_col, slot_time in pareto_slots | |
| for component in ( | |
| slot_group, | |
| slot_title, | |
| slot_note, | |
| slot_layout, | |
| slot_price_col, | |
| slot_price, | |
| slot_time_col, | |
| slot_time, | |
| ) | |
| ], | |
| ] | |
| view_inputs = [ | |
| platform, | |
| owner, | |
| optimized, | |
| prompt_count, | |
| seed_state, | |
| view_state, | |
| ] | |
| view_outputs = [ | |
| lb_note, | |
| ranking, | |
| *pareto_outputs, | |
| gallery, | |
| samples_panel, | |
| ] | |
| filter_inputs = [dataset_dd, metric_dd, models_dd, *view_inputs] | |
| dataset_outputs = [ | |
| dataset_dd, | |
| metric_dd, | |
| models_dd, | |
| lb_controls, | |
| platform, | |
| owner, | |
| optimized, | |
| *view_outputs, | |
| pp_tab, | |
| sm_tab, | |
| main_tabs, | |
| view_state, | |
| ] | |
| dataset_dd.change( | |
| on_dataset, | |
| inputs=filter_inputs, | |
| outputs=dataset_outputs, | |
| **_VIEW_EVENTS, | |
| ) | |
| metric_outputs = [ | |
| metric_dd, | |
| *view_outputs, | |
| view_state, | |
| ] | |
| metric_dd.change( | |
| on_metric, | |
| inputs=filter_inputs, | |
| outputs=metric_outputs, | |
| **_VIEW_EVENTS, | |
| ) | |
| models_outputs = [ | |
| models_dd, | |
| *view_outputs, | |
| view_state, | |
| ] | |
| models_dd.change( | |
| on_models, | |
| inputs=filter_inputs, | |
| outputs=models_outputs, | |
| **_VIEW_EVENTS, | |
| ) | |
| tab_outputs = [ | |
| filters_host, | |
| dataset_dd, | |
| metric_dd, | |
| models_dd, | |
| lb_controls, | |
| platform, | |
| owner, | |
| optimized, | |
| *view_outputs, | |
| main_tabs, | |
| view_state, | |
| ] | |
| for tab, tab_item in ( | |
| (TAB_LEADERBOARDS, lb_tab), | |
| (TAB_PARETO, pp_tab), | |
| (TAB_SAMPLES, sm_tab), | |
| (TAB_ABOUT, about_tab), | |
| ): | |
| tab_item.select( | |
| _on_tab(tab), | |
| inputs=filter_inputs, | |
| outputs=tab_outputs, | |
| show_progress="hidden", | |
| ) | |
| for component in (platform, owner, optimized): | |
| component.change( | |
| on_leaderboard_filters, | |
| inputs=[ | |
| dataset_dd, | |
| metric_dd, | |
| models_dd, | |
| platform, | |
| owner, | |
| optimized, | |
| view_state, | |
| ], | |
| outputs=[ranking, view_state], | |
| show_progress="hidden", | |
| ) | |
| prompt_count.change( | |
| on_samples_controls, | |
| inputs=[dataset_dd, models_dd, prompt_count, seed_state], | |
| outputs=gallery, | |
| show_progress="hidden", | |
| ) | |
| shuffle_button.click( | |
| on_shuffle, | |
| inputs=[dataset_dd, models_dd, prompt_count, seed_state], | |
| outputs=[seed_state, gallery], | |
| show_progress="hidden", | |
| ) | |
| def render_about(): | |
| with gr.Row(elem_classes="about-layout", equal_height=False): | |
| with gr.Column(scale=1, min_width=340, elem_classes="about-col"): | |
| gr.Markdown(ABOUT_OVERVIEW_CONTENT, elem_classes="about-copy") | |
| with gr.Column(scale=1, min_width=340, elem_classes="about-col"): | |
| gr.Markdown(ABOUT_DETAILS_CONTENT, elem_classes="about-copy") | |
| def render_footer(): | |
| gr.HTML(FOOTER_CONTENT, padding=False, elem_classes="page-footer-host") | |
| with gr.Accordion( | |
| "Citation", | |
| open=False, | |
| elem_classes="citation-accordion", | |
| ): | |
| gr.Markdown(CITATION_CONTENT) | |