imwemans commited on
Commit
fe41b50
·
1 Parent(s): f4a3d41

made code simpler and less repetitive for the plots

Browse files
Files changed (2) hide show
  1. app.py +6 -7
  2. src/display/plotting.py +64 -78
app.py CHANGED
@@ -8,7 +8,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
10
  from src.populate import get_leaderboard_df_from_hf_dataset, summarize_model_task_type_performance
11
- from src.display.plotting import extract_mean_std, prepare_leaderboard_df, make_wrapped_task_plot, make_wrapped_model_plot, plot_task_metric, plot_model_across_tasks
12
 
13
  def restart_space():
14
  API.restart_space(repo_id=REPO_ID)
@@ -27,8 +27,9 @@ try:
27
  LEADERBOARD_DF = get_leaderboard_df_from_hf_dataset(EVAL_RESULTS_PATH)
28
  LEADERBOARD_DF = summarize_model_task_type_performance(LEADERBOARD_DF)
29
  LEADERBOARD_DF = prepare_leaderboard_df(LEADERBOARD_DF)
30
- wrapped_task_plot = make_wrapped_task_plot(LEADERBOARD_DF)
31
- wrapped_model_plot = make_wrapped_model_plot(LEADERBOARD_DF)
 
32
 
33
  except Exception:
34
  restart_space()
@@ -80,11 +81,10 @@ with demo:
80
  gr.Markdown("Visualize model performance for a specific task and metric")
81
 
82
  task_choices = sorted(LEADERBOARD_DF["Task"].dropna().unique())
83
- metric_choices = ["Accuracy", "MCC"]
84
 
85
  with gr.Row():
86
  task_dropdown = gr.Dropdown(choices=task_choices, label="Select Task")
87
- metric_dropdown = gr.Dropdown(choices=metric_choices, value="Accuracy", label="Select Metric")
88
 
89
  performance_plot = gr.Plot()
90
 
@@ -96,11 +96,10 @@ with demo:
96
  gr.Markdown("Visualize model performance across tasks according to a specific metric")
97
 
98
  model_choices = sorted(LEADERBOARD_DF["Model"].dropna().unique())
99
- metric_choices = ["Accuracy", "MCC"]
100
 
101
  with gr.Row():
102
  model_dropdown = gr.Dropdown(choices=model_choices, label="Select Model")
103
- metric2_dropdown = gr.Dropdown(choices=metric_choices, value="Accuracy", label="Select Metric")
104
 
105
  model_plot = gr.Plot()
106
 
 
8
  from src.display.utils import BenchRawColumn, fields
9
  from src.envs import API, EVAL_RESULTS_PATH, REPO_ID, RESULTS_REPO, TOKEN
10
  from src.populate import get_leaderboard_df_from_hf_dataset, summarize_model_task_type_performance
11
+ from src.display.plotting import METRICS_FOR_PLOTS, extract_mean_std, prepare_leaderboard_df, make_plot_wrapper, plot_metric_bar
12
 
13
  def restart_space():
14
  API.restart_space(repo_id=REPO_ID)
 
27
  LEADERBOARD_DF = get_leaderboard_df_from_hf_dataset(EVAL_RESULTS_PATH)
28
  LEADERBOARD_DF = summarize_model_task_type_performance(LEADERBOARD_DF)
29
  LEADERBOARD_DF = prepare_leaderboard_df(LEADERBOARD_DF)
30
+
31
+ wrapped_task_plot = make_plot_wrapper(leaderboard_df=LEADERBOARD_DF, group_by="Model", filter_col="Task", orientation="h")
32
+ wrapped_model_plot = make_plot_wrapper(leaderboard_df=LEADERBOARD_DF, group_by="Task", filter_col="Model", orientation="v")
33
 
34
  except Exception:
35
  restart_space()
 
81
  gr.Markdown("Visualize model performance for a specific task and metric")
82
 
83
  task_choices = sorted(LEADERBOARD_DF["Task"].dropna().unique())
 
84
 
85
  with gr.Row():
86
  task_dropdown = gr.Dropdown(choices=task_choices, label="Select Task")
87
+ metric_dropdown = gr.Dropdown(choices=METRICS_FOR_PLOTS, value="Accuracy", label="Select Metric")
88
 
89
  performance_plot = gr.Plot()
90
 
 
96
  gr.Markdown("Visualize model performance across tasks according to a specific metric")
