import gradio as gr import pandas as pd import plotly.express as px DATA_URL = "https://huggingface.co/datasets/marin-community/token-counts/resolve/main/data/token_counts.csv" PALETTE = [ "#1877F2", "#F0701A", "#5A24C7", "#E42C97", "#00487C", "#0EAC96", "#AB76FF", "#B50550", "#0099E6", "#22085F", "#783301", ] GRIDCOLOR = "rgba(220,220,220,0.5)" LEGEND = dict(bgcolor="rgba(255,255,255,0.8)", bordercolor="rgba(200,200,200,0.5)") ANNOTATION_FONT_COLOR = "rgba(60,60,60,0.9)" LAYOUT_DEFAULTS = dict( template="plotly_white", font=dict(color=ANNOTATION_FONT_COLOR), legend=LEGEND, xaxis=dict(gridcolor=GRIDCOLOR), yaxis=dict(gridcolor=GRIDCOLOR), ) PROVENANCE_ORDER = ["Organic", "Machine Translated", "Synthetic"] PROVENANCE_COLORS = { "Organic": PALETTE[0], "Machine Translated": PALETTE[5], "Synthetic": PALETTE[1], } def load_data(): df = pd.read_csv(DATA_URL) df["tokens_b"] = df["marin_tokens"] / 1e9 df["provenance"] = df.apply( lambda r: "Machine Translated" if r.get("machine_translated") else ("Synthetic" if r.get("synthetic") else "Organic"), axis=1, ) return df def make_category_bar(df): grouped = df.groupby("category")["tokens_b"].sum().reset_index() grouped = grouped.sort_values("tokens_b", ascending=True) fig = px.bar( grouped, x="tokens_b", y="category", orientation="h", labels={"tokens_b": "Tokens (B)", "category": ""}, title="Tokens by Category", color="category", color_discrete_sequence=PALETTE, ) fig.update_layout(**LAYOUT_DEFAULTS, showlegend=False, height=300, margin=dict(l=0, r=20, t=40, b=0)) return fig def make_provenance_pie(df): counts = df.groupby("provenance")["tokens_b"].sum().reset_index() fig = px.pie( counts, values="tokens_b", names="provenance", title="Data Provenance", category_orders={"provenance": PROVENANCE_ORDER}, color="provenance", color_discrete_map=PROVENANCE_COLORS, ) fig.update_layout(template="plotly_white", font=dict(color=ANNOTATION_FONT_COLOR), legend=LEGEND, height=300, margin=dict(l=0, r=0, t=40, b=0)) return fig def apply_filters(df, category_filter, provenance_filter): filtered = df.copy() if category_filter and category_filter != "All": filtered = filtered[filtered["category"] == category_filter] if provenance_filter and provenance_filter != "All": filtered = filtered[filtered["provenance"] == provenance_filter] return filtered def make_dataset_bar(df, category_filter, provenance_filter): filtered = apply_filters(df, category_filter, provenance_filter) filtered = filtered.sort_values("tokens_b", ascending=True).tail(30) fig = px.bar( filtered, x="tokens_b", y="marin_name", orientation="h", labels={"tokens_b": "Tokens (B)", "marin_name": ""}, title="Tokens by Dataset", color="category", color_discrete_sequence=PALETTE, ) fig.update_layout(**LAYOUT_DEFAULTS, height=max(400, len(filtered) * 22), margin=dict(l=0, r=20, t=40, b=0)) return fig def make_table(df, category_filter, provenance_filter): filtered = apply_filters(df, category_filter, provenance_filter) cols = ["marin_name", "tokens_b", "category", "provenance", "pdf", "hf_repo", "hf_subset", "license", "transform_tldr"] display = filtered[[c for c in cols if c in filtered.columns]].copy() display["tokens_b"] = display["tokens_b"].round(2) display = display.rename(columns={"tokens_b": "tokens (B)"}) display = display.sort_values("tokens (B)", ascending=False) return display def build_ui(): df = load_data() total_b = df["tokens_b"].sum() categories = ["All"] + sorted(df["category"].unique().tolist()) header = f"# Marin Dataset Catalog\n\n**{total_b:,.1f}B** tokens across **{len(df)}** datasets" return df, header, categories with gr.Blocks(title="Marin Dataset Catalog") as demo: header_md = gr.Markdown("") with gr.Row(): cat_plot = gr.Plot() prov_plot = gr.Plot() with gr.Row(): category_dd = gr.Dropdown(["All"], value="All", label="Category") provenance_dd = gr.Dropdown(["All"] + PROVENANCE_ORDER, value="All", label="Provenance") dataset_plot = gr.Plot() table = gr.Dataframe() @demo.load def on_load(): df, header, categories = build_ui() return ( header, make_category_bar(df), make_provenance_pie(df), gr.Dropdown(choices=categories, value="All"), make_dataset_bar(df, "All", "All"), make_table(df, "All", "All"), ) on_load_outputs = [header_md, cat_plot, prov_plot, category_dd, dataset_plot, table] demo.load(on_load, outputs=on_load_outputs) def update(cat, prov): df = load_data() return make_dataset_bar(df, cat, prov), make_table(df, cat, prov) category_dd.change(update, [category_dd, provenance_dd], [dataset_plot, table]) provenance_dd.change(update, [category_dd, provenance_dd], [dataset_plot, table]) demo.launch()