pefidias commited on
Commit
9f99eeb
·
verified ·
1 Parent(s): 2e5a9eb
Files changed (9) hide show
  1. .gitignore +21 -6
  2. Makefile +2 -10
  3. README.md +26 -19
  4. app.py +40 -14
  5. pyproject.toml +1 -7
  6. requirements.txt +1 -10
  7. src/about.py +30 -47
  8. src/display/css_html_js.py +42 -83
  9. src/display/plotting.py +37 -15
.gitignore CHANGED
@@ -1,16 +1,31 @@
1
- auto_evals/
2
- venv/
3
  __pycache__/
 
 
 
 
 
 
4
  .env
5
- .ipynb_checkpoints
6
- *ipynb
7
  .vscode/
 
 
 
8
 
 
 
 
 
 
9
  eval-queue/
10
  eval-results/
11
  eval-queue-bk/
12
  eval-results-bk/
13
  logs/
 
14
 
15
-
16
- *.pem
 
 
1
+ # Python
 
2
  __pycache__/
3
+ *.pyc
4
+ *.pyo
5
+ *.pyd
6
+ .Python
7
+ venv/
8
+ .venv/
9
  .env
10
+
11
+ # IDEs
12
  .vscode/
13
+ .idea/
14
+ *.swp
15
+ *.swo
16
 
17
+ # Jupyter
18
+ .ipynb_checkpoints/
19
+ *.ipynb
20
+
21
+ # Project specific
22
  eval-queue/
23
  eval-results/
24
  eval-queue-bk/
25
  eval-results-bk/
26
  logs/
27
+ auto_evals/
28
 
29
+ # Other
30
+ *.pem
31
+ .DS_Store
Makefile CHANGED
@@ -1,16 +1,8 @@
1
- .PHONY: style format
2
-
3
 
4
  style:
5
  python -m black --line-length 119 .
6
- python -m isort .
7
  ruff check --fix .
8
 
9
-
10
- quality:
11
- python -m black --check --line-length 119 .
12
- python -m isort --check-only .
13
- ruff check .
14
-
15
- make app:
16
  gradio app.py
 
1
+ .PHONY: style run
 
2
 
3
  style:
4
  python -m black --line-length 119 .
 
5
  ruff check --fix .
6
 
7
+ run:
 
 
 
 
 
 
8
  gradio app.py
README.md CHANGED
@@ -1,5 +1,5 @@
1
  ---
2
- title: Dna Benchmark
3
  emoji: 🥇
4
  colorFrom: green
5
  colorTo: indigo
@@ -7,40 +7,47 @@ sdk: gradio
7
  app_file: app.py
8
  pinned: true
9
  license: apache-2.0
10
- short_description: DNA Foundational Models leaderboard!
11
  sdk_version: 5.19.0
12
  ---
13
 
14
- # Start the configuration
15
 
16
- Most of the variables to change for a default leaderboard are in `src/env.py` (replace the path for your leaderboard) and `src/about.py` (for tasks).
17
 
18
- Results files should have the following format and be stored as json files:
 
 
 
 
 
 
 
 
19
  ```json
20
  {
21
  "config": {
22
- "model_dtype": "torch.float16", # or torch.bfloat16 or 8bit or 4bit
23
- "model_name": "path of the model on the hub: org/model",
24
- "model_sha": "revision on the hub",
25
  },
26
  "results": {
27
  "task_name": {
28
- "metric_name": score,
29
- },
30
- "task_name2": {
31
- "metric_name": score,
32
  }
33
  }
34
  }
35
  ```
36
 
37
- Request files are created automatically by this tool.
38
 
39
- If you encounter problem on the space, don't hesitate to restart it to remove the create eval-queue, eval-queue-bk, eval-results and eval-results-bk created folder.
 
 
40
 
41
- # Code logic for more complex edits
 
42
 
