TTS / app.py
Manar312's picture
Update app.py
bcd36a0 verified
Raw
History Blame Contribute Delete
21.8 kB
# -*- coding: utf-8 -*-
"""
RAFIQ - Arabic Speech-to-Text API
Egyptian Arabic ASR using custom character-level Transformer
"""
import os
import io
import math
import gc
import re
import string
import tempfile
from contextlib import asynccontextmanager
from typing import Optional
import numpy as np
import librosa
import noisereduce as nr
import torch
import torch.nn as nn
import torch.nn.functional as F
from num2words import num2words
from fastapi import FastAPI, File, UploadFile, HTTPException
from fastapi.responses import JSONResponse
from pydantic import BaseModel
# ─────────────────────────────────────────────
# Constants (must match training config)
# ─────────────────────────────────────────────
MAX_TEXT_LEN = 70
MAX_SEQ_LEN = 70
N_MELS = 128
SAMPLE_RATE = 16000
HOP_LENGTH = 512
N_FFT = 2048
CHUNK_LENGTH = 15
N_SAMPLES = SAMPLE_RATE * CHUNK_LENGTH
N_FRAMES = math.ceil(N_SAMPLES / HOP_LENGTH)
NEG_INFTY = -1e9
# ─────────────────────────────────────────────
# Vocabulary (identical to training)
# ─────────────────────────────────────────────
special_tokens = ['<PAD>', '<UNK>', '<SOS>', '<EOS>']
english_characters = list(string.ascii_lowercase + ' ')
arabic_characters = list("ابتثجحخدذرزسشصضطظعغفقكلمنهويئءىةؤ")
characters = english_characters + arabic_characters
vocab = special_tokens + characters
char2idx = {char: idx for idx, char in enumerate(vocab)}
idx2char = {idx: char for idx, char in enumerate(vocab)}
vocab_size = len(vocab)
# ─────────────────────────────────────────────
# Model architecture (copy from training script)
# ─────────────────────────────────────────────
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len):
super().__init__()
self.max_len = max_len
self.d_model = d_model
def forward(self):
pos = torch.arange(self.max_len, dtype=torch.float).unsqueeze(1)
_i = torch.arange(self.d_model, dtype=torch.float).unsqueeze(0)
_i = 1 / torch.pow(torch.tensor(10000.0), (2 * (_i // 2)) / self.d_model)
angles = pos * _i
angles[:, 0::2] = torch.sin(angles[:, 0::2])
angles[:, 1::2] = torch.cos(angles[:, 1::2])
return angles[:self.max_len].unsqueeze(0)
def scaled_dot_product_attention(q, k, v, s_mask=None):
d_k = q.shape[-1]
scaled = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k)
if s_mask is not None:
scaled = scaled.permute(1, 0, 2, 3) + s_mask
scaled = scaled.permute(1, 0, 2, 3)
attention = F.softmax(scaled, dim=-1)
return torch.matmul(attention, v), attention
class MultiHeadAttention(nn.Module):
def __init__(self, _input_dim, d_model, _num_heads=8):
super().__init__()
assert d_model % _num_heads == 0
self.num_heads = _num_heads
self.d_model = d_model
self.head_dim = d_model // _num_heads
self.qkv_layer = nn.Linear(_input_dim, 3 * d_model)
self.output_layer = nn.Linear(d_model, d_model)
def forward(self, mha_x, self_mask=None):
seq_len, batch_size = mha_x.shape[1], mha_x.shape[0]
qkv = self.qkv_layer(mha_x)
qkv = qkv.reshape(batch_size, seq_len, self.num_heads, 3 * self.head_dim).permute(0, 2, 1, 3)
q, k, v = qkv.chunk(3, dim=-1)
values, _ = scaled_dot_product_attention(q, k, v, self_mask)
values = values.permute(0, 2, 1, 3).reshape(batch_size, seq_len, self.head_dim * self.num_heads)
return self.output_layer(values)
class MultiHeadCrossAttention(nn.Module):
def __init__(self, d_model, _num_heads=8):
super().__init__()
assert d_model % _num_heads == 0
self.num_heads = _num_heads
self.head_dim = d_model // _num_heads
self.kv_layer = nn.Linear(d_model, 2 * d_model)
self.q_layer = nn.Linear(d_model, d_model)
self.output_layer = nn.Linear(d_model, d_model)
def forward(self, mhca_input, encoder_output, cross_mask=None):
dec_len, enc_len = mhca_input.shape[1], encoder_output.shape[1]
batch_size = mhca_input.shape[0]
kv = self.kv_layer(encoder_output).reshape(batch_size, enc_len, self.num_heads, 2 * self.head_dim).permute(0, 2, 1, 3)
q = self.q_layer(mhca_input).reshape(batch_size, dec_len, self.num_heads, self.head_dim).permute(0, 2, 1, 3)
k, v = kv.chunk(2, dim=-1)
values, _ = scaled_dot_product_attention(q, k, v, cross_mask)
values = values.permute(0, 2, 1, 3).reshape(batch_size, dec_len, self.head_dim * self.num_heads)
return self.output_layer(values)
class LayerNorm(nn.Module):
def __init__(self, normalization_shape, eps=1e-6):
super().__init__()
self.eps = eps
self.gamma = nn.Parameter(torch.ones(normalization_shape))
self.beta = nn.Parameter(torch.zeros(normalization_shape))
def forward(self, layer):
dim = [-(i + 1) for i in range(len(self.gamma.shape))]
mean = layer.mean(dim=dim, keepdim=True)
std = (layer.var(dim=dim, keepdim=True) + self.eps).sqrt()
return self.gamma * (layer - mean) / std + self.beta
class PositionWiseFeedForward(nn.Module):
def __init__(self, d_model, hidden, dropout=0.1):
super().__init__()
self.linear1 = nn.Linear(d_model, hidden)
self.linear2 = nn.Linear(hidden, d_model)
self.gelu = nn.GELU()
self.dropout = nn.Dropout(p=dropout)
def forward(self, x):
return self.linear2(self.dropout(self.gelu(self.linear1(x))))
class EncoderLayer(nn.Module):
def __init__(self, d_model, num_heads, ffn_hidden, dropout):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, d_model, num_heads)
self.layer_norm = LayerNorm([d_model])
self.dropout = nn.Dropout(dropout)
self.ffn = PositionWiseFeedForward(d_model, ffn_hidden, dropout)
def forward(self, x, mask=None):
res = x
x = self.layer_norm(self.dropout(self.self_attn(x, mask)) + res)
res = x
x = self.layer_norm(self.dropout(self.ffn(x)) + res)
return x
class SequentialEncoder(nn.Sequential):
def forward(self, *inputs):
x, mask = inputs
for module in self._modules.values():
x = module(x, mask)
return x
class Encoder(nn.Module):
def __init__(self, d_model, ffn_hidden, num_heads, dropout, num_layers):
super().__init__()
self.layers = SequentialEncoder(*[EncoderLayer(d_model, num_heads, ffn_hidden, dropout) for _ in range(num_layers)])
def forward(self, x, mask):
return self.layers(x, mask)
class DecoderLayer(nn.Module):
def __init__(self, d_model, num_heads, ffn_hidden=2048, dropout=0.1):
super().__init__()
self.self_attn = MultiHeadAttention(d_model, d_model, num_heads)
self.cross_attn = MultiHeadCrossAttention(d_model, num_heads)
self.layer_norm = LayerNorm([d_model])
self.dropout = nn.Dropout(dropout)
self.ffn = PositionWiseFeedForward(d_model, ffn_hidden, dropout)
def forward(self, x, enc_out, self_mask, cross_mask):
res = x; x = self.layer_norm(self.dropout(self.self_attn(x, self_mask)) + res)
res = x; x = self.layer_norm(self.dropout(self.cross_attn(x, enc_out, cross_mask)) + res)
res = x; x = self.layer_norm(self.dropout(self.ffn(x)) + res)
return x
class SequentialDecoder(nn.Sequential):
def forward(self, *inputs):
x, enc_out, self_mask, cross_mask = inputs
for module in self._modules.values():
x = module(x, enc_out, self_mask, cross_mask)
return x
class Decoder(nn.Module):
def __init__(self, d_model, ffn_hidden, num_heads, dropout, num_layers):
super().__init__()
self.layers = SequentialDecoder(*[DecoderLayer(d_model, num_heads, ffn_hidden, dropout) for _ in range(num_layers)])
def forward(self, x, enc_out, self_mask=None, cross_mask=None):
return self.layers(x, enc_out, self_mask, cross_mask)
class Transformer(nn.Module):
def __init__(self, d_model, ffn_hidden, num_heads, drop_prob, num_encoder_layers, num_decoder_layers):
super().__init__()
self.encoder = Encoder(d_model, ffn_hidden, num_heads, drop_prob, num_encoder_layers)
self.decoder = Decoder(d_model, ffn_hidden, num_heads, drop_prob, num_decoder_layers)
def forward(self, src, tgt, enc_self_mask=None, dec_self_mask=None, dec_cross_mask=None):
src = self.encoder(src, enc_self_mask)
return self.decoder(tgt, src, dec_self_mask, dec_cross_mask)
def get_conv_Lout(L_in, conv):
return math.floor((L_in + 2 * conv.padding[0] - conv.dilation[0] * (conv.kernel_size[0] - 1) - 1) / conv.stride[0] + 1)
class MMS(nn.Module):
def __init__(self, vocab_size, d_model=512, nhead=8, num_encoder_layers=6,
num_decoder_layers=6, dim_feedforward=2048,
max_encoder_seq_len=100, max_decoder_seq_len=100,
n_mels=N_MELS, dropout=0.1):
super().__init__()
self.transformer = Transformer(d_model, dim_feedforward, nhead, dropout,
num_encoder_layers, num_decoder_layers)
self.en_positional_encoding = PositionalEncoding(d_model, max_encoder_seq_len)
self.de_positional_encoding = PositionalEncoding(d_model, max_decoder_seq_len)
self.conv1 = nn.Conv1d(n_mels, d_model, kernel_size=3, padding=1)
self.conv2 = nn.Conv1d(d_model, d_model, kernel_size=3, stride=2, padding=1)
self.gelu = nn.GELU()
self.embedding = nn.Embedding(vocab_size, d_model, padding_idx=0)
self.d_model = d_model
self.ff = nn.Linear(d_model, vocab_size)
def get_encoder_seq_len(self, L_in):
return get_conv_Lout(get_conv_Lout(L_in, self.conv1), self.conv2)
def forward(self, audio, text, enc_self_mask, dec_self_mask, dec_cross_mask, device):
audio = self.gelu(self.conv1(audio))
audio = self.gelu(self.conv2(audio))
audio = audio.permute(0, 2, 1)
audio += self.en_positional_encoding().to(device)
text = self.embedding(text) + self.de_positional_encoding().to(device)
out = self.transformer(audio, text, enc_self_mask, dec_self_mask, dec_cross_mask)
return self.ff(out)
# ─────────────────────────────────────────────
# Helpers
# ─────────────────────────────────────────────
def pad_or_trim(array, length=N_SAMPLES, axis=-1, padding=True):
if array.shape[axis] > length:
array = array.take(indices=range(length), axis=axis)
if padding and array.shape[axis] < length:
pad_widths = [(0, 0)] * array.ndim
pad_widths[axis] = (0, length - array.shape[axis])
array = np.pad(array, pad_widths)
return array
def preprocess_audio(audio_bytes: bytes):
audio_data, _ = librosa.load(io.BytesIO(audio_bytes), sr=SAMPLE_RATE)
audio_data = nr.reduce_noise(y=audio_data, sr=SAMPLE_RATE)
duration = librosa.get_duration(y=audio_data, sr=SAMPLE_RATE)
orig_len = audio_data.shape[-1]
modified = pad_or_trim(audio_data, padding=False)
sgram = librosa.stft(y=modified, n_fft=N_FFT, hop_length=HOP_LENGTH)
sgram_mag, _ = librosa.magphase(sgram)
mel = librosa.feature.melspectrogram(S=sgram_mag, sr=SAMPLE_RATE,
n_fft=N_FFT, hop_length=HOP_LENGTH, n_mels=N_MELS)
mel_db = librosa.amplitude_to_db(mel, ref=np.min)
padded = np.pad(mel_db, ((0, 0), (0, N_FRAMES - mel_db.shape[-1])))
return padded, orig_len, duration
def nlp_preprocessing(sentence: str) -> str:
sentence = sentence.lower().replace("\n", " ")
sentence = re.sub(r'[إأآ]', 'ا', sentence)
sentence = re.sub(r'[^a-zA-Zء-ي\s\d]', '', sentence)
sentence = re.sub(r'[\u0617-\u061A\u064B-\u065F]', '', sentence)
sentence = re.sub(r'([a-zA-Z])([ء-ي])|([ء-ي])([a-zA-Z])', r'\1\3 \2\4', sentence)
sentence = re.sub(r'\s+', ' ', sentence)
sentence = re.sub(r'\d+', lambda x: num2words(int(x.group()), lang='ar'), sentence)
return sentence
def tokenize_text(text, max_len=MAX_SEQ_LEN, start_token=True, end_token=True):
tokens = [char2idx.get(c, char2idx['<UNK>']) for c in text]
max_len -= (start_token + end_token)
text_len = len(tokens)
if text_len < max_len:
if end_token:
tokens += [char2idx['<EOS>']]
tokens += [char2idx['<PAD>']] * (max_len - text_len)
else:
tokens = tokens[:max_len]
if end_token:
tokens += [char2idx['<EOS>']]
if start_token:
tokens.insert(0, char2idx['<SOS>'])
return tokens
def text_decoder(token_list):
out = ''
for token in token_list:
if isinstance(token, torch.Tensor):
token = token.item()
char = idx2char[token]
if char == '<EOS>':
return out
if char not in special_tokens:
out += char
return out
def generate_padding_masks(transcription, audio_original_len, conv_func, frames=N_FRAMES):
batch_size, seq_len = transcription.size()
audio_len = conv_func(frames)
look_ahead_mask = torch.triu(torch.full((seq_len, seq_len), True), diagonal=1)
enc_mask = torch.full([batch_size, audio_len, audio_len], False)
dec_self = torch.full([batch_size, seq_len, seq_len], False)
dec_cross = torch.full([batch_size, seq_len, audio_len], False)
for i in range(batch_size):
new_len = conv_func(audio_original_len[i])
enc_mask[i, new_len:, :] = True
enc_mask[i, :, new_len:] = True
dec_cross[i, :, new_len:] = True
zeros = np.where(transcription[0].cpu().numpy() == 0)[0]
if len(zeros) > 0:
idx = zeros[0]
dec_self[i, idx:, :] = True
dec_self[i, :, idx:] = True
dec_cross[i, idx:, :] = True
dec_self_mask = torch.where(look_ahead_mask + dec_self, NEG_INFTY, 0.0)
dec_cross_mask = torch.where(dec_cross, NEG_INFTY, 0.0)
enc_self_mask = torch.where(enc_mask, NEG_INFTY, 0.0)
return enc_self_mask, dec_self_mask, dec_cross_mask
def greedy_decode(audio_tensor, orig_len, model, device, max_len=MAX_TEXT_LEN):
model.eval()
transcription = ""
audio_tensor = audio_tensor.to(device)
with torch.no_grad():
for i in range(max_len):
tgt = torch.tensor(tokenize_text(transcription, max_len=MAX_TEXT_LEN, end_token=False),
dtype=torch.long).unsqueeze(0).to(device)
enc_mask, dec_self, dec_cross = generate_padding_masks(tgt, [orig_len], model.get_encoder_seq_len)
out = model(audio_tensor, tgt,
enc_mask.to(device), dec_self.to(device), dec_cross.to(device), device)
next_tok = torch.argmax(F.softmax(out, dim=-1)[0, i, :], dim=-1).unsqueeze(0)
transcription += text_decoder(next_tok)
if next_tok.item() == char2idx['<EOS>']:
break
return transcription
# ─────────────────────────────────────────────
# App startup / model loading
# ─────────────────────────────────────────────
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model_state = {"model": None}
import gdown
MODEL_PATH = "/app/entire_model.pth"
if not os.path.exists(MODEL_PATH):
print("Downloading model from Google Drive...")
gdown.download(
"https://drive.google.com/uc?id=16i7ui-yYIWppaylOG2ofkpRlPQWeE4nZ",
MODEL_PATH,
quiet=False
)
@asynccontextmanager
async def lifespan(app: FastAPI):
if os.path.exists(MODEL_PATH):
try:
import sys
sys.modules['__main__'].MMS = MMS
sys.modules['__main__'].PositionalEncoding = PositionalEncoding
sys.modules['__main__'].Transformer = Transformer
sys.modules['__main__'].Encoder = Encoder
sys.modules['__main__'].Decoder = Decoder
sys.modules['__main__'].EncoderLayer = EncoderLayer
sys.modules['__main__'].DecoderLayer = DecoderLayer
sys.modules['__main__'].MultiHeadAttention = MultiHeadAttention
sys.modules['__main__'].MultiHeadCrossAttention = MultiHeadCrossAttention
sys.modules['__main__'].PositionWiseFeedForward = PositionWiseFeedForward
sys.modules['__main__'].LayerNorm = LayerNorm
sys.modules['__main__'].SequentialEncoder = SequentialEncoder
sys.modules['__main__'].SequentialDecoder = SequentialDecoder
m = torch.load(MODEL_PATH, map_location=device, weights_only=False)
m.to(device)
m.eval()
model_state["model"] = m
print(f"✅ Model loaded from {MODEL_PATH} on {device}")
except Exception as e:
print(f"⚠️ Could not load model: {e}")
else:
print(f"⚠️ Model file not found at {MODEL_PATH}. Upload your .pth file.")
yield
del model_state["model"]
gc.collect()
# Cleanup
del model_state["model"]
gc.collect()
# ─────────────────────────────────────────────
# FastAPI app
# ─────────────────────────────────────────────
app = FastAPI(
title="RAFIQ – Arabic Speech-to-Text API",
description=(
"**Egyptian Arabic ASR** powered by a custom character-level Transformer.\n\n"
"Upload an audio file (WAV / MP3 / OGG, ≤ 15 s) and receive the Arabic transcription.\n\n"
"Built for the **RAFIQ** autism-support platform – FCAI, Beni-Suef University 2026."
),
version="1.0.0",
lifespan=lifespan,
)
# ── Response schemas ──────────────────────────
class TranscribeResponse(BaseModel):
transcription: str
duration_seconds: float
device: str
model_loaded: bool
class HealthResponse(BaseModel):
status: str
model_loaded: bool
device: str
vocab_size: int
sample_rate: int
max_audio_seconds: int
# ── Endpoints ────────────────────────────────
@app.get("/", tags=["Info"])
def root():
return {
"message": "RAFIQ Arabic ASR API is running 🎙️",
"docs": "/docs",
"health": "/health",
}
@app.get("/health", response_model=HealthResponse, tags=["Info"])
def health():
return HealthResponse(
status="ok",
model_loaded=model_state["model"] is not None,
device=str(device),
vocab_size=vocab_size,
sample_rate=SAMPLE_RATE,
max_audio_seconds=CHUNK_LENGTH,
)
@app.post(
"/transcribe",
response_model=TranscribeResponse,
tags=["ASR"],
summary="Transcribe Egyptian Arabic audio",
description=(
"Upload a WAV / MP3 / OGG audio file (max 15 seconds).\n"
"The API applies noise reduction, extracts mel-spectrograms, "
"and runs greedy decoding through the Arabic Transformer model."
),
)
async def transcribe(
file: UploadFile = File(..., description="Audio file: WAV, MP3, or OGG. Max 15 seconds.")
):
if model_state["model"] is None:
raise HTTPException(
status_code=503,
detail="Model not loaded. Make sure 'entire_model.pth' is present in the Space.",
)
allowed = {"audio/wav", "audio/x-wav", "audio/mpeg", "audio/mp3", "audio/ogg", "audio/flac"}
if file.content_type and file.content_type not in allowed:
raise HTTPException(
status_code=415,
detail=f"Unsupported file type: {file.content_type}. Use WAV, MP3, or OGG.",
)
audio_bytes = await file.read()
if len(audio_bytes) > 50 * 1024 * 1024: # 50 MB guard
raise HTTPException(status_code=413, detail="File too large (max 50 MB).")
try:
mel, orig_len, duration = preprocess_audio(audio_bytes)
except Exception as e:
raise HTTPException(status_code=422, detail=f"Audio preprocessing failed: {e}")
audio_tensor = torch.tensor(mel, dtype=torch.float32).unsqueeze(0) # (1, N_MELS, N_FRAMES)
try:
transcription = greedy_decode(audio_tensor, orig_len, model_state["model"], device)
except Exception as e:
raise HTTPException(status_code=500, detail=f"Inference failed: {e}")
return TranscribeResponse(
transcription=transcription,
duration_seconds=round(duration, 2),
device=str(device),
model_loaded=True,
)