Spaces:
Paused
Paused
Download app.py from naicoi/GLM-ASR-Nano: direct link, hf CLI and curl.
- Browser
- Download file 5.39 kB
-
https://huggingface.co/spaces/naicoi/GLM-ASR-Nano/resolve/main/app.py
- Command line
-
hf download hf://spaces/naicoi/GLM-ASR-Nano/app.py
-
curl -L -o app.py https://huggingface.co/spaces/naicoi/GLM-ASR-Nano/resolve/main/app.py
5.39 kB
| import os | |
| import tempfile | |
| import time | |
| import wave | |
| import gradio as gr | |
| import spaces | |
| import torch | |
| from transformers import AutoProcessor, GlmAsrForConditionalGeneration | |
| CHECKPOINT_DIR = "zai-org/GLM-ASR-Nano-2512" | |
| MAX_NEW_TOKENS = 128 | |
| ENABLE_AOT = os.environ.get("ENABLE_AOT", "0") == "1" | |
| # --------------------------------------------------------------------------- | |
| # Model Loading (CPU — ZeroGPU moves to CUDA per @spaces.GPU call) | |
| # --------------------------------------------------------------------------- | |
| print(f"Loading model from {CHECKPOINT_DIR}...") | |
| LOAD_ERROR = None | |
| try: | |
| processor = AutoProcessor.from_pretrained(CHECKPOINT_DIR) | |
| model = GlmAsrForConditionalGeneration.from_pretrained( | |
| CHECKPOINT_DIR, | |
| torch_dtype=torch.bfloat16, | |
| ) | |
| model.eval() | |
| except Exception as e: | |
| LOAD_ERROR = f"Failed to load model: {e}" | |
| print(LOAD_ERROR) | |
| # --------------------------------------------------------------------------- | |
| # AoT Compilation (opt-in — saves 1.3–1.8× inference time on H200) | |
| # --------------------------------------------------------------------------- | |
| AOT_COMPILED = False | |
| if LOAD_ERROR is None and ENABLE_AOT and hasattr(spaces, "aoti_capture"): | |
| def _create_dummy_audio(duration_sec=2, sample_rate=16000): | |
| tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) | |
| with wave.open(tmp.name, "w") as wf: | |
| wf.setnchannels(1) | |
| wf.setsampwidth(2) | |
| wf.setframerate(sample_rate) | |
| wf.writeframes(b"\x00\x00" * (duration_sec * sample_rate)) | |
| return tmp.name | |
| def _compile_model(): | |
| global AOT_COMPILED | |
| dummy_path = _create_dummy_audio() | |
| try: | |
| inputs = processor.apply_transcription_request(dummy_path) | |
| inputs = inputs.to(model.device, dtype=model.dtype) | |
| # 1. Capture representative inputs | |
| with spaces.aoti_capture(model) as call: | |
| model(**inputs) | |
| # 2. Export computation graph | |
| exported = torch.export.export(model, args=call.args, kwargs=call.kwargs) | |
| # 3. Compile for H200 | |
| compiled = spaces.aoti_compile(exported) | |
| # 4. Patch model.forward with compiled version | |
| spaces.aoti_apply(compiled, model) | |
| AOT_COMPILED = True | |
| print("✓ AoT compilation successful") | |
| except Exception as e: | |
| print(f"✗ AoT compilation failed (falling back to eager): {e}") | |
| finally: | |
| os.unlink(dummy_path) | |
| _compile_model() | |
| # --------------------------------------------------------------------------- | |
| # Transcription (with timing) | |
| # --------------------------------------------------------------------------- | |
| def transcribe_wrapper(audio_file_path): | |
| if LOAD_ERROR is not None: | |
| raise gr.Error(LOAD_ERROR) | |
| if model is None or processor is None: | |
| raise gr.Error("Model failed to load. Check Space logs for details.") | |
| if audio_file_path is None: | |
| return "[Please upload an audio file or record one.]", "" | |
| t_start = time.perf_counter() | |
| try: | |
| inputs = processor.apply_transcription_request(audio_file_path) | |
| t_process = time.perf_counter() | |
| inputs = inputs.to(model.device, dtype=model.dtype) | |
| t_device = time.perf_counter() | |
| with torch.inference_mode(): | |
| outputs = model.generate( | |
| **inputs, do_sample=False, max_new_tokens=MAX_NEW_TOKENS | |
| ) | |
| t_generate = time.perf_counter() | |
| transcript = processor.batch_decode( | |
| outputs[:, inputs.input_ids.shape[1] :], | |
| skip_special_tokens=True, | |
| )[0].strip() | |
| t_end = time.perf_counter() | |
| mode = "AoT" if AOT_COMPILED else "eager" | |
| timing = ( | |
| f"Process: {t_process - t_start:.3f}s | " | |
| f"ToDevice: {t_device - t_process:.3f}s | " | |
| f"Generate: {t_generate - t_device:.3f}s | " | |
| f"Decode: {t_end - t_generate:.3f}s | " | |
| f"Total: {t_end - t_start:.3f}s [{mode}]" | |
| ) | |
| print(timing) | |
| return transcript or "[Empty transcription]", timing | |
| except Exception as e: | |
| elapsed = time.perf_counter() - t_start | |
| print(f"Transcription error ({elapsed:.3f}s): {e}") | |
| return f"An error occurred during transcription: {e}", "" | |
| # --------------------------------------------------------------------------- | |
| # Gradio Interface | |
| # --------------------------------------------------------------------------- | |
| title = "✨ GLM-ASR-Nano-2512 Transcription Demo" | |
| description = ( | |
| "This demo uses the sota new GLM-ASR Nano model to transcribe audio files with great accuracy! " | |
| "The architecture is simple and efficient, composed of a whisper encoder and an llm. " | |
| "Upload an audio file (or record one) to transcribe it into text using the model." | |
| ) | |
| audio_input = gr.Audio( | |
| type="filepath", | |
| label="Audio Input (WAV/MP3)", | |
| sources=["upload", "microphone"], | |
| ) | |
| output_text = gr.Textbox(label="Transcription Result", lines=5) | |
| timing_text = gr.Textbox(label="Performance", lines=2) | |
| demo = gr.Interface( | |
| fn=transcribe_wrapper, | |
| inputs=[audio_input], | |
| outputs=[output_text, timing_text], | |
| title=title, | |
| description=description, | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch() | |