43
- You'll find
44
- - the main table' columns names and properties in `src/display/utils.py`
45
- - the logic to read all results and request files, then convert them in dataframe lines, in `src/leaderboard/read_evals.py`, and `src/populate.py`
46
- - the logic to allow or filter submissions in `src/submission/submit.py` and `src/submission/check_validity.py`
 
1
  ---
2
+ title: DNA Benchmark
3
  emoji: 🥇
4
  colorFrom: green
5
  colorTo: indigo
 
7
  app_file: app.py
8
  pinned: true
9
  license: apache-2.0
10
+ short_description: DNA Foundational Models leaderboard
11
  sdk_version: 5.19.0
12
  ---
13
 
14
+ # DNA Benchmark Leaderboard
15
 
16
+ A simple leaderboard for evaluating DNA foundational models on genomics tasks.
17
 
18
+ ## Configuration
19
+
20
+ - `src/envs.py` - Configure repository paths and environment variables
21
+ - `src/about.py` - Define tasks and leaderboard text
22
+ - `src/display/utils.py` - Configure table columns and metrics
23
+
24
+ ## Results Format
25
+
26
+ Results files should be stored as JSON with the following structure:
27
  ```json
28
  {
29
  "config": {
30
+ "model_dtype": "torch.float16",
31
+ "model_name": "org/model",
32
+ "model_sha": "revision"
33
  },
34
  "results": {
35
  "task_name": {
36
+ "metric_name": score
 
 
 
37
  }
38
  }
39
  }
40
  ```
41
 
42
+ ## Development
43
 
44
+ ```bash
45
+ # Install dependencies
46
+ pip install -r requirements.txt
47
 
48
+ # Run locally
49
+ gradio app.py
50
 
51
+ # Format code
52
+ make style
53
+ ```
 
app.py CHANGED
@@ -3,7 +3,7 @@ from apscheduler.schedulers.background import BackgroundScheduler
3
  from gradio_leaderboard import ColumnFilter, Leaderboard, SelectColumns
4
  from huggingface_hub import snapshot_download
5
 
6
- from src.about import CITATION_BUTTON_LABEL, CITATION_BUTTON_TEXT, INTRODUCTION_TEXT, LLM_BENCHMARKS_TEXT, TITLE
7
  from src.display.css_html_js import custom_css
8
  from src.display.utils import BenchRawColumn, fields
9
  from src.envs import API, EVAL_RESULTS_PATH, REPO_ID, RESULTS_REPO, TOKEN
@@ -47,7 +47,6 @@ def init_leaderboard(dataframe):
47
  cant_deselect=[c.name for c in fields(BenchRawColumn) if c.never_hidden],
48
  label="Select Columns to Display:",
49
  ),
50
- search_columns=[BenchRawColumn.model.name],
51
  filter_columns=[
52
  ColumnFilter(BenchRawColumn.task.name, type="checkboxgroup", label=BenchRawColumn.task.name),
53
  # ColumnFilter(BenchRawColumn.precision.name, type="checkboxgroup", label="Precision"),
@@ -74,9 +73,46 @@ def init_leaderboard(dataframe):
74
  label="Select the Max Context Length (bp)",
75
  ),
76
  ],
 
77
  )
78
 
79
- demo = gr.Blocks(css=custom_css)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
80
  with demo:
81
  gr.HTML(TITLE)
82
  gr.Markdown(INTRODUCTION_TEXT, elem_classes="markdown-text")
@@ -120,18 +156,8 @@ with demo:
120
  with gr.TabItem("📝 About", elem_id="llm-benchmark-about", id=2):
121
  gr.Markdown(LLM_BENCHMARKS_TEXT, elem_classes="markdown-text")
122
 
123
- with gr.Row():
124
- with gr.Accordion("📙 Citation", open=False):
125
- citation_button = gr.Textbox(
126
- value=CITATION_BUTTON_TEXT,
127
- label=CITATION_BUTTON_LABEL,
128
- lines=20,
129
- elem_id="citation-button",
130
- show_copy_button=True,
131
- )
132
-
133
  scheduler = BackgroundScheduler()
