audiotim1 / app.py
Atikarahmanda's picture
Update app.py
a8e5691 verified
Raw
History Blame Contribute Delete
11.5 kB
"""
Whisper Transcription API β€” HuggingFace Spaces (FastAPI)
Model: Atikarahmanda/whispertim1 (Fine-tuned Whisper Bahasa Indonesia)
Strategi load model:
1. Coba pipeline() langsung ke fine-tuned repo
2. Jika gagal (repo tidak punya preprocessor_config.json),
load WhisperProcessor dari base model + fine-tuned weights secara manual
"""
import io
import os
import time
import tempfile
import warnings
import traceback
from pathlib import Path
from contextlib import asynccontextmanager
import torch
import numpy as np
import soundfile as sf
import librosa
from fastapi import FastAPI, UploadFile, File, HTTPException
from fastapi.responses import JSONResponse
from transformers import (
pipeline,
WhisperFeatureExtractor,
WhisperForConditionalGeneration,
WhisperTokenizer,
)
from transformers.utils import logging as hf_logging
warnings.filterwarnings("ignore")
hf_logging.set_verbosity_error()
# ─────────────────────────────────────────────
# KONFIGURASI
# ─────────────────────────────────────────────
FINETUNED_MODEL_ID = "Atikarahmanda/whisperteam1"
CHUNK_SEC = 30
STRIDE_SEC = 5
GLOBAL_BATCH_SIZE = 8
SAMPLE_RATE = 16_000
LANGUAGE = "indonesian"
TASK = "transcribe"
SUPPORTED_EXTS = {".wav", ".mp3", ".m4a", ".flac", ".ogg", ".webm", ".opus"}
# ─────────────────────────────────────────────
# GLOBAL STATE (model di-load sekali saat startup)
# ─────────────────────────────────────────────
state: dict = {}
def _build_pipeline(device_id: int):
"""
Load pipeline ASR dengan bypass total AutoFeatureExtractor.
Masalah: transformers >= 4.40 selalu mencari 'preprocessor_config.json'
melalui AutoFeatureExtractor / WhisperProcessor, bahkan jika repo punya
'processor_config.json' (nama lama yang valid).
Solusi: load WhisperFeatureExtractor dan WhisperTokenizer LANGSUNG
tanpa melalui Auto class β€” keduanya tidak butuh preprocessor_config.json.
"""
generate_kwargs = {
"language": LANGUAGE,
"task": TASK,
"max_new_tokens": 440,
"no_repeat_ngram_size": 3,
"repetition_penalty": 1.3,
"temperature": 0.2,
}
# WhisperFeatureExtractor.from_pretrained membaca processor_config.json
# secara langsung tanpa melewati AutoFeatureExtractor.
print(f" Memuat WhisperFeatureExtractor dari '{FINETUNED_MODEL_ID}'...")
feature_extractor = WhisperFeatureExtractor.from_pretrained(FINETUNED_MODEL_ID)
# Tokenizer Whisper tidak berubah saat fine-tuning β€” aman load dari base model.
# Repo whispertim1 tidak punya vocab.json (dibutuhkan slow tokenizer),
# jadi kita load dari openai/whisper-large-v2 yang pasti ada.
# Catatan: base model harus sesuai ukuran yang dipakai fine-tuning.
BASE_TOKENIZER = "openai/whisper-small"
print(f" Memuat tokenizer dari '{BASE_TOKENIZER}'...")
tokenizer = WhisperTokenizer.from_pretrained(
BASE_TOKENIZER,
language=LANGUAGE,
task=TASK,
)
print(f" Memuat model weights dari '{FINETUNED_MODEL_ID}'...")
model = WhisperForConditionalGeneration.from_pretrained(FINETUNED_MODEL_ID)
asr = pipeline(
"automatic-speech-recognition",
model=model,
tokenizer=tokenizer,
feature_extractor=feature_extractor,
chunk_length_s=CHUNK_SEC,
stride_length_s=STRIDE_SEC,
batch_size=GLOBAL_BATCH_SIZE,
device=device_id,
return_timestamps=True,
generate_kwargs=generate_kwargs,
)
print(" Pipeline berhasil dibangun.")
return asr
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Load model saat startup, unload saat shutdown."""
print("πŸ”„ Memuat model Whisper...")
t0 = time.time()
device = "cuda" if torch.cuda.is_available() else "cpu"
device_id = 0 if torch.cuda.is_available() else -1
print(f" Device: {device.upper()}")
state["transcriber"] = _build_pipeline(device_id)
print(f"βœ… Whisper siap β€” {time.time() - t0:.1f} detik")
print("πŸ”„ Memuat model Silero VAD...")
t1 = time.time()
vad_model, vad_utils = torch.hub.load(
repo_or_dir="snakers4/silero-vad",
model="silero_vad",
force_reload=False,
trust_repo=True,
)
(
state["get_speech_timestamps"],
_save_audio,
state["read_audio"],
_VADIterator,
_collect_chunks,
) = vad_utils
state["vad_model"] = vad_model
print(f"βœ… VAD siap β€” {time.time() - t1:.1f} detik")
yield # app berjalan di sini
state.clear()
print("πŸ›‘ Model di-unload.")
app = FastAPI(
title="Whisper Transcription API",
description="ASR Bahasa Indonesia β€” model Atikarahmanda/whispertim1",
version="1.0.0",
lifespan=lifespan,
)
# ─────────────────────────────────────────────
# HELPER
# ─────────────────────────────────────────────
def load_audio_bytes(data: bytes, filename: str) -> np.ndarray:
"""
Baca bytes audio β†’ numpy array mono 16 kHz.
Mendukung wav, mp3, m4a, flac, ogg, webm, opus via soundfile/librosa.
"""
ext = Path(filename).suffix.lower()
with tempfile.NamedTemporaryFile(suffix=ext, delete=False) as tmp:
tmp.write(data)
tmp_path = tmp.name
try:
# librosa handles virtually every format (calls ffmpeg/soundfile internally)
wav, sr = librosa.load(tmp_path, sr=SAMPLE_RATE, mono=True)
finally:
os.unlink(tmp_path)
return wav.astype(np.float32)
def apply_vad(wav: np.ndarray) -> np.ndarray:
"""
Jalankan Silero VAD, potong keheningan di awal.
Kembalikan numpy array yang sudah di-trim.
"""
# VAD butuh tensor
wav_tensor = torch.from_numpy(wav)
timestamps = state["get_speech_timestamps"](
wav_tensor,
state["vad_model"],
sampling_rate=SAMPLE_RATE,
)
if timestamps:
start_sample = timestamps[0]["start"]
cut_point = max(0, start_sample - int(0.3 * SAMPLE_RATE)) # 300ms buffer
wav_trimmed = wav[cut_point:]
vad_status = f"speech ditemukan, dipotong {cut_point / SAMPLE_RATE:.2f} detik dari awal"
else:
wav_trimmed = wav
vad_status = "tidak ada speech terdeteksi, audio utuh"
return wav_trimmed, vad_status
# ─────────────────────────────────────────────
# ENDPOINTS
# ─────────────────────────────────────────────
@app.get("/", tags=["health"])
def root():
return {"status": "ok", "model": FINETUNED_MODEL_ID}
@app.get("/health", tags=["health"])
def health():
ready = "transcriber" in state and "vad_model" in state
return {"ready": ready, "device": "cuda" if torch.cuda.is_available() else "cpu"}
@app.post("/transcribe", tags=["transcription"])
async def transcribe(file: UploadFile = File(...)):
"""
Kirim satu file audio β†’ terima teks transkripsi.
- **file**: File audio (wav, mp3, m4a, flac, ogg, webm, opus)
Response:
```json
{
"filename": "audio.wav",
"transkripsi": "...",
"vad_status": "speech ditemukan, dipotong 0.80 detik dari awal",
"durasi_detik": 12.3,
"waktu_proses_detik": 4.1
}
```
"""
if "transcriber" not in state:
raise HTTPException(503, "Model belum siap, tunggu sebentar lagi.")
ext = Path(file.filename or "audio").suffix.lower()
if ext not in SUPPORTED_EXTS:
raise HTTPException(
400,
f"Format tidak didukung: '{ext}'. "
f"Gunakan salah satu: {sorted(SUPPORTED_EXTS)}",
)
t_start = time.time()
try:
raw = await file.read()
wav = load_audio_bytes(raw, file.filename or "audio.wav")
except Exception as e:
raise HTTPException(422, f"Gagal membaca audio: {e}")
durasi = len(wav) / SAMPLE_RATE
try:
wav_vad, vad_status = apply_vad(wav)
except Exception as e:
# VAD gagal β†’ pakai audio asli
wav_vad = wav
vad_status = f"VAD error ({e}), audio utuh"
try:
result = state["transcriber"](wav_vad)
if isinstance(result, list):
result = result[0]
teks = result["text"].strip()
except Exception as e:
raise HTTPException(500, f"Transkripsi gagal: {e}\n{traceback.format_exc()}")
waktu_proses = round(time.time() - t_start, 2)
return JSONResponse({
"filename": file.filename,
"transkripsi": teks,
"vad_status": vad_status,
"durasi_detik": round(durasi, 2),
"waktu_proses_detik": waktu_proses,
})
@app.post("/transcribe/batch", tags=["transcription"])
async def transcribe_batch(files: list[UploadFile] = File(...)):
"""
Kirim banyak file audio sekaligus β†’ terima list transkripsi.
Maksimum 20 file per request.
"""
if "transcriber" not in state:
raise HTTPException(503, "Model belum siap.")
if len(files) > 20:
raise HTTPException(400, "Maksimum 20 file per request.")
t_start = time.time()
audio_inputs = []
meta = []
for f in files:
ext = Path(f.filename or "audio").suffix.lower()
try:
raw = await f.read()
wav = load_audio_bytes(raw, f.filename or "audio.wav")
wav_vad, vad_status = apply_vad(wav)
audio_inputs.append(wav_vad)
meta.append({
"filename": f.filename,
"durasi_detik": round(len(wav) / SAMPLE_RATE, 2),
"vad_status": vad_status,
})
except Exception as e:
meta.append({
"filename": f.filename,
"durasi_detik": 0,
"vad_status": f"error: {e}",
"transkripsi": f"ERROR: {e}",
})
audio_inputs.append(None)
results = []
try:
valid_idx = [i for i, a in enumerate(audio_inputs) if a is not None]
valid_audio = [audio_inputs[i] for i in valid_idx]
if valid_audio:
batch_out = state["transcriber"](valid_audio)
if isinstance(batch_out, dict):
batch_out = [batch_out]
out_iter = iter(batch_out)
for i, m in enumerate(meta):
if i in valid_idx:
res = next(out_iter)
results.append({**m, "transkripsi": res["text"].strip()})
else:
results.append(m) # already has "transkripsi" error key
except Exception as e:
raise HTTPException(500, f"Batch transcription gagal: {e}")
return JSONResponse({
"total_file": len(files),
"waktu_proses_detik": round(time.time() - t_start, 2),
"hasil": results,
})