File size: 5,411 Bytes
2652037
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
"""Gradio web UI for the phonetic-reduction analyzer.

Run with:
    .venv/bin/python app.py
Then open the URL Gradio prints (defaults to http://127.0.0.1:7860).
"""
from __future__ import annotations

import os
import tempfile

import gradio as gr
import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import pandas as pd

from reduction import pipeline, stt, tts


# --- model cache ---------------------------------------------------------
_models: dict[str, object] = {}
_voice = None


def _get_model(size: str):
    if size not in _models:
        _models[size] = stt.load_model(size=size)
    return _models[size]


def _get_voice():
    global _voice
    if _voice is None:
        _voice = tts.load_voice()
    return _voice


# --- plot ----------------------------------------------------------------
def _make_plot(df: pd.DataFrame) -> plt.Figure:
    n = len(df)
    fig, ax = plt.subplots(figsize=(max(8, n * 0.55), 4.5))
    x = list(range(n))
    labels = df["word"].tolist()

    # Colour bars by STT confidence: low-confidence → red, normal → steelblue
    probs = df["stt_probability"].fillna(1.0).tolist()
    bar_colors = ["#cc3333" if p < 0.6 else "#4a90c4" for p in probs]
    ax.bar(x, df["combined_score"], color=bar_colors, alpha=0.55, label="combined")
    ax.plot(x, df["duration_score"], "o-", color="#2c7a2c", lw=1.4, label="duration", alpha=0.85)
    ax.plot(x, df["spectral_score"], "s-", color="#d97726", lw=1.4, label="spectral", alpha=0.85)
    ax.axhline(1.0, color="gray", linestyle="--", linewidth=0.8)

    ax.set_xticks(x)
    ax.set_xticklabels(labels, rotation=40, ha="right", fontsize=10)
    ax.set_ylabel("skóre  (1 = průměr,  <1 redukováno)")
    ax.set_xlabel("token")
    ax.set_title("Per-word reduction scores")
    ax.legend(loc="upper right", framealpha=0.9)
    ax.grid(True, axis="y", alpha=0.25)
    fig.tight_layout()
    return fig


# --- main analyze callback ----------------------------------------------
def analyze_audio(audio_path: str | None, whisper_size: str, progress=gr.Progress()):
    if not audio_path:
        return "❌ Nahraj nebo nahrej zvuk.", None, None, None

    progress(0.1, desc=f"Načítám Whisper {whisper_size}…")
    model = _get_model(whisper_size)
    progress(0.3, desc="Načítám Piper hlas…")
    voice = _get_voice()

    progress(0.5, desc="Analyzuji…")
    rows, whisper_text, _ = pipeline.analyze(
        audio_path, model, voice, language="cs", groundtruth=None, trim=True
    )

    progress(0.9, desc="Renderuji graf a CSV…")
    df = pd.DataFrame(rows)

    # Pretty display dataframe (round numerics, drop internal cols)
    display_cols = [
        "word",
        "stt_probability",
        "orig_duration",
        "tts_duration",
        "duration_ratio",
        "spectral_distance",
        "duration_score",
        "spectral_score",
        "combined_score",
    ]
    display_df = df[display_cols].copy()
    for c in display_cols[1:]:
        display_df[c] = display_df[c].round(3)

    fig = _make_plot(df)

    # CSV download (full row, not abbreviated)
    csv_path = os.path.join(tempfile.gettempdir(), "reduction_result.csv")
    df.to_csv(csv_path, index=False)

    return whisper_text or "(žádný transkript)", display_df, fig, csv_path


# --- UI -----------------------------------------------------------------
HEADER_MD = """
# Měření fonetické redukce — prototyp
Per-word reduction scores via STT → per-word TTS canonical → duration + spectral DTW.
Skóre **< 1** = redukováno oproti průměru nahrávky, **> 1** = plně artikulováno.
Červené sloupce = nízká STT confidence (pravděpodobně přeslech / halucinace).
"""

with gr.Blocks(title="Phonetic reduction analyzer") as demo:
    gr.Markdown(HEADER_MD)
    with gr.Row():
        with gr.Column(scale=2):
            audio_in = gr.Audio(
                sources=["upload", "microphone"],
                type="filepath",
                label="Nahrávka (drag & drop wav, nebo klikni na mikrofon a namluv)",
            )
        with gr.Column(scale=1):
            whisper_dd = gr.Dropdown(
                choices=["small", "medium", "large-v3"],
                value="large-v3",
                label="Whisper model",
                info="large-v3 = nejlepší kvalita, ~6s/věta na CPU",
            )
            run_btn = gr.Button("Analyzuj", variant="primary", size="lg")

    transcript_out = gr.Textbox(label="Whisper transkript", interactive=False, lines=2)
    plot_out = gr.Plot(label="Skóre redukce")
    table_out = gr.Dataframe(label="Per-word detaily", interactive=False)
    csv_out = gr.File(label="Stáhnout CSV", interactive=False)

    run_btn.click(
        analyze_audio,
        inputs=[audio_in, whisper_dd],
        outputs=[transcript_out, table_out, plot_out, csv_out],
    )


if __name__ == "__main__":
    # On HuggingFace Spaces, $SPACE_ID is set; on local dev it isn't.
    on_hf = bool(os.environ.get("SPACE_ID"))
    if on_hf:
        # HF provides PORT and expects 0.0.0.0; let HF's launcher own model preloading
        demo.launch(server_name="0.0.0.0")
    else:
        print("[startup] preloading large-v3 + Piper voice (one-time, ~10s)…")
        _get_model("large-v3")
        _get_voice()
        print("[startup] ready, launching Gradio…")
        demo.launch(server_name="127.0.0.1", server_port=7860, inbrowser=False)