134
  scheduler.add_job(restart_space, "interval", seconds=1800)
135
  scheduler.start()
136
 
137
- demo.queue(default_concurrency_limit=40).launch(share=True)
 
3
  from gradio_leaderboard import ColumnFilter, Leaderboard, SelectColumns
4
  from huggingface_hub import snapshot_download
5
 
6
+ from src.about import INTRODUCTION_TEXT, LLM_BENCHMARKS_TEXT, TITLE
7
  from src.display.css_html_js import custom_css
8
  from src.display.utils import BenchRawColumn, fields
9
  from src.envs import API, EVAL_RESULTS_PATH, REPO_ID, RESULTS_REPO, TOKEN
 
47
  cant_deselect=[c.name for c in fields(BenchRawColumn) if c.never_hidden],
48
  label="Select Columns to Display:",
49
  ),
 
50
  filter_columns=[
51
  ColumnFilter(BenchRawColumn.task.name, type="checkboxgroup", label=BenchRawColumn.task.name),
52
  # ColumnFilter(BenchRawColumn.precision.name, type="checkboxgroup", label="Precision"),
 
73
  label="Select the Max Context Length (bp)",
74
  ),
75
  ],
76
+ search_columns=[BenchRawColumn.model.name],
77
  )
78
 
79
+ # Create an elegant dark theme with custom colors for checkboxes and sliders
80
+ theme = gr.themes.Soft(
81
+ primary_hue="violet",
82
+ secondary_hue="purple",
83
+ neutral_hue="slate",
84
+ spacing_size="md",
85
+ radius_size="lg",
86
+ text_size="md",
87
+ font=[gr.themes.GoogleFont("Inter"), "ui-sans-serif", "system-ui", "sans-serif"],
88
+ ).set(
89
+ checkbox_background_color_selected="*primary_600",
90
+ checkbox_border_color_selected="*primary_600",
91
+ checkbox_background_color_selected_dark="*primary_500",
92
+ checkbox_border_color_selected_dark="*primary_500",
93
+ slider_color="*primary_600",
94
+ slider_color_dark="*primary_500",
95
+ button_primary_background_fill="*primary_600",
96
+ button_primary_background_fill_hover="*primary_700",
97
+ button_primary_background_fill_dark="*primary_500",
98
+ button_primary_background_fill_hover_dark="*primary_600",
99
+ )
100
+
101
+ demo = gr.Blocks(
102
+ css=custom_css,
103
+ theme=theme,
104
+ js="""
105
+ () => {
106
+ const theme = localStorage.getItem('theme');
107
+ if (theme === null) {
108
+ localStorage.setItem('theme', 'dark');
109
+ document.body.classList.add('dark');
110
+ }
111
+ }
112
+ """,
113
+ title="DNA Benchmark",
114
+ head="<link rel='icon' href='data:image/svg+xml,<svg xmlns=\"http://www.w3.org/2000/svg\" viewBox=\"0 0 100 100\"><text y=\".9em\" font-size=\"90\">🧬</text></svg>'>"
115
+ )
116
  with demo:
117
  gr.HTML(TITLE)
118
  gr.Markdown(INTRODUCTION_TEXT, elem_classes="markdown-text")
 
156
  with gr.TabItem("📝 About", elem_id="llm-benchmark-about", id=2):
157
  gr.Markdown(LLM_BENCHMARKS_TEXT, elem_classes="markdown-text")
158
 
 
 
 
 
 
 
 
 
 
 
159
  scheduler = BackgroundScheduler()
160
  scheduler.add_job(restart_space, "interval", seconds=1800)
161
  scheduler.start()
162
 
163
+ demo.queue(default_concurrency_limit=40).launch(share=True)
pyproject.toml CHANGED
@@ -1,13 +1,7 @@
1
  [tool.ruff]
2
- # Enable pycodestyle (`E`) and Pyflakes (`F`) codes by default.
3
  select = ["E", "F"]
4
- ignore = ["E501"] # line too long (black is taking care of this)
5
  line-length = 119
