main
py 498 lines 19 KB
Raw
1 import streamlit as st
2 import torch
3 import os
4 import tempfile
5 import torchaudio
6 from nemo.collections.asr.models import ASRModel
7 from pydub import AudioSegment
8 import numpy as np
9 import csv
10 import datetime
11 import pandas as pd
12 import time
13 import gc
14
15 # Set page config with a modern theme
16 st.set_page_config(
17 page_title="Parakeet ASR Demo",
18 page_icon="🎙️",
19 layout="wide",
20 initial_sidebar_state="collapsed"
21 )
22
23 # Custom CSS for modern UI
24 st.markdown("""
25 <style>
26 .main {
27 padding: 1rem 2rem;
28 }
29 .stButton>button {
30 width: 100%;
31 background-color: #4CAF50;
32 color: white;
33 padding: 0.5rem 1rem;
34 border: none;
35 border-radius: 4px;
36 font-size: 1rem;
37 font-weight: 500;
38 }
39 .stButton>button:hover {
40 background-color: #45a049;
41 }
42 .stDataFrame {
43 border-radius: 8px;
44 box-shadow: 0 2px 4px rgba(0,0,0,0.1);
45 }
46 h1 {
47 color: #1E88E5;
48 font-size: 2.5rem !important;
49 font-weight: 700 !important;
50 margin: -1rem 0 1rem 0 !important;
51 }
52 h3 {
53 color: #1E88E5;
54 font-size: 1.5rem !important;
55 font-weight: 600 !important;
56 margin-top: 1.5rem !important;
57 }
58 .stProgress > div > div > div > div {
59 background-color: #4CAF50;
60 }
61 .info-box {
62 background-color: rgba(30, 136, 229, 0.1);
63 border-left: 5px solid #1E88E5;
64 padding: 0.5rem;
65 border-radius: 4px;
66 margin-bottom: 0.5rem;
67 font-size: 0.9rem;
68 }
69 .success-box {
70 background-color: rgba(76, 175, 80, 0.1);
71 border-left: 5px solid #4CAF50;
72 padding: 0.5rem;
73 border-radius: 4px;
74 margin-bottom: 0.5rem;
75 font-size: 0.9rem;
76 line-height: 1.2;
77 }
78 .warning-box {
79 background-color: rgba(255, 152, 0, 0.1);
80 border-left: 5px solid #FF9800;
81 padding: 0.5rem;
82 border-radius: 4px;
83 margin-bottom: 0.5rem;
84 font-size: 0.9rem;
85 }
86 .error-box {
87 background-color: rgba(244, 67, 54, 0.1);
88 border-left: 5px solid #F44336;
89 padding: 0.5rem;
90 border-radius: 4px;
91 margin-bottom: 0.5rem;
92 font-size: 0.9rem;
93 }
94 </style>
95 """, unsafe_allow_html=True)
96
97 # Constants
98 MODEL_NAME = "nvidia/parakeet-tdt-0.6b-v2"
99 DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
100 SUPPORTED_FORMATS = ['wav', 'flac']
101 MAX_RECOMMENDED_DURATION = 30 * 60 # 30 minutes in seconds
102 LONG_AUDIO_THRESHOLD = 480 # 8 minutes in seconds
103
104 # Initialize session state variables
105 if 'transcription_history' not in st.session_state:
106 st.session_state.transcription_history = []
107
108 def custom_info(text):
109 st.markdown(f'<div class="info-box">{text}</div>', unsafe_allow_html=True)
110
111 def custom_success(text):
112 st.markdown(f'<div class="success-box">{text}</div>', unsafe_allow_html=True)
113
114 def custom_warning(text):
115 st.markdown(f'<div class="warning-box">{text}</div>', unsafe_allow_html=True)
116
117 def custom_error(text):
118 st.markdown(f'<div class="error-box">{text}</div>', unsafe_allow_html=True)
119
120 @st.cache_resource
121 def load_model():
122 """Load the ASR model with proper error handling and visualization."""
123 try:
124 with st.spinner("Loading Parakeet TDT model... This may take a minute."):
125 model = ASRModel.from_pretrained(model_name=MODEL_NAME)
126 model.eval()
127
128 # Move to appropriate device
129 model = model.to(DEVICE)
130
131 # Use mixed precision for better performance when using GPU
132 if DEVICE == "cuda":
133 model = model.to(torch.bfloat16)
134
135 custom_success(f"Model loaded successfully on {DEVICE.upper()}!")
136 return model
137 except Exception as e:
138 custom_error(f"Error loading model: {str(e)}")
139 if "CUDA" in str(e) or "GPU" in str(e):
140 custom_info("GPU issues detected. Try running with CPU by setting device to 'cpu'.")
141 return None
142
143 def process_audio(audio_path):
144 """Process audio file for transcription with progress reporting."""
145 try:
146 # Load audio file
147 audio = AudioSegment.from_file(audio_path)
148 duration_sec = audio.duration_seconds
149
150 if duration_sec > MAX_RECOMMENDED_DURATION:
151 custom_warning(f"Audio is very long ({duration_sec/60:.1f} minutes). Transcription may take a while and could encounter memory issues.")
152
153 # Progress reporting
154 progress_bar = st.progress(0)
155
156
157 # Resample to 16kHz if needed
158 if audio.frame_rate != 16000:
159 progress_bar.progress(0.2)
160 audio = audio.set_frame_rate(16000)
161
162 progress_bar.progress(0.5)
163
164 # Convert to mono if stereo
165 if audio.channels > 1:
166 audio = audio.set_channels(1)
167
168 progress_bar.progress(0.8)
169
170 # Save processed audio
171 temp_dir = tempfile.gettempdir()
172 processed_path = os.path.join(temp_dir, "processed_audio.wav")
173 audio.export(processed_path, format="wav")
174
175 progress_bar.progress(1.0)
176 custom_success("Audio processed successfully!")
177
178 return processed_path, duration_sec
179 except Exception as e:
180 custom_error(f"Error processing audio: {str(e)}")
181
182 # Provide more specific error messages based on common issues
183 if "No such file" in str(e):
184 custom_info("The audio file could not be found. Please upload it again.")
185 elif "Unsupported format" in str(e) or "unknown format" in str(e):
186 custom_info(f"File format not supported. Please upload one of these formats: {', '.join(SUPPORTED_FORMATS)}")
187 elif "memory" in str(e).lower():
188 custom_info("Memory error occurred. Try with a shorter audio file or restart the application.")
189
190 return None, None
191
192 def format_time(seconds):
193 """Convert seconds to HH:MM:SS format."""
194 return str(datetime.timedelta(seconds=seconds)).split('.')[0]
195
196 def transcribe_audio(audio_path, show_progress=True):
197 """Transcribe audio file using the model with detailed progress reporting."""
198 start_time = time.time()
199 try:
200 model = load_model()
201 if model is None:
202 return None
203
204 processed_path, duration_sec = process_audio(audio_path)
205 if processed_path is None:
206 return None
207
208 # Apply long audio settings if needed
209 long_audio_settings_applied = False
210 if duration_sec > LONG_AUDIO_THRESHOLD:
211 try:
212 custom_info(f"Audio longer than {LONG_AUDIO_THRESHOLD/60:.1f} minutes. Applying optimized settings for long transcription.")
213 model.change_attention_model("rel_pos_local_attn", [256, 256])
214 model.change_subsampling_conv_chunking_factor(1)
215 long_audio_settings_applied = True
216 except Exception as e:
217 custom_warning(f"Could not apply long audio settings: {str(e)}")
218
219 try:
220 if show_progress:
221 with st.spinner("Transcribing audio... This may take a while for longer files."):
222 progress_bar = st.progress(0)
223
224 # Create a progress indicator that updates based on estimated time
225 # This is an estimation since we can't track actual progress
226 estimated_total_time = duration_sec * 0.5 # Rough estimate: processing takes ~50% of audio duration
227
228 def update_progress():
229 for i in range(1, 101):
230 elapsed = time.time() - start_time
231 if elapsed >= estimated_total_time:
232 break
233 progress = min(elapsed / estimated_total_time, 0.95) # Max at 95% until completion
234 progress_bar.progress(progress)
235 time.sleep(estimated_total_time / 100)
236
237 # Start progress update in separate thread
238 import threading
239 progress_thread = threading.Thread(target=update_progress)
240 progress_thread.daemon = True
241 progress_thread.start()
242
243 # Actual transcription
244 output = model.transcribe([processed_path], timestamps=True)
245
246 # Complete progress
247 progress_bar.progress(1.0)
248 else:
249 output = model.transcribe([processed_path], timestamps=True)
250
251 if not output or not isinstance(output, list) or not output[0] or not hasattr(output[0], 'timestamp'):
252 custom_error("Transcription failed or produced unexpected output format.")
253 return None
254
255 segment_timestamps = output[0].timestamp['segment']
256
257 # Generate CSV content with better headers
258 csv_data = [["From (s)", "To (s)", "From (time)", "To (time)", "Duration", "Transcription"]]
259
260 for ts in segment_timestamps:
261 start_s = ts['start']
262 end_s = ts['end']
263 start_formatted = format_time(start_s)
264 end_formatted = format_time(end_s)
265 duration = end_s - start_s
266
267 csv_data.append([
268 f"{start_s:.2f}",
269 f"{end_s:.2f}",
270 start_formatted,
271 end_formatted,
272 f"{duration:.2f}",
273 ts['segment']
274 ])
275
276 processing_time = time.time() - start_time
277 custom_success(f"Transcription completed in {processing_time:.1f} seconds!")
278
279 # Save to history
280 filename = os.path.basename(audio_path)
281 st.session_state.transcription_history.append({
282 'filename': filename,
283 'duration': duration_sec,
284 'timestamp': datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
285 'csv_data': csv_data
286 })
287
288 return csv_data
289
290 finally:
291 # Cleanup
292 if long_audio_settings_applied:
293 try:
294 model.change_attention_model("rel_pos")
295 model.change_subsampling_conv_chunking_factor(-1)
296 except Exception as e:
297 custom_warning(f"Issue reverting model settings: {str(e)}")
298
299 if os.path.exists(processed_path):
300 os.remove(processed_path)
301
302 # Force garbage collection
303 torch.cuda.empty_cache()
304 gc.collect()
305
306 except Exception as e:
307 custom_error(f"Error during transcription: {str(e)}")
308
309 # Offer more helpful advice based on error
310 if "CUDA out of memory" in str(e):
311 custom_info("GPU ran out of memory. Try processing a shorter audio file, or restart the application.")
312 elif "timeout" in str(e).lower():
313 custom_info("The operation timed out. This could be due to the file size or server load.")
314
315 return None
316
317 def export_to_formats(csv_data):
318 """Create exportable data in multiple formats"""
319 # For CSV
320 csv_string = "\n".join([",".join([f'"{cell}"' if ',' in cell else cell for cell in row]) for row in csv_data])
321
322 # For plain text (just the transcription)
323 text_string = "\n\n".join([row[5] for row in csv_data[1:]])
324
325 # For SRT (subtitle format)
326 srt_string = ""
327 for i, row in enumerate(csv_data[1:], 1):
328 start_s = float(row[0])
329 end_s = float(row[1])
330
331 # Convert to SRT time format (HH:MM:SS,mmm)
332 start_srt = f"{int(start_s//3600):02d}:{int((start_s%3600)//60):02d}:{int(start_s%60):02d},{int((start_s%1)*1000):03d}"
333 end_srt = f"{int(end_s//3600):02d}:{int((end_s%3600)//60):02d}:{int(end_s%60):02d},{int((end_s%1)*1000):03d}"
334
335 srt_string += f"{i}\n{start_srt} --> {end_srt}\n{row[5]}\n\n"
336
337 return csv_string, text_string, srt_string
338
339 # Main UI
340 st.title("🎙️ Speech Transcription with Parakeet TDT")
341
342 # Sidebar for navigation
343 with st.sidebar:
344 st.title("Navigation")
345 page = st.radio("Go to", ["Transcribe", "About"])
346
347 # Main content based on selected page
348 if page == "Transcribe":
349 # Description
350 st.markdown("""
351 This demo showcases [parakeet-tdt-0.6b-v2](https://huggingface.co/nvidia/parakeet-tdt-0.6b-v2), a 600-million-parameter model designed for high-quality English speech recognition.
352
353 **Key Features:**
354 - Automatic punctuation and capitalization
355 - Accurate word-level timestamps
356 - Efficiently transcribes long audio segments
357 - Robust performance on spoken numbers and song lyrics transcription
358 """)
359
360 # Create two columns for better layout
361 col1, col2 = st.columns([2, 1])
362
363 with col1:
364 # File uploader with expanded format support
365 uploaded_file = st.file_uploader(f"Upload an audio file (WAV or FLAC)",
366 type=SUPPORTED_FORMATS)
367
368 with col2:
369 if uploaded_file:
370 st.markdown("### Audio Preview")
371 st.audio(uploaded_file)
372
373 # Display basic audio info
374 file_info = f"**File:** {uploaded_file.name}<br>"
375 file_info += f"**Size:** {uploaded_file.size / (1024*1024):.2f} MB<br>"
376 st.markdown(f'<div class="info-box">{file_info}</div>', unsafe_allow_html=True)
377
378 if uploaded_file:
379 # Save uploaded file
380 temp_dir = tempfile.gettempdir()
381 audio_path = os.path.join(temp_dir, uploaded_file.name)
382 with open(audio_path, "wb") as f:
383 f.write(uploaded_file.getbuffer())
384
385 if st.button("🎯 Transcribe Audio", type="primary"):
386 with st.spinner("Processing and transcribing audio..."):
387 csv_data = transcribe_audio(audio_path)
388
389 if csv_data:
390 # Display transcription results
391 st.markdown("### 📝 Transcription Results")
392
393 # Create a dataframe for display
394 df = pd.DataFrame(csv_data[1:], columns=csv_data[0])
395
396 # Display as table with modern styling
397 st.dataframe(
398 df,
399 column_config={
400 "From (s)": st.column_config.NumberColumn(format="%.2f", width="small"),
401 "To (s)": st.column_config.NumberColumn(format="%.2f", width="small"),
402 "From (time)": st.column_config.TextColumn(width="small"),
403 "To (time)": st.column_config.TextColumn(width="small"),
404 "Duration": st.column_config.NumberColumn(format="%.2f", width="small"),
405 "Transcription": st.column_config.TextColumn(width="large")
406 },
407 hide_index=True,
408 use_container_width=True
409 )
410
411 # Create export strings for different formats
412 csv_string, text_string, srt_string = export_to_formats(csv_data)
413
414 # Show export options
415 st.markdown("### 📥 Export Options")
416 col1, col2, col3 = st.columns(3)
417
418 with col1:
419 st.download_button(
420 "📄 Download as CSV",
421 data=csv_string,
422 file_name=f"{os.path.splitext(uploaded_file.name)[0]}_transcript.csv",
423 mime="text/csv",
424 use_container_width=True
425 )
426
427 with col2:
428 st.download_button(
429 "📝 Download as Text",
430 data=text_string,
431 file_name=f"{os.path.splitext(uploaded_file.name)[0]}_transcript.txt",
432 mime="text/plain",
433 use_container_width=True
434 )
435
436 with col3:
437 st.download_button(
438 "🎬 Download as SRT",
439 data=srt_string,
440 file_name=f"{os.path.splitext(uploaded_file.name)[0]}_subtitle.srt",
441 mime="text/plain",
442 use_container_width=True
443 )
444
445 # Word count analysis
446 total_words = sum(len(row[5].split()) for row in csv_data[1:])
447 total_duration = float(csv_data[-1][1]) - float(csv_data[1][0])
448 words_per_minute = (total_words / total_duration) * 60 if total_duration > 0 else 0
449
450 st.markdown("### 📊 Analysis")
451 col1, col2, col3 = st.columns(3)
452 col1.metric("Total Words", f"{total_words}")
453 col2.metric("Speech Duration", f"{format_time(total_duration)}")
454 col3.metric("Words per Minute", f"{words_per_minute:.1f}")
455
456 # Cleanup
457 if os.path.exists(audio_path):
458 os.remove(audio_path)
459
460 elif page == "About":
461 st.title("About this Application")
462
463 st.markdown("""
464 ## Parakeet Speech Recognition
465
466 This application uses NVIDIA's Parakeet-TDT, a powerful speech recognition model designed for accurate transcription with timestamps and punctuation.
467
468 ### Model Details
469
470 - **Model**: [nvidia/parakeet-tdt-0.6b-v2](https://huggingface.co/nvidia/parakeet-tdt-0.6b-v2)
471 - **Parameters**: 600 million
472 - **Features**: Automatic punctuation, capitalization, and word-level timestamps
473 - **Language**: English
474
475
476 ### About Speech Recognition Technology
477
478 Modern speech recognition systems like Parakeet use advanced neural networks trained on thousands of hours of speech data. These systems analyze audio waveforms to predict the most likely sequence of words being spoken, taking into account language patterns and context.
479 """)
480
481 # System information
482 st.subheader("System Information")
483 col1, col2 = st.columns(2)
484
485 with col1:
486 st.markdown(f"**System Platform:** {os.name.upper()}")
487 st.markdown(f"**PyTorch Version:** {torch.__version__}")
488
489 with col2:
490 if torch.cuda.is_available():
491 st.markdown(f"**CUDA Available:** Yes (Version {torch.version.cuda})")
492 st.markdown(f"**GPU:** {torch.cuda.get_device_name(0)}")
493 else:
494 st.markdown("**CUDA Available:** No (Using CPU)")
495
496 # Footer
497 st.markdown("---")
498 st.markdown("Made with NVIDIA's Parakeet-TDT model")