import pandas as pd
import jiwer
import plotly.express as px
import plotly.graph_objects as go
from constants import CSV_FILE, TABLE_HEADERS
# Ses işleme için (Gradio otomatik halletse de importlar iyi olur)
import os
# --- YARDIMCI FONKSİYONLAR (Madalya, Badge, Link) ---
def get_medal(rank):
if rank == 1: return "🥇"
if rank == 2: return "🥈"
if rank == 3: return "🥉"
return f"{rank}"
def get_badge(value, high_is_good=False, is_rtf=False):
try:
val = float(value)
if is_rtf: # RTF için farklı eşikler (Düşük iyi)
if val < 0.05: css = "badge-excellent"
elif val < 0.10: css = "badge-good"
elif val < 0.20: css = "badge-average"
else: css = "badge-poor"
return f"{val:.2f}x"
else: # WER/CER için (Düşük iyi)
if val < 10.0: css = "badge-excellent"
elif val < 20.0: css = "badge-good"
elif val < 50.0: css = "badge-average"
else: css = "badge-poor"
return f"{val:.2f}%"
except: return str(value)
def make_clickable_model(model_name, link):
if pd.isna(link) or link == "" or str(link).lower() == "nan": return model_name
return f"{model_name}"
def parse_params_to_num(param_str):
"""'1.55B', '769M' gibi stringleri sayıya çevirir."""
try:
if pd.isna(param_str): return 0
param_str = str(param_str).upper()
if param_str.endswith('B'): return float(param_str[:-1]) * 1e9
if param_str.endswith('M'): return float(param_str[:-1]) * 1e6
if param_str.endswith('K'): return float(param_str[:-1]) * 1e3
return float(param_str)
except: return 0
# --- VERİ YÜKLEME VE İŞLEME ---
def load_raw_data():
try:
df = pd.read_csv(CSV_FILE)
# Sayısal sütunları temizle
for col in df.columns:
if "WER" in col or "CER" in col or "RTF" in col:
df[col] = pd.to_numeric(df[col], errors='coerce')
# Parametreleri sayıya çevir (Grafik için)
df['Params_Num'] = df['Params'].apply(parse_params_to_num)
# Ortalama WER'e göre sırala
if "Average_WER" in df.columns:
df = df.sort_values(by="Average_WER", ascending=True).reset_index(drop=True)
return df
except Exception as e:
print(f"Data loading error: {e}")
return pd.DataFrame()
# --- GELİŞMİŞ FİLTRELEME ---
def filter_data(search_query, model_types, max_params):
df = load_raw_data()
if df.empty: return df
# 1. Metin Arama
if search_query:
df = df[df.apply(lambda row: row.astype(str).str.contains(search_query, case=False).any(), axis=1)]
# 2. Model Tipi Filtresi
if model_types:
df = df[df['Model_Type'].isin(model_types)]
# 3. Parametre Boyutu Filtresi (Slider değeri Milyon cinsinden)
if max_params < 2000: # Slider maksimumda değilse filtrele
max_params_val = max_params * 1e6
df = df[df['Params_Num'] <= max_params_val]
return df.reset_index(drop=True)
# --- HTML TABLO OLUŞTURMA ---
def dataframe_to_html(df_filtered):
if df_filtered.empty: return "
No models found matching criteria.
"
# Orijinal linkleri almak için ham veriyi oku
df_raw = pd.read_csv(CSV_FILE)
html = ""
for header in TABLE_HEADERS: html += f"| {header} | "
html += "
"
for index, row in df_filtered.iterrows():
rank = index + 1
medal = get_medal(rank)
model_name = row['Model']
try:
link = df_raw.loc[df_raw['Model'] == model_name, 'Model_Link'].values[0]
model_link = make_clickable_model(model_name, link)
except: model_link = model_name
# Rozetleri oluştur
cv_wer = get_badge(row.get('CommonVoice_WER', '-'))
cv_cer = row.get('CommonVoice_CER', '-')
fl_wer = get_badge(row.get('FLEURS_WER', '-'))
fl_cer = row.get('FLEURS_CER', '-')
avg_wer = get_badge(row.get('Average_WER', '-'))
rtf = get_badge(row.get('RTF', '-'), is_rtf=True)
html += f"""
| {medal} | {model_link} |
{cv_wer} | {cv_cer} | {fl_wer} | {fl_cer} |
{avg_wer} | {row.get('Params','-')} | {rtf} |
{row.get('Model_Type','-')} | {row.get('License','-')} |
"""
html += "
"
return html
# --- PLOTLY GRAFİKLERİ ---
def create_plots(df):
if df.empty: return None, None
# 1. Boyut vs. Başarım (Bubble Chart)
fig_size = px.scatter(
df, x="Params_Num", y="Average_WER",
size="Params_Num", color="Model_Type",
hover_name="Model", text="Model",
labels={"Params_Num": "Parameters (Log Scale)", "Average_WER": "Average WER (%)", "Model_Type": "Type"},
title="🤖 Model Size vs. Performance (Lower-Left is Better)",
log_x=True, height=400,
color_discrete_sequence=px.colors.qualitative.Prism # Mavi-Yeşil temaya uygun renkler
)
fig_size.update_traces(textposition='top center', marker=dict(opacity=0.7, line=dict(width=1, color='DarkSlateGrey')))
fig_size.update_layout(
plot_bgcolor='rgba(0,0,0,0)', paper_bgcolor='rgba(0,0,0,0)',
font_color="var(--text-primary)",
xaxis=dict(showgrid=True, gridcolor='var(--border-color)'),
yaxis=dict(showgrid=True, gridcolor='var(--border-color)')
)
# 2. Hız vs. Başarım (Scatter Plot)
fig_speed = px.scatter(
df, x="RTF", y="Average_WER",
color="Model_Type", hover_name="Model", text="Model",
labels={"RTF": "Speed (RTF - Lower is Faster)", "Average_WER": "Average WER (%)"},
title="⚡ Speed vs. Performance (Lower-Left is Better)",
height=400, color_discrete_sequence=px.colors.qualitative.Prism
)
fig_speed.update_traces(textposition='top center', marker=dict(size=12, opacity=0.8))
fig_speed.update_layout(
plot_bgcolor='rgba(0,0,0,0)', paper_bgcolor='rgba(0,0,0,0)',
font_color="var(--text-primary)",
xaxis=dict(showgrid=True, gridcolor='var(--border-color)'),
yaxis=dict(showgrid=True, gridcolor='var(--border-color)')
)
return fig_size, fig_speed
# --- METİN SANDBOX HESAPLAMA ---
def calculate_metrics_text(ref, hyp, norm=True):
if not ref or not hyp: return "-","-","Please fill both areas."
if norm:
trans = jiwer.Compose([jiwer.ToLowerCase(), jiwer.RemovePunctuation(), jiwer.RemoveMultipleSpaces(), jiwer.Strip()])
ref = trans(ref); hyp = trans(hyp)
try:
wer = jiwer.wer(ref, hyp)*100; cer = jiwer.cer(ref, hyp)*100
return f"{wer:.2f}%", f"{cer:.2f}%", f"**Ref:** {ref}\n\n**Hyp:** {hyp}"
except: return "Error", "Error", "Calculation failed."
# --- SESLİ SANDBOX (DEMO MODEL) ---
# Gerçek bir model yüklemek CPU'da yavaş olabilir, bu yüzden 'tiny' kullanıyoruz.
from transformers import pipeline
# Modeli global olarak bir kere yükle (Hafif model)
try:
asr_pipeline = pipeline("automatic-speech-recognition", model="openai/whisper-tiny", device="cpu")
except:
asr_pipeline = None
print("Warning: Could not load Whisper-tiny for audio sandbox.")
def transcribe_audio(audio_path):
if not asr_pipeline: return "ASR Model not loaded successfully."
if audio_path is None: return "Please record or upload audio."
try:
# Gradio audio path'ini doğrudan pipeline'a ver
result = asr_pipeline(audio_path)
return result["text"]
except Exception as e:
return f"Transcription Error: {e}"