6
- fixable = ["A", "B", "C", "D", "E", "F", "G", "I", "N", "Q", "S", "T", "W", "ANN", "ARG", "BLE", "COM", "DJ", "DTZ", "EM", "ERA", "EXE", "FBT", "ICN", "INP", "ISC", "NPY", "PD", "PGH", "PIE", "PL", "PT", "PTH", "PYI", "RET", "RSE", "RUF", "SIM", "SLF", "TCH", "TID", "TRY", "UP", "YTT"]
7
-
8
- [tool.isort]
9
- profile = "black"
10
- line_length = 119
11
 
12
  [tool.black]
13
  line-length = 119
 
1
  [tool.ruff]
 
2
  select = ["E", "F"]
3
+ ignore = ["E501"]
4
  line-length = 119
 
 
 
 
 
5
 
6
  [tool.black]
7
  line-length = 119
requirements.txt CHANGED
@@ -1,18 +1,9 @@
1
  APScheduler
2
- black
3
  datasets
4
  gradio==5.41.0
5
  gradio[oauth]
6
  gradio_leaderboard==0.0.13
7
- gradio_client
8
  huggingface-hub>=0.18.0
9
  plotly>=5.0.0
10
- matplotlib
11
- numpy
12
  pandas
13
- python-dateutil
14
- tqdm
15
- transformers
16
- tokenizers>=0.15.0
17
- sentencepiece
18
- python_dotenv==1.1.0
 
1
  APScheduler
 
2
  datasets
3
  gradio==5.41.0
4
  gradio[oauth]
5
  gradio_leaderboard==0.0.13
 
6
  huggingface-hub>=0.18.0
7
  plotly>=5.0.0
 
 
8
  pandas
9
+ python-dotenv==1.1.0
 
 
 
 
 
src/about.py CHANGED
@@ -9,66 +9,49 @@ class Task:
9
  col_name: str
10
 
11
 
12
- # Select your tasks here
13
- # ---------------------------------------------------
14
  class Tasks(Enum):
15
- # task_key in the json file, metric_key in the json file, name to display in the leaderboard
16
  task0 = Task("anli_r1", "acc", "ANLI")
17
  task1 = Task("logiqa", "acc_norm", "LogiQA")
18
 
19
 
20
- NUM_FEWSHOT = 0 # Change with your few shot
21
- # ---------------------------------------------------
22
 
23
-
24
- # Your leaderboard name
25
- TITLE = """<h1 align="center" id="space-title">DNA Benchmark</h1>"""
26
-
27
- # What does your leaderboard evaluate?
28
- INTRODUCTION_TEXT = """
29
- Intro text
 
 
30
  """
31
 
32
- # Which evaluations are you running? how can people reproduce what you have?
33
- LLM_BENCHMARKS_TEXT = f"""
34
- ##
35
- This leaderboard evaluates the performance of LLMs on various tasks.
36
-
37
- ## Reproducibility
38
- To reproduce our results, here is the commands you can run:
39
-
40
  """
41
 
42
- EVALUATION_QUEUE_TEXT = """
43
- ## Some good practices before submitting a model
44
-
45
- ### 1) Make sure you can load your model and tokenizer using AutoClasses:
46
- ```python
47
- from transformers import AutoConfig, AutoModel, AutoTokenizer
48
- config = AutoConfig.from_pretrained("your model name", revision=revision)
49
- model = AutoModel.from_pretrained("your model name", revision=revision)
50
- tokenizer = AutoTokenizer.from_pretrained("your model name", revision=revision)
51
- ```
52
- If this step fails, follow the error messages to debug your model before submitting it. It's likely your model has been improperly uploaded.
53
 
54
- Note: make sure your model is public!
55
- Note: if your model needs `use_remote_code=True`, we do not support this option yet but we are working on adding it, stay posted!
56
 
57
- ### 2) Convert your model weights to [safetensors](https://huggingface.co/docs/safetensors/index)
58
- It's a new format for storing weights which is safer and faster to load and use. It will also allow us to add the number of parameters of your model to the `Extended Viewer`!
59
 
60
- ### 3) Make sure your model has an open license!
61
- This is a leaderboard for Open LLMs, and we'd love for as many people as possible to know they can use your model 🤗
 
 
 
62
 
63
- ### 4) Fill up your model card
64
- When we add extra information about models to the leaderboard, it will be automatically taken from the model card
65
-
66
- ## In case of model failure
67
- If your model is displayed in the `FAILED` category, its execution stopped.
68
- Make sure you have followed the above steps first.
69
- If everything is done, check you can launch the EleutherAIHarness on your model locally, using the above command without modifications (you can add `--limit` to limit the number of examples per task).
70
- """
71
 
