adelevett/Flashcard2Audio
0
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(' ', ' ').replace('>', '>').replace('<', '<').replace('&', '&')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 