Team Ai
Apppublic

adelevett/Flashcard2Audio

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
app.py545 linesDownload Raw Back to root
1import gradio as gr2import pandas as pd3import genanki4from pocket_tts import TTSModel5import tempfile6import os7import shutil8import random9import zipfile10import sqlite311import re12import time13import json14import torch15import scipy.io.wavfile16from pathlib import Path17from concurrent.futures import ThreadPoolExecutor, as_completed18from pydub import AudioSegment19 20# --- Configuration ---21MAX_WORKERS = 4        # Keep low for HF Spaces (CPU/RAM constraint)22PREVIEW_LIMIT = 100    # UI safety cap23PROGRESS_THROTTLE = 1.0 # Seconds between UI updates24 25# --- Helpers ---26 27def clean_text_for_tts(text):28    """Deep cleaning for TTS input only."""29    if pd.isna(text): return ""30    text = str(text)31    # Remove HTML tags32    text = re.sub(re.compile('<.*?>'), '', text)33    # Remove Anki sound tags34    text = re.sub(r'\[sound:.*?\]', '', text)35    # Remove mustache templates36    text = re.sub(r'\{\{.*?\}\}', '', text)37    return text.strip()38 39def has_existing_audio(text):40    """Check if text already contains an Anki sound tag."""41    if pd.isna(text): return False42    return bool(re.search(r'\[sound:.*?\]', str(text)))43 44print("Loading TTS Model...")45try:46    TTS_MODEL = TTSModel.load_model()47    print("Model Loaded Successfully.")48except Exception as e:49    print(f"CRITICAL ERROR loading model: {e}")50    TTS_MODEL = None51    52# Get default voice state53VOICE_STATE = None54if TTS_MODEL:55    try:56        VOICE_STATE = TTS_MODEL.get_state_for_audio_prompt("alba")  # Default voice57    except Exception as e:58        print(f"Warning: Could not load default voice: {e}")59 60def wav_to_mp3(src_wav, dst_mp3):61    AudioSegment.from_wav(src_wav).export(dst_mp3, format="mp3", bitrate="64k")62 63def generate_audio_for_row(q_text, a_text, idx, tmpdir, mode):64    """65    Generates audio. Returns (path_q, path_a). 66    Returns 'SKIP' if audio exists and we are preserving it.67    """68    q_out, a_out = None, None69    70    # Logic for handling modes71    # Mode 0: Smart Fill (Preserve Existing)72    # Mode 1: Overwrite All73    74    overwrite = (mode == "Generate all new audio (Overwrite)")75    76    # --- Question Processing ---77    if not overwrite and has_existing_audio(q_text):78        q_out = "SKIP"79    else:80        q_wav = os.path.join(tmpdir, f"q_{idx}.wav")81        try:82            clean = clean_text_for_tts(q_text)83            if clean and TTS_MODEL and VOICE_STATE:84                # Generate audio using new API85                audio_tensor = TTS_MODEL.generate_audio(VOICE_STATE, clean)86                # Convert tensor to numpy and save as wav87                scipy.io.wavfile.write(q_wav, TTS_MODEL.sample_rate, audio_tensor.numpy())88                q_out = q_wav89            else:90                AudioSegment.silent(duration=500).export(q_wav, format="wav")91                q_out = q_wav92        except Exception as e:93            print(f"TTS Error Q row {idx}: {e}")94            # Fallback to silence to keep deck integrity95            AudioSegment.silent(duration=500).export(q_wav, format="wav")96            q_out = q_wav97 98    # --- Answer Processing ---99    if not overwrite and has_existing_audio(a_text):100        a_out = "SKIP"101    else:102        a_wav = os.path.join(tmpdir, f"a_{idx}.wav")103        try:104            clean = clean_text_for_tts(a_text)105            if clean and TTS_MODEL and VOICE_STATE:106                # Generate audio using new API107                audio_tensor = TTS_MODEL.generate_audio(VOICE_STATE, clean)108                # Convert tensor to numpy and save as wav109                scipy.io.wavfile.write(a_wav, TTS_MODEL.sample_rate, audio_tensor.numpy())110                a_out = a_wav111            else:112                AudioSegment.silent(duration=500).export(a_wav, format="wav")113                a_out = a_wav114        except Exception as e:115            print(f"TTS Error A row {idx}: {e}")116            AudioSegment.silent(duration=500).export(a_wav, format="wav")117            a_out = a_wav118            119    return q_out, a_out120 121def strip_html_for_display(text):122    """Remove HTML tags for preview readability."""123    if pd.isna(text) or text == "": return ""124    text = str(text)125    # Remove HTML tags126    text = re.sub(r'<[^>]+>', '', text)127    # Decode HTML entities128    text = text.replace('&nbsp;', ' ').replace('&gt;', '>').replace('&lt;', '<').replace('&amp;', '&')129    # Limit length for display130    if len(text) > 50:131        text = text[:50] + '...'132    return text.strip()133 134def extract_unique_tags(df):135    """Extract all unique tags from the Tags column."""136    if df is None or 'Tags' not in df.columns:137        return ["All"]138    139    all_tags = set()140    for tag_str in df['Tags']:141        if tag_str:142            # Tags are space-separated, e.g., " MK_MathematicsKnowledge "143            tags = [t.strip() for t in tag_str.split() if t.strip()]144            all_tags.update(tags)145    146    return ["All"] + sorted(list(all_tags))147 148def parse_file(file_obj):149    if file_obj is None:150        return None, None, None, "No file uploaded", "", None151    152    ext = Path(file_obj.name).suffix.lower()153    df = pd.DataFrame()154    extract_root = None # Directory where we keep original media155    has_media = False156    157    try:158        if ext == ".csv":159            df = pd.read_csv(file_obj.name)160            if len(df.columns) < 2:161                df = pd.read_csv(file_obj.name, header=None)162            if len(df.columns) < 2:163                return None, None, None, "CSV error: Need 2 columns", "", None164                165            df = df.iloc[:, :2]166            df.columns = ["Question", "Answer"]167            df['Tags'] = ""  # CSV files don't have tags168            169        elif ext == ".apkg" or ext == ".zip":170            # Extract to a PERSISTENT temp dir (passed to state)171            extract_root = tempfile.mkdtemp()172            with zipfile.ZipFile(file_obj.name, 'r') as z:173                z.extractall(extract_root)174            175            col_path = os.path.join(extract_root, "collection.anki2")176            if not os.path.exists(col_path):177                shutil.rmtree(extract_root)178                return None, None, None, "Invalid APKG: No collection.anki2", "", None179                180            conn = sqlite3.connect(col_path)181            cur = conn.cursor()182            cur.execute("SELECT flds, tags FROM notes")183            rows = cur.fetchall()184            185            data = []186            audio_count = 0  # Count cards with existing audio187            for r in rows:188                flds = r[0].split('\x1f')189                tags = r[1].strip() if len(r) > 1 else ""190                q = flds[0] if len(flds) > 0 else ""191                a = flds[1] if len(flds) > 1 else ""192                # Check if either field has audio tags193                if re.search(r'\[sound:.*?\]', q) or re.search(r'\[sound:.*?\]', a):194                    audio_count += 1195                data.append([q, a, tags])196            197            df = pd.DataFrame(data, columns=["Question", "Answer", "Tags"])198            conn.close()199            200            # has_media means existing AUDIO, not images201            has_media = audio_count > 0202            203        else:204            return None, None, None, "Unsupported file type", "", None205        206        df = df.fillna("")207        208        msg = f"โœ… Loaded {len(df)} cards."209        if has_media:210            msg += f" ๐ŸŽต {audio_count} cards have existing audio."211            212        return df, has_media, df.head(PREVIEW_LIMIT), msg, estimate_time(len(df), has_media), extract_root213        214    except Exception as e:215        if extract_root and os.path.exists(extract_root):216            shutil.rmtree(extract_root)217        return None, None, None, f"Error: {str(e)}", "", None218 219def estimate_time(num_cards, has_existing_media=False, mode="Smart Fill (Preserve Existing)"):220    """221    Estimate based on benchmark: ~4.7s per card for full generation.222    Adjusts for Smart Fill mode when existing media is present.223    """224    if num_cards == 0:225        return "0s"226    227    # Base benchmark: 4.7s per card for full audio generation228    seconds_per_card = 4.7229    230    # If using Smart Fill with existing media, assume ~50% speedup (many cards already have audio)231    if has_existing_media and "Smart Fill" in mode:232        seconds_per_card *= 0.5233    234    seconds = num_cards * seconds_per_card235    236    if seconds < 60: 237        return f"~{int(seconds)}s"238    elif seconds < 3600:239        return f"~{int(seconds/60)} min"240    else:241        hours = int(seconds / 3600)242        mins = int((seconds % 3600) / 60)243        return f"~{hours}h {mins}m" if mins > 0 else f"~{hours}h"244 245def process_dataframe(df_full, search_term, extract_root, mode, search_in, selected_tag, progress=gr.Progress()):246    if df_full is None or len(df_full) == 0:247        return None, "No data"248    249    # Start with full dataframe250    df = df_full.copy()251    252    # Apply tag filter first253    if selected_tag and selected_tag != "All":254        df = df[df['Tags'].str.contains(selected_tag, na=False, case=False)]255    256    # Apply text search filter257    if search_term:258        if search_in == "Question Only":259            mask = df['Question'].str.contains(search_term, case=False, na=False)260        elif search_in == "Answer Only":261            mask = df['Answer'].str.contains(search_term, case=False, na=False)262        else:  # Both263            mask = df.astype(str).apply(lambda x: x.str.contains(search_term, case=False, na=False)).any(axis=1)264        df = df[mask]265 266    if len(df) == 0:267        return None, "No matching cards"268 269    # Setup270    work_dir = tempfile.mkdtemp()271    media_files = [] 272    273    try:274        # --- Media Preservation Logic ---275        if extract_root:276            media_map_path = os.path.join(extract_root, "media")277            if os.path.exists(media_map_path) and os.path.getsize(media_map_path) > 0:278                try:279                    with open(media_map_path, 'r') as f:280                        # Fix: Handle potentially malformed JSON gracefully281                        content = f.read().strip()282                        if content:283                            media_map = json.loads(content) # {"0": "my_audio.mp3", ...}284                            285                            # Rename files in extract_root back to original names286                            for k, v in media_map.items():287                                src = os.path.join(extract_root, k)288                                dst = os.path.join(extract_root, v)289                                if os.path.exists(src):290                                    # Rename enables genanki to find them by name291                                    os.rename(src, dst)292                                    media_files.append(dst)293                        else:294                            print("Warning: Media map file is empty.")295                except Exception as e:296                    print(f"Warning: Could not restore existing media: {e}")297 298        # --- Genanki Setup ---299        model_id = random.randrange(1 << 30, 1 << 31)300        my_model = genanki.Model(301            model_id, 'PocketTTS Model',302            fields=[{'name': 'Question'}, {'name': 'Answer'}],303            templates=[{304                'name': 'Card 1',305                'qfmt': '{{Question}}',306                'afmt': '{{FrontSide}}<hr id="answer">{{Answer}}',307            }])308        my_deck = genanki.Deck(random.randrange(1 << 30, 1 << 31), 'Pocket TTS Deck')309 310        # --- Execution ---311        total = len(df)312        completed = 0313        last_update_time = 0314        315        with ThreadPoolExecutor(max_workers=MAX_WORKERS) as exe:316            futures = {}317            for idx, row in df.iterrows():318                f = exe.submit(generate_audio_for_row, str(row['Question']), str(row['Answer']), idx, work_dir, mode)319                futures[f] = idx320 321            for future in as_completed(futures):322                idx = futures[future]323                try:324                    q_res, a_res = future.result()325                    326                    # --- Field Construction (Corrected) ---327                    q_original = str(df.loc[idx, 'Question'])328                    q_field = q_original329                    330                    # Update Question331                    if q_res and q_res != "SKIP":332                        q_mp3 = str(Path(q_res).with_suffix('.mp3'))333                        wav_to_mp3(q_res, q_mp3)334                        os.remove(q_res) # clean wav335                        media_files.append(q_mp3)336                        337                        # Remove OLD sound tags first to avoid duplicates338                        q_field = re.sub(r'\[sound:.*?\]', '', q_field)339                        q_field = q_field.strip() + f"<br>[sound:{os.path.basename(q_mp3)}]"340                    341                    # Update Answer342                    a_original = str(df.loc[idx, 'Answer'])343                    a_field = a_original344                    345                    if a_res and a_res != "SKIP":346                        a_mp3 = str(Path(a_res).with_suffix('.mp3'))347                        wav_to_mp3(a_res, a_mp3)348                        os.remove(a_res) # clean wav349                        media_files.append(a_mp3)350                        351                        # Remove OLD sound tags first352                        a_field = re.sub(r'\[sound:.*?\]', '', a_field)353                        a_field = a_field.strip() + f"<br>[sound:{os.path.basename(a_mp3)}]"354 355                    # Add Note356                    note = genanki.Note(357                        model=my_model,358                        fields=[q_field, a_field]359                    )360                    my_deck.add_note(note)361                    362                except Exception as e:363                    print(f"Row {idx} failed: {e}")364                365                # --- Throttled Progress ---366                completed += 1367                current_time = time.time()368                if completed == total or (current_time - last_update_time) > PROGRESS_THROTTLE:369                    progress(completed / total, desc=f"Processed {completed}/{total}")370                    last_update_time = current_time371 372        # --- Package ---373        package = genanki.Package(my_deck)374        # Deduplicate media files list375        package.media_files = list(set(media_files))376        377        raw_out = os.path.join(work_dir, "output.apkg")378        package.write_to_file(raw_out)379        380        final_out = os.path.join(tempfile.gettempdir(), f"pocket_deck_{random.randint(1000,9999)}.apkg")381        shutil.copy(raw_out, final_out)382        383        return final_out, f"โœ… Done! Packaged {len(package.media_files)} audio files."384 385    except Exception as e:386        return None, f"Critical Error: {str(e)}"387 388    finally:389        # --- Guaranteed Cleanup ---390        if os.path.exists(work_dir):391            shutil.rmtree(work_dir)392        # Also clean up the input extraction root if it exists393        if extract_root and os.path.exists(extract_root):394            shutil.rmtree(extract_root)395 396# --- UI ---397 398with gr.Blocks(title="Pocket TTS Anki") as app:399    gr.Markdown("## ๐ŸŽด Pocket TTS Anki Generator")400    gr.Markdown("Offline Neural Audio. Supports CSV and APKG (smart media preservation).")401    402    # State variables403    full_df_state = gr.State()404    extract_root_state = gr.State() # Holds path to unzipped APKG405 406    with gr.Row():407        file_input = gr.File(label="Upload (CSV/APKG)", file_types=[".csv", ".apkg", ".zip"])408        status = gr.Textbox(label="Status", interactive=False)409        eta_box = gr.Textbox(label="Est. Time", interactive=False)410 411    with gr.Row():412        search_box = gr.Textbox(label="Search Text", placeholder="Enter text to search...")413        search_field = gr.Radio(414            choices=["Both", "Question Only", "Answer Only"],415            value="Both",416            label="Search In"417        )418        tag_dropdown = gr.Dropdown(419            label="Filter by Tag",420            choices=["All"],421            value="All",422            interactive=True423        )424    425    with gr.Row():426        # New 3-Way Toggle427        mode_radio = gr.Radio(428            choices=[429                "Smart Fill (Preserve Existing)", 430                "Generate all new audio (Overwrite)",431                "Only generate missing (Same as Smart Fill)" 432            ],433            value="Smart Fill (Preserve Existing)",434            label="Generation Mode"435        )436 437    preview_table = gr.Dataframe(438        label="Preview (First 100)", 439        interactive=False,440        column_widths=["30%", "45%", "25%"]441    )442    443    with gr.Row():444        btn = gr.Button("๐Ÿš€ Generate Deck", variant="primary")445        dl = gr.File(label="Download")446    447    result_lbl = gr.Textbox(label="Result", interactive=False)448 449    has_media_state = gr.State(False)450    451    def on_upload(file):452        # Returns: df, has_media, preview, msg, eta, extract_path453        df, has_media, preview, msg, eta, ext_path = parse_file(file)454        455        # Extract tags and create cleaned preview456        tag_choices = extract_unique_tags(df)457        458        if df is not None:459            display_df = df.copy()460            display_df['Question'] = display_df['Question'].apply(strip_html_for_display)461            display_df['Answer'] = display_df['Answer'].apply(strip_html_for_display)462            clean_preview = display_df.head(PREVIEW_LIMIT)463        else:464            clean_preview = preview465        466        return (467            df,                                                      # full_df_state468            has_media,                                               # has_media_state  469            clean_preview,                                           # preview_table470            msg,                                                     # status471            eta,                                                     # eta_box472            ext_path,                                                # extract_root_state473            gr.Dropdown(choices=tag_choices, value="All")           # tag_dropdown474        )475 476    file_input.upload(on_upload, inputs=file_input, 477                      outputs=[full_df_state, has_media_state, preview_table, status, eta_box, extract_root_state, tag_dropdown])478 479    def on_clear():480        """Reset all fields when file is cleared."""481        return (482            None,                                    # full_df_state483            False,                                   # has_media_state484            None,                                    # preview_table485            "",                                      # status486            "",                                      # eta_box487            None,                                    # extract_root_state488            gr.Dropdown(choices=["All"], value="All"),  # tag_dropdown489            "",                                      # search_box490            None,                                    # dl (download file)491            ""                                       # result_lbl492        )493 494    file_input.clear(on_clear, inputs=[], 495                     outputs=[full_df_state, has_media_state, preview_table, status, eta_box, extract_root_state, tag_dropdown, search_box, dl, result_lbl])496 497    def on_search(term, df, has_media, mode, search_in, selected_tag):498        if df is None: return None, "No data"499        500        filtered_df = df.copy()501        502        # Apply tag filter first503        if selected_tag and selected_tag != "All":504            filtered_df = filtered_df[filtered_df['Tags'].str.contains(selected_tag, na=False, case=False)]505        506        # Apply text search507        if term:508            if search_in == "Question Only":509                mask = filtered_df['Question'].str.contains(term, case=False, na=False)510            elif search_in == "Answer Only":511                mask = filtered_df['Answer'].str.contains(term, case=False, na=False)512            else:  # Both513                mask = filtered_df.astype(str).apply(lambda x: x.str.contains(term, case=False, na=False)).any(axis=1)514            filtered_df = filtered_df[mask]515        516        # Create cleaned display version517        display_df = filtered_df.copy()518        display_df['Question'] = display_df['Question'].apply(strip_html_for_display)519        display_df['Answer'] = display_df['Answer'].apply(strip_html_for_display)520        521        return display_df.head(PREVIEW_LIMIT), estimate_time(len(filtered_df), has_media, mode)522 523    search_box.change(on_search, inputs=[search_box, full_df_state, has_media_state, mode_radio, search_field, tag_dropdown], 524                     outputs=[preview_table, eta_box])525    526    search_field.change(on_search, inputs=[search_box, full_df_state, has_media_state, mode_radio, search_field, tag_dropdown],527                     outputs=[preview_table, eta_box])528    529    tag_dropdown.change(on_search, inputs=[search_box, full_df_state, has_media_state, mode_radio, search_field, tag_dropdown],530                     outputs=[preview_table, eta_box])531    532    mode_radio.change(on_search, inputs=[search_box, full_df_state, has_media_state, mode_radio, search_field, tag_dropdown],533                     outputs=[preview_table, eta_box])534 535    btn.click(process_dataframe, 536              inputs=[full_df_state, search_box, extract_root_state, mode_radio, search_field, tag_dropdown], 537              outputs=[dl, result_lbl])538 539if __name__ == "__main__":540    app.queue(max_size=2).launch(541        server_name="0.0.0.0",542        server_port=7860,543        ssr_mode=False544    )545