72
- CITATION_BUTTON_LABEL = "Copy the following snippet to cite these results"
73
- CITATION_BUTTON_TEXT = r"""
74
  """
 
9
  col_name: str
10
 
11
 
 
 
12
  class Tasks(Enum):
 
13
  task0 = Task("anli_r1", "acc", "ANLI")
14
  task1 = Task("logiqa", "acc_norm", "LogiQA")
15
 
16
 
17
+ NUM_FEWSHOT = 0
 
18
 
19
+ TITLE = """
20
+ <div style="text-align: center; background: linear-gradient(135deg, #667eea 0%, #764ba2 100%); padding: 3rem 2rem; border-radius: 15px; margin-bottom: 2rem; box-shadow: 0 10px 30px rgba(0,0,0,0.2);">
21
+ <h1 style="color: white; font-size: 3.5rem; margin: 0; font-weight: 800; text-shadow: 2px 2px 4px rgba(0,0,0,0.3);">
22
+ 🧬 DNA Benchmark
23
+ </h1>
24
+ <p style="color: rgba(255,255,255,0.95); font-size: 1.3rem; margin-top: 1rem; font-weight: 300;">
25
+ Evaluating DNA Foundational Models Performance
26
+ </p>
27
+ </div>
28
  """
29
 
30
+ INTRODUCTION_TEXT = """
31
+ <div style="background: linear-gradient(135deg, #f5f7fa 0%, #c3cfe2 100%); padding: 1.5rem; border-radius: 10px; border-left: 5px solid #667eea; margin-bottom: 1.5rem;">
32
+ <p style="font-size: 1.1rem; color: #2d3748; margin: 0; line-height: 1.6;">
33
+ Compare and analyze the performance of state-of-the-art DNA foundational models across various genomics tasks including histone modification, splicing, promoter identification, enhancer detection, and SNP classification.
34
+ </p>
35
+ </div>
 
 
36
  """
37
 
38
+ LLM_BENCHMARKS_TEXT = """
39
+ <div style="background: var(--background-fill-primary, #fff);">
 
 
 
 
 
 
 
 
 
40
 
41
+ ## 📊 About This Leaderboard
 
42
 
43
+ This leaderboard provides comprehensive benchmarking of DNA foundational models across multiple genomics tasks:
 
44
 
45
+ - **🧬 Histone Modifications**: H2AFZ, H3K27ac, H3K27me3, H3K36me3, H3K4me1/2/3, H3K9ac, H3K9me3, H4K20me1
46
+ - **✂️ Splicing**: Donor sites, Acceptor sites, All splice sites
47
+ - **📍 Promoter Identification**: TATA-containing, Non-TATA, All promoters
48
+ - **🎯 Enhancer Detection**: Enhancer identification and classification
49
+ - **🔬 SNP Classification**: eQTL causality, ClinVar pathogenic variants, OMIM pathogenic variants
50
 
51
+ ### 📈 Key Metrics
52
+ - **Accuracy**: Overall prediction accuracy
53
+ - **MCC**: Matthews Correlation Coefficient
54
+ - **Weighted F1**: F1-score weighted by class support
 
 
 
 
55
 
56
+ </div>
 
57
  """