97
 
98
  model_choices = sorted(LEADERBOARD_DF["Model"].dropna().unique())
 
99
 
100
  with gr.Row():
101
  model_dropdown = gr.Dropdown(choices=model_choices, label="Select Model")
102
+ metric2_dropdown = gr.Dropdown(choices=METRICS_FOR_PLOTS, value="Accuracy", label="Select Metric")
103
 
104
  model_plot = gr.Plot()
105
 
src/display/plotting.py CHANGED
@@ -1,6 +1,8 @@
1
  import plotly.express as px
2
  import pandas as pd
3
 
 
 
4
  def extract_mean_std(col):
5
  if isinstance(col, str) and '±' in col:
6
  try:
@@ -13,108 +15,92 @@ def extract_mean_std(col):
13
  return None, None
14
 
15
  def prepare_leaderboard_df(df):
16
- for metric in ["Accuracy", "MCC"]:
17
  means, stds = zip(*df[metric].apply(extract_mean_std))
18
  df[f"{metric}_mean"] = means
19
  df[f"{metric}_std"] = stds
20
  return df
21
 
22
- def make_wrapped_task_plot(leaderboard_df):
23
- def wrapped_task_plot_fn(task, metric):
24
- return plot_task_metric(leaderboard_df, task, metric)
25
- return wrapped_task_plot_fn
26
-
27
- def make_wrapped_model_plot(leaderboard_df):
28
- def wrapped_model_plot_fn(model, metric):
29
- return plot_model_across_tasks(leaderboard_df, model, metric)
30
- return wrapped_model_plot_fn
31
-
32
- # Plot 1: Model Performance per Task
33
- def plot_task_metric(df, task, metric):
 
34
  df = df.copy()
35
- df = df[df["Task"] == task]
36
 
37
  y_col = f"{metric}_mean"
38
  std_col = f"{metric}_std"
39
 
40
  if df.empty or y_col not in df.columns:
41
- return px.bar(title="No models found for selected task.")
42
 
43
  df = df.sort_values(by=y_col, ascending=False)
44
- df["Model_wrapped"] = df["Model"]
45
-
46
- fig = px.bar(
47
- df,
48
- y="Model_wrapped",
49
- x=y_col,
50
- orientation='h',
51
- title=f"{metric} per Model on {task}",
52
- labels={y_col: metric, "Model_wrapped": "Model"},
53
- text_auto=".2f",
54
- height=20 * len(df) + 200
55
- )
56
- fig.update_traces(
57
- marker_color="#F97316",
58
- hovertemplate=f"<b>%{{y}}</b><br>{metric}: %{{x:.2f}}<br>Std Dev: %{{customdata[0]:.3f}}",
59
- customdata=df[[std_col]].values
60
- )
61
-
62
- fig.update_layout(
63
- title={
64
- 'text': f"{metric} per Model on {task} tasks",
65
- 'x': 0.5,
66
- 'xanchor': 'center',
67
- 'yanchor': 'top',
68
- 'font': dict(size=20),
69
- 'pad': dict(t=0, b=0),
70
- },
71
- margin=dict(t=10, b=10, l=150, r=10),
72
- yaxis_tickfont_size=10,
73
- xaxis=dict(range=[0, 1])
74
- )
75
-
76
- return fig
77
-
78
- # Plot 2: Overview of the performance of a model
79
- import plotly.express as px
80
-
81
- def plot_model_across_tasks(df, model_name, metric):
82
- df = df.copy()
83
- y_col = f"{metric}_mean"
84
- std_col = f"{metric}_std"
85
-
86
- model_df = df[df["Model"] == model_name]
87
- if model_df.empty or y_col not in model_df.columns:
88
- return px.bar(title="No data found for selected model.")
89
 
90
- model_df = model_df.sort_values(by=y_col, ascending=False)
 
91
 
92
  fig = px.bar(
93
- model_df,
94
- x="Task",
95
- y=y_col,
96
- title=f"{metric} for {model_name} across Tasks",
97
- labels={y_col: metric, "Task": "Task"},
98
  text_auto=".2f",
 
 
 
99
  )
100
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
101
  fig.update_traces(
102
  marker_color="#F97316",
103
- hovertemplate=f"<b>%{{x}}</b><br>{metric}: %{{y:.2f}}<br>Std Dev: %{{customdata[0]:.3f}}",
104
- customdata=model_df[[std_col]].values
105
  )
106
 
107
- fig.update_layout(
108
  title={
109
- 'text': f"{metric} for {model_name} across Tasks",
110
- 'x': 0.5,
111
- 'xanchor': 'center',
112
- 'yanchor': 'top',
113
- 'font': dict(size=20),
114
- 'pad': dict(t=0, b=0),
115
- },
116
- yaxis=dict(range=[0, 1]),
117
- margin=dict(t=50, b=100)
 
 
 
 
118
  )
119
 
 
 
 
 
 
 
 
120
  return fig
 
1
  import plotly.express as px
2
  import pandas as pd
3
 
4
+ METRICS_FOR_PLOTS = ["Accuracy", "MCC"] # TO DO: Add F1 later (to leaderboard as a column) and also here to make it available for plots
5
+
6
  def extract_mean_std(col):
7
  if isinstance(col, str) and '±' in col:
8
  try:
 
15
  return None, None
16
 
17
  def prepare_leaderboard_df(df):
18
+ for metric in METRICS_FOR_PLOTS:
19
  means, stds = zip(*df[metric].apply(extract_mean_std))
20
  df[f"{metric}_mean"] = means
21
  df[f"{metric}_std"] = stds
22
  return df
23
 
24
+ def make_plot_wrapper(leaderboard_df, group_by, filter_col, orientation):
25
+ def plot_fn(filter_val, metric):
26
+ return plot_metric_bar(
27
+ df=leaderboard_df,
28
+ group_by=group_by,
29
+ filter_col=filter_col,
30
+ filter_val=filter_val,
31
+ metric=metric,
32
+ orientation=orientation
33
+ )
34
+ return plot_fn
35
+
36
+ def plot_metric_bar(df, group_by, filter_col, filter_val, metric, orientation="v"):
37
  df = df.copy()
38
+ df = df[df[filter_col] == filter_val]
39
 
40
  y_col = f"{metric}_mean"
41
  std_col = f"{metric}_std"
42
 
43
  if df.empty or y_col not in df.columns:
44
+ return px.bar(title="No data found.")
45
 
46
  df = df.sort_values(by=y_col, ascending=False)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
47
 
48
+ x_axis, y_axis = (y_col, group_by) if orientation == "h" else (group_by, y_col)
49
+ height = 20 * len(df) + 200 if orientation == "h" else None
50
 
51
  fig = px.bar(
52
+ df,
53
+ x=x_axis,
54
+ y=y_axis,
55
+ orientation=orientation,
 
56
  text_auto=".2f",
57
+ labels={y_col: metric, group_by: group_by},
58
+ title=f"{metric} for {filter_val}",
59
+ height=height,
60
  )
61
 
62
+ if orientation == "h":
63
+ hovertemplate = (
64
+ "<b>%{y}</b><br>"
65
+ f"{metric}: %{{x:.2f}}<br>"
66
+ "Std Dev: %{customdata[0]:.3f}"
67
+ )
68
+ else:
69
+ hovertemplate = (
70
+ "<b>%{x}</b><br>"
71
+ f"{metric}: %{{y:.2f}}<br>"
72
+ "Std Dev: %{customdata[0]:.3f}"
73
+ )
74
+
75
+
76
  fig.update_traces(
77
  marker_color="#F97316",
78
+ hovertemplate=hovertemplate,
79
+ customdata=df[[std_col]].values
80
  )
81
 
82
+ layout_args = dict(
83
  title={
84
+ 'text': f"{metric} for {filter_val}",
85
+ 'x': 0.5,
86
+ 'xanchor': 'center',
87
+ 'yanchor': 'top',
88
+ 'font': dict(size=20),
89
+ 'pad': dict(t=0, b=0),
90
+ },
91
+ margin=dict(t=10, b=100, l=150, r=10),
92
+ hoverlabel=dict(
93
+ font=dict(color="white"),
94
+ bgcolor="#F97316",
95
+ bordercolor="#c5580d"
96
+ )
97
  )
98
 
99
+ if orientation == "h":
100
+ layout_args["xaxis"] = dict(range=[0, 1])
101
+ layout_args["yaxis_tickfont_size"] = 10
102
+ else:
103
+ layout_args["yaxis"] = dict(range=[0, 1])
104
+
105
+ fig.update_layout(**layout_args)
106
  return fig