# -*- 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 = ['', '', '', ''] 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['']) for c in text] max_len -= (start_token + end_token) text_len = len(tokens) if text_len < max_len: if end_token: tokens += [char2idx['']] tokens += [char2idx['']] * (max_len - text_len) else: tokens = tokens[:max_len] if end_token: tokens += [char2idx['']] if start_token: tokens.insert(0, char2idx['']) 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 == '': 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['']: 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, )