src/display/css_html_js.py CHANGED
@@ -1,105 +1,64 @@
1
  custom_css = """
2
-
3
- .markdown-text {
4
- font-size: 16px !important;
5
- }
6
-
7
- #models-to-add-text {
8
- font-size: 18px !important;
9
- }
10
-
11
- #citation-button span {
12
- font-size: 16px !important;
13
  }
14
 
15
- #citation-button textarea {
16
- font-size: 16px !important;
 
17
  }
18
 
19
- #citation-button > label > button {
20
- margin: 6px;
21
- transform: scale(1.3);
22
  }
23
 
24
- #leaderboard-table {
25
- margin-top: 15px
 
 
 
26
  }
27
 
28
- #leaderboard-table-lite {
29
- margin-top: 15px
 
 
30
  }
31
-
32
- #search-bar-table-box > div:first-child {
33
- background: none;
34
- border: none;
35
  }
36
-
37
- #search-bar {
38
- padding: 0px;
 
 
 
39
  }
40
 
41
- /* Limit the width of the first AutoEvalColumn so that names don't expand too much */
42
- #leaderboard-table td:nth-child(2),
43
- #leaderboard-table th:nth-child(2) {
44
- max-width: 400px;
45
- overflow: auto;
46
- white-space: nowrap;
47
  }
48
 
49
- .tab-buttons button {
50
- font-size: 20px;
51
  }
52
 
53
- #scale-logo {
54
- border-style: none !important;
55
- box-shadow: none;
56
- display: block;
57
- margin-left: auto;
58
- margin-right: auto;
59
- max-width: 600px;
60
  }
61
 
62
- #scale-logo .download {
63
- display: none;
64
- }
65
- #filter_type{
66
- border: 0;
67
- padding-left: 0;
68
- padding-top: 0;
69
- }
70
- #filter_type label {
71
- display: flex;
72
- }
73
- #filter_type label > span{
74
- margin-top: var(--spacing-lg);
75
- margin-right: 0.5em;
76
- }
77
- #filter_type label > .wrap{
78
- width: 103px;
79
  }
80
- #filter_type label > .wrap .wrap-inner{
81
- padding: 2px;
82
- }
83
- #filter_type label > .wrap .wrap-inner input{
84
- width: 1px
85
- }
86
- #filter-columns-type{
87
- border:0;
88
- padding:0.5;
89
- }
90
- #filter-columns-size{
91
- border:0;
92
- padding:0.5;
93
- }
94
- #box-filter > .form{
95
- border: 0
96
  }
97
  """
98
-
99
- get_window_url_params = """
100
- function(url_params) {
101
- const params = new URLSearchParams(window.location.search);
102
- url_params = Object.fromEntries(params);
103
- return url_params;
104
- }
105
- """
 
1
  custom_css = """
2
+ .gradio-container {
3
+ max-width: 1400px !important;
4
+ margin: 0 auto !important;
5
+ font-family: 'Inter', -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif !important;
 
 
 
 
 
 
 
6
  }
7
 
8
+ .markdown-text {
9
+ font-size: 16px !important;
10
+ line-height: 1.7 !important;
11
  }
12
 
13
+ body:not(.dark) .markdown-text {
14
+ color: #2d3748 !important;
 
15
  }
16
 
17
+ body:not(.dark) .markdown-text h2 {
18
+ color: #667eea !important;
19
+ font-weight: 700 !important;
20
+ margin-top: 1.5rem !important;
21
+ margin-bottom: 1rem !important;
22
  }
23
 
24
+ /* Table: only font weights */
25
+ table tbody td,
26
+ table tbody th {
27
+ font-weight: 400 !important;
28
  }
29
+ table thead th {
30
+ font-weight: 600 !important;
 
 
31
  }
32
+ """
33
+ custom_css = """
34
+ .gradio-container {
35
+ max-width: 1400px !important;
36
+ margin: 0 auto !important;
37
+ font-family: 'Inter', -apple-system, BlinkMacSystemFont, 'Segoe UI', sans-serif !important;
38
  }
39
 
40
+ .markdown-text {
41
+ font-size: 16px !important;
42
+ line-height: 1.7 !important;
 
 
 
43
  }
44
 
45
+ body:not(.dark) .markdown-text {
46
+ color: #2d3748 !important;
47
  }
48
 
49
+ body:not(.dark) .markdown-text h2 {
50
+ color: #667eea !important;
51
+ font-weight: 700 !important;
52
+ margin-top: 1.5rem !important;
53
+ margin-bottom: 1rem !important;
 
 
54
  }
55
 
56
+ /* Table: only font weights */
57
+ table tbody td,
58
+ table tbody th {
59
+ font-weight: 400 !important;
 
 
 
 
 
 
 
 
 
 
 
 
 
60
  }
61
+ table thead th {
62
+ font-weight: 600 !important;
 
 
 
 
 
 
 
 
 
 
 
 
 
 
63
  }
64
  """
 
 
 
 
 
 
 
 
src/display/plotting.py CHANGED
@@ -1,8 +1,17 @@
1
  import plotly.express as px
 
2
  import pandas as pd
3
 
4
  METRICS_FOR_PLOTS = ["Accuracy", "MCC", "Weighted F1"]
5
 
 
 
 
 
 
 
 
 
6
  def extract_mean_std(col):
7
  if isinstance(col, str) and '±' in col:
8
  try:
@@ -45,7 +54,7 @@ def make_model_all_datasets_wrapper(leaderboard_df):
45
 
46
  def plot_metric_bar_all_datasets(df, model, metric, orientation="v"):
47
  df = df.copy()
48
- df = df[df['Model'] == model]
49
 
50
  y_col = f"{metric}_mean"
51
  std_col = f"{metric}_std"
@@ -62,12 +71,13 @@ def plot_metric_bar_all_datasets(df, model, metric, orientation="v"):
62
  df,
63
  x=x_axis,
64
  y=y_axis,
65
- color='Dataset Name',
66
  orientation=orientation,
67
  text_auto=".2f",
68
  labels={y_col: metric, "Task": "Task"},
69
  title=f"{metric} for {model} across all datasets",
70
  height=height,
 
71
  )
72
 
73
  if orientation == "h":
@@ -96,20 +106,25 @@ def plot_metric_bar_all_datasets(df, model, metric, orientation="v"):
96
  'x': 0.5,
97
  'xanchor': 'center',
98
  'yanchor': 'top',
99
- 'font': dict(size=20),
100
- 'pad': dict(t=0, b=0),
101
  },
102
- margin=dict(t=10, b=100, l=150, r=10),
 
 
 
103
  )
104
 
105
  if orientation == "h":
106
- layout_args["xaxis"] = dict(range=[0, 1])
107
  layout_args["yaxis_tickfont_size"] = 10
 
108
  else:
109
- layout_args["yaxis"] = dict(range=[0, 1])
 
110
 
111
  fig.update_layout(**layout_args)
112
-
113
  return fig
114
 
115
  def plot_metric_bar(df, group_by, filter_col, filter_val, dataset, metric, orientation="v"):
@@ -155,7 +170,9 @@ def plot_metric_bar(df, group_by, filter_col, filter_val, dataset, metric, orien
155
  )
156
 
157
  fig.update_traces(
158
- marker_color="#F97316",
 
 
159
  hovertemplate=hovertemplate,
160
  customdata=df[[std_col]].values
161
  )
@@ -166,18 +183,23 @@ def plot_metric_bar(df, group_by, filter_col, filter_val, dataset, metric, orien
166
  'x': 0.5,
167
  'xanchor': 'center',
168
  'yanchor': 'top',
169
- 'font': dict(size=20),
170
- 'pad': dict(t=0, b=0),
171
  },
172
- margin=dict(t=10, b=100, l=150, r=10),
 
 
 
173
  )
174
 
175
  if orientation == "h":
176
- layout_args["xaxis"] = dict(range=[0, 1])
177
  layout_args["yaxis_tickfont_size"] = 10
 
178
  else:
179
- layout_args["yaxis"] = dict(range=[0, 1])
 
180
 
181
  fig.update_layout(**layout_args)
182
-
183
  return fig
 
1
  import plotly.express as px
2
+ import plotly.graph_objects as go
3
  import pandas as pd
4
 
5
  METRICS_FOR_PLOTS = ["Accuracy", "MCC", "Weighted F1"]
6
 
7
+ # Color scheme matching the UI theme
8
+ COLORS = {
9
+ 'primary': '#667eea',
10
+ 'secondary': '#764ba2',
11
+ 'accent': '#F97316',
12
+ 'gradient': ['#667eea', '#764ba2', '#F97316', '#06b6d4', '#10b981']
13
+ }
14
+
15
  def extract_mean_std(col):
16
  if isinstance(col, str) and '±' in col:
17
  try:
 
54
 
55
  def plot_metric_bar_all_datasets(df, model, metric, orientation="v"):
56
  df = df.copy()
57
+ df = df[df['Model'] == model]
58
 
59
  y_col = f"{metric}_mean"
60
  std_col = f"{metric}_std"
 
71
  df,
72
  x=x_axis,
73
  y=y_axis,
74
+ color='Dataset Name',
75
  orientation=orientation,
76
  text_auto=".2f",
77
  labels={y_col: metric, "Task": "Task"},
78
  title=f"{metric} for {model} across all datasets",
79
  height=height,
80
+ color_discrete_sequence=COLORS['gradient']
81
  )
82
 
83
  if orientation == "h":
 
106
  'x': 0.5,
107
  'xanchor': 'center',
108
  'yanchor': 'top',
109
+ 'font': dict(size=22, color='#2d3748', family='Inter, sans-serif'),
110
+ 'pad': dict(t=20, b=10),
111
  },
112
+ margin=dict(t=80, b=100, l=150, r=10),
113
+ plot_bgcolor='rgba(0,0,0,0)',
114
+ paper_bgcolor='white',
115
+ font=dict(family='Inter, sans-serif', color='#2d3748'),
116
  )
117
 
118
  if orientation == "h":
119
+ layout_args["xaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
120
  layout_args["yaxis_tickfont_size"] = 10
121
+ layout_args["yaxis"] = dict(gridcolor='#e2e8f0')
122
  else:
123
+ layout_args["yaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
124
+ layout_args["xaxis"] = dict(gridcolor='#e2e8f0')
125
 
126
  fig.update_layout(**layout_args)
127
+
128
  return fig
129
 
130
  def plot_metric_bar(df, group_by, filter_col, filter_val, dataset, metric, orientation="v"):
 
170
  )
171
 
172
  fig.update_traces(
173
+ marker_color=COLORS['primary'],
174
+ marker_line_color=COLORS['secondary'],
175
+ marker_line_width=1.5,
176
  hovertemplate=hovertemplate,
177
  customdata=df[[std_col]].values
178
  )
 
183
  'x': 0.5,
184
  'xanchor': 'center',
185
  'yanchor': 'top',
186
+ 'font': dict(size=22, color='#2d3748', family='Inter, sans-serif'),
187
+ 'pad': dict(t=20, b=10),
188
  },
189
+ margin=dict(t=80, b=100, l=150, r=10),
190
+ plot_bgcolor='rgba(0,0,0,0)',
191
+ paper_bgcolor='white',
192
+ font=dict(family='Inter, sans-serif', color='#2d3748'),
193
  )
194
 
195
  if orientation == "h":
196
+ layout_args["xaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
197
  layout_args["yaxis_tickfont_size"] = 10
198
+ layout_args["yaxis"] = dict(gridcolor='#e2e8f0')
199
  else:
200
+ layout_args["yaxis"] = dict(range=[0, 1], gridcolor='#e2e8f0')
201
+ layout_args["xaxis"] = dict(gridcolor='#e2e8f0')
202
 
203
  fig.update_layout(**layout_args)
204
+
205
  return fig