from operator import is_
import pandas as pd
import gradio as gr
import os
import requests
from dotenv import load_dotenv
from matplotlib.colors import LinearSegmentedColormap
import plotly.graph_objects as go
import numpy as np
from huggingface_hub import HfApi
from huggingface_hub.hf_api import HTTPError
from huggingface_hub.utils import GatedRepoError
from gradio_rangeslider import RangeSlider
import datetime
from title import css, TITLE_HTML, SUBTITLE_HTML
from data_manager import DataManager
load_dotenv()
webhook_url = os.environ.get("WEBHOOK_URL")
metric_list = [
"Compression Rate (%)",
"Bits Per Character (BPC)",
"Bits Per Byte (BPB)",
]
model_size_list = [
"~14B",
"~9B",
"~7B",
"~3B",
"~1.5B",
"Other",
]
metric_to_sheet = {
"Compression Rate (%)": "cr",
"Bits Per Character (BPC)": "bpc",
"Bits Per Byte (BPB)": "bpb",
}
model_size_to_file_name = {
"~14B": "14b",
"~9B": "9b",
"~7B": "7b",
"~3B": "3b",
"~1.5B": "1b5",
"Other": "other",
}
def read_about_md():
with open("about.md", "r", encoding="utf-8") as f:
return f.read()
def update_table(
data_manager: DataManager,
period: str,
models_size: list,
metric: str,
visible_columns: list,
color_columns: list,
size_range: list,
midpoint: float = 0.5,
ascending: bool = True,
request: gr.Request = None,
):
is_dark_mode = request.is_dark if request else False
print(
f"Updating - time: {datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')}, period: {period}, models: {models_size}, metric: {metric}, visible_columns: {visible_columns}, color_columns: {color_columns}, size_range: {size_range}, ascending: {ascending}, is_dark: {is_dark_mode}\n"
)
target_file_name = [model_size_to_file_name[model] for model in models_size]
metric_code = metric_to_sheet[metric]
# 过滤掉不在当前 period 可用列中的列名,避免错误
if visible_columns:
available_columns = data_manager.get_available_columns(period)
visible_columns = [col for col in visible_columns if col in available_columns]
filtered_data = data_manager.query(
period=period,
metric_code=metric_code,
param_range=(size_range[0], size_range[1]),
model_groups=target_file_name,
visible_columns=visible_columns,
)
if len(filtered_data) == 0:
return "No data available for the selected models and period."
colors = ["#2ca02c", "#2b2b2b", "#d62728"] if is_dark_mode else ["#63be7b", "#ffffff", "#f8696b"]
vmin, vmax, vmid = {}, {}, {}
for column in filtered_data.columns:
if column in ["Name", "Params (B)"]:
continue
col_values = filtered_data[column].dropna()
if len(col_values) > 1:
sorted_values = np.sort(col_values)
vmin[column] = sorted_values.min()
vmax[column] = sorted_values.max()
idx = int(len(sorted_values) * midpoint)
vmid[column] = sorted_values[idx]
def custom_background_gradient(series, cmap, vmin_val, vmax_val, vmid_val):
if len(series) == 0:
return series
def normalize(x):
if pd.isna(x):
return 0.5 # Neutral for NaN
if vmid_val == vmin_val and x <= vmid_val:
return 0.0
if vmid_val == vmax_val and x >= vmid_val:
return 1.0
if vmid_val == vmin_val or vmid_val == vmax_val:
return 0.5
if x <= vmid_val:
return 0.5 * (x - vmin_val) / (vmid_val - vmin_val)
else:
return 0.5 + 0.5 * (x - vmid_val) / (vmax_val - vmid_val)
normed = series.apply(normalize)
cmap_colors = [cmap(x) for x in normed]
return ["background-color: rgba({}, {}, {}, {}); color: black;".format(*[int(255 * c) for c in color[:3]], color[3]) for color in cmap_colors]
target_color_columns = []
if "Average" in color_columns:
target_color_columns.append("Average (lower=better)")
if "Individual Tests" in color_columns:
target_color_columns.extend([col for col in filtered_data.columns if col not in ["Name", "Params (B)", "Average (lower=better)"]])
def color_params_column_dynamic(value):
if not pd.notna(value):
return "default"
if is_dark_mode:
return "background-color: #4b4936; color: #f0f0f0;"
else:
return "background-color: #fffdd0; color: black;"
formatter = {col: "{:.3f}" for col in filtered_data.columns if filtered_data[col].dtype in ["float64", "float32"]}
styler = filtered_data.style.format(formatter)
styler = styler.map(color_params_column_dynamic, subset=["Params (B)"])
for column in target_color_columns:
if column in vmin:
custom_cmap = LinearSegmentedColormap.from_list("custom_cmap", colors)
styler = styler.apply(
custom_background_gradient, cmap=custom_cmap, vmin_val=vmin[column], vmax_val=vmax[column], vmid_val=vmid[column], subset=[column]
)
styler = styler.hide(axis="index")
widths = [250, 80, 80, 70, 70, 70, 70, 70, 70, 70, 70, 70, 70, 70]
table_styles = []
table_styles.append(
{
"selector": "th",
"props": [
("background-color", "var(--background-fill-secondary)"),
("color", "var(--body-text-color)"),
("padding", "8px"),
("font-weight", "bold"),
],
}
)
table_styles.append({"selector": "table", "props": [("border-collapse", "collapse"), ("border", f"1px solid var(--border-color-primary)")]})
for i, w in enumerate(widths):
table_styles.append(
{
"selector": f"th.col{i}, td.col{i}",
"props": [
("min-width", f"{w}px"),
("max-width", f"{w}px"),
("text-align", "center"),
("border", f"1px solid var(--border-color-primary)"),
],
}
)
styler = styler.set_table_styles(table_styles)
return styler.to_html()
def check_model_exists(model_id):
api = HfApi()
try:
model_info = api.model_info(model_id)
return "Exists and is accessible"
except GatedRepoError:
return "Exists but is restricted"
except HTTPError as e:
if e.response.status_code == 404:
return "Does not exist"
else:
return "Error: " + str(e)
def submit_model(name):
if "Exists" not in check_model_exists(name):
return f"# ERROR: Model {name} does not exist on Hugging Face!"
try:
response = requests.post(webhook_url, json={"content": name})
if response.status_code == 200:
response_data = response.json()
if response_data.get("status") == "success":
return "# SUCCESS: We will check the model as soon as possible. Thank you for your submission!"
else:
return f"# ERROR: {response_data.get('message', 'Unknown error')}"
else:
return f"# ERROR: Failed to submit model {name}. Server returned status code {response.status_code}."
except requests.exceptions.HTTPError:
return "# ERROR: Network error while contacting queue. Please try again in a few minutes."
except Exception as e:
print(e)
return "ERROR: Unexpected error. Please try again later."
def create_scaling_plot(data_manager: DataManager, period: str):
new_df = data_manager.query(
period=period,
metric_code="cr",
param_range=(0, 40),
model_groups=None,
visible_columns=None,
)
if len(new_df) == 0:
fig = go.Figure()
fig.update_layout(title={"text": "Compression Rate Scaling Law", "x": 0.5}, width=800, height=600)
return fig
x_values = new_df["Params (B)"].astype(float).tolist()
y_values = new_df["Average (lower=better)"].astype(float).tolist()
names = new_df["Name"].tolist()
x_min, x_max = np.log10(min(x_values)), np.log10(max(x_values))
y_min, y_max = np.log10(min(y_values)), np.log10(max(y_values))
x_dtick = (x_max - x_min) / 4
y_dtick = (y_max - y_min) / 4
fig = go.Figure()
fig.add_trace(
go.Scatter(
x=x_values,
y=y_values,
mode="markers",
name="Models",
marker=dict(size=12, color="#39C5BB", opacity=0.8),
text=names,
customdata=list(zip(x_values, y_values)),
hovertemplate=(
"%{text}
" + "Params: %{customdata[0]:.2f}B
" + "Compression Rate: %{customdata[1]:.2f}%
" + ""
),
)
)
fig.update_layout(
title={"text": "Compression Rate Scaling Law", "x": 0.5, "xanchor": "center", "yanchor": "top"},
width=800,
height=600,
showlegend=True,
xaxis=dict(
title="Parameters (B)",
showgrid=True,
zeroline=False,
type="log",
dtick=x_dtick,
tickformat=".2f",
range=[x_min - 0.1, x_max + 0.1],
),
yaxis=dict(
title="Compression Rate (%)",
showgrid=True,
zeroline=False,
type="log",
dtick=y_dtick,
tickformat=".2f",
range=[y_min - 0.1, y_max + 0.1],
autorange="reversed",
),
)
return fig
if __name__ == "__main__":
data_manager = DataManager("data")
time_list = data_manager.get_available_periods()
last_period = time_list[-1]
initial_fig = create_scaling_plot(data_manager, last_period) if last_period else go.Figure()
initial_metric = metric_list[0]
initial_columns = data_manager.get_available_columns(last_period)
initial_colors = ["Average", "Individual Tests"]
initial_size_range = [0, 40]
initial_data = update_table(data_manager, last_period, model_size_list, initial_metric, initial_columns, initial_colors, initial_size_range)
theme = gr.themes.Default()
with gr.Blocks(theme=theme, css=css) as demo:
gr.HTML(TITLE_HTML)
gr.HTML(SUBTITLE_HTML)
with gr.Tabs() as tabs:
with gr.Tab("🏆 Leaderboard"):
with gr.Row():
with gr.Column():
period_selector = gr.Dropdown(label="Period", choices=time_list, value=last_period)
model_selector = gr.CheckboxGroup(label="Model Size", choices=model_size_list, value=model_size_list)
size_range_slider = RangeSlider(minimum=0, maximum=40, value=[0, 40], step=0.1, label="Model Size Range")
metric_selector = gr.Dropdown(label="Metric", choices=metric_list, value=initial_metric)
with gr.Column():
midpoint_slider = gr.Slider(minimum=0.1, maximum=0.9, value=0.5, step=0.01, label="Color Gradient Midpoint")
color_selector = gr.CheckboxGroup(label="Colored Columns", choices=["Average", "Individual Tests"], value=initial_colors)
colfilter = gr.CheckboxGroup(label="Data Source", choices=initial_columns, value=initial_columns)
table = gr.HTML(initial_data)
def update_table_wrapper(period, models_size, metric, visible_columns, color_columns, size_range, midpoint):
return update_table(data_manager, period, models_size, metric, visible_columns, color_columns, size_range, midpoint)
def update_column_choices(period, current_selected):
if not period:
return gr.update(choices=[], value=[])
columns = data_manager.get_available_columns(period)
# 只保留在新 choices 中存在的已选择值
if current_selected:
valid_selected = [col for col in current_selected if col in columns]
# 如果过滤后为空,默认选择所有列(保持默认行为)
if not valid_selected:
valid_selected = columns
else:
# 如果没有当前选择,默认选择所有列(保持默认行为)
valid_selected = columns
return gr.update(choices=columns, value=valid_selected)
shared_inputs = [period_selector, model_selector, metric_selector, colfilter, color_selector, size_range_slider, midpoint_slider]
period_selector.change(update_column_choices, inputs=[period_selector, colfilter], outputs=colfilter)
period_selector.change(update_table_wrapper, inputs=shared_inputs, outputs=table)
model_selector.change(update_table_wrapper, inputs=shared_inputs, outputs=table)
metric_selector.change(update_table_wrapper, inputs=shared_inputs, outputs=table)
colfilter.change(update_table_wrapper, inputs=shared_inputs, outputs=table)
color_selector.change(update_table_wrapper, inputs=shared_inputs, outputs=table)
size_range_slider.change(update_table_wrapper, inputs=shared_inputs, outputs=table)
midpoint_slider.change(update_table_wrapper, inputs=shared_inputs, outputs=table)
with gr.Tab("📚 Long Context"):
gr.Markdown("## Coming soon...")
with gr.Tab("📈 Scaling Law"):
period_selector_2 = gr.Dropdown(label="Period", choices=time_list, value=last_period)
def update_plot(period):
new_fig = create_scaling_plot(data_manager, period)
return new_fig
plot = gr.Plot(initial_fig)
period_selector_2.change(update_plot, inputs=period_selector_2, outputs=plot)
with gr.Tab("ℹ️ About"):
gr.Markdown(read_about_md())
with gr.Tab("🚀 Submit"):
with gr.Group():
with gr.Row():
model_name = gr.Textbox(max_lines=1, placeholder="Enter model name...", show_label=False, scale=4)
submit = gr.Button("Submit", variant="primary", scale=0)
output = gr.Markdown("# Enter a public HF repo id, then hit Submit to add it to the evaluation queue.")
submit.click(fn=submit_model, inputs=model_name, outputs=output)
demo.launch(share=False)