GLM-ASR-Nano / app.py
naicoi's picture
Update app.py
15f813a verified
Raw History Blame Contribute Delete
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
@spaces.GPU(duration=15)
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)
# ---------------------------------------------------------------------------
@spaces.GPU(duration=5)
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()