nahw-api / main.py
Yxxi4's picture
Update main.py
d4dc133 verified
Raw
History Blame Contribute Delete
8.69 kB
import math
import torch
import torch.nn as nn
from fastapi import FastAPI
from pydantic import BaseModel
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, pipeline
# ==========================================
# 1. الإعدادات والثوابت (Constants)
# ==========================================
DEVICE = torch.device('cpu') # إجبار العمل على المعالج العادي للسيرفر المجاني
FATHA = '\u064E'
DAMMA = '\u064F'
KASRA = '\u0650'
SUKUN = '\u0652'
SHADDA = '\u0651'
FATHATAN = '\u064B'
DAMMATAN = '\u064C'
KASRATAN = '\u064D'
ALL_DIACRITICS = set([FATHA, DAMMA, KASRA, SUKUN, SHADDA, FATHATAN, DAMMATAN, KASRATAN])
DIACRITIC_CLASSES = [
'', FATHA, DAMMA, KASRA, SUKUN, SHADDA,
SHADDA + FATHA, SHADDA + DAMMA, SHADDA + KASRA,
FATHATAN, DAMMATAN, KASRATAN,
SHADDA + FATHATAN, SHADDA + DAMMATAN, SHADDA + KASRATAN,
]
IDX_TO_DIAC = {i: d for i, d in enumerate(DIACRITIC_CLASSES)}
NUM_CLASSES = 15
def remove_diacritics(text):
return ''.join(c for c in text if c not in ALL_DIACRITICS)
# ==========================================
# 2. بنية نموذج التشكيل (Architecture)
# ==========================================
class MultiHeadSelfAttention(nn.Module):
def __init__(self, hidden_dim, num_heads=8, dropout=0.1):
super().__init__()
assert hidden_dim % num_heads == 0
self.num_heads = num_heads
self.head_dim = hidden_dim // num_heads
self.q_proj = nn.Linear(hidden_dim, hidden_dim)
self.k_proj = nn.Linear(hidden_dim, hidden_dim)
self.v_proj = nn.Linear(hidden_dim, hidden_dim)
self.out_proj = nn.Linear(hidden_dim, hidden_dim)
self.dropout = nn.Dropout(dropout)
self.scale = math.sqrt(self.head_dim)
def forward(self, x, mask=None):
B, T, C = x.shape
Q = self.q_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
K = self.k_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
V = self.v_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
attn = torch.matmul(Q, K.transpose(-2, -1)) / self.scale
if mask is not None:
attn_mask = mask.unsqueeze(1).unsqueeze(2)
attn = attn.masked_fill(attn_mask == 0, float('-inf'))
attn = torch.softmax(attn, dim=-1)
attn = self.dropout(attn)
out = torch.matmul(attn, V)
out = out.transpose(1, 2).contiguous().view(B, T, C)
return self.out_proj(out)
class AttentionBlock(nn.Module):
def __init__(self, hidden_dim, num_heads=8, dropout=0.1):
super().__init__()
self.norm1 = nn.LayerNorm(hidden_dim)
self.attn = MultiHeadSelfAttention(hidden_dim, num_heads, dropout)
self.norm2 = nn.LayerNorm(hidden_dim)
self.ff = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim * 2),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_dim * 2, hidden_dim),
nn.Dropout(dropout)
)
def forward(self, x, mask=None):
x = x + self.attn(self.norm1(x), mask)
x = x + self.ff(self.norm2(x))
return x
class BiLSTMAttentionDiacritizer(nn.Module):
def __init__(self, vocab_size, embed_dim=256, hidden_dim=256,
num_lstm_layers=3, num_attn_layers=2, num_heads=8,
num_classes=15, dropout=0.3):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
self.embed_dropout = nn.Dropout(dropout)
self.lstm = nn.LSTM(
embed_dim, hidden_dim, num_lstm_layers,
batch_first=True, bidirectional=True,
dropout=dropout if num_lstm_layers > 1 else 0
)
lstm_out_dim = hidden_dim * 2
self.attn_layers = nn.ModuleList([
AttentionBlock(lstm_out_dim, num_heads, dropout)
for _ in range(num_attn_layers)
])
self.final_norm = nn.LayerNorm(lstm_out_dim)
self.dropout = nn.Dropout(dropout)
self.classifier = nn.Sequential(
nn.Linear(lstm_out_dim, hidden_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_dim, num_classes)
)
def forward(self, input_ids, attention_mask=None):
x = self.embed_dropout(self.embedding(input_ids))
lstm_out, _ = self.lstm(x)
for attn_layer in self.attn_layers:
lstm_out = attn_layer(lstm_out, attention_mask)
lstm_out = self.dropout(self.final_norm(lstm_out))
return self.classifier(lstm_out)
# ==========================================
# 3. تحميل النماذج (Model Loading)
# ==========================================
print("جاري تحميل نموذج التشكيل...")
# تأكد من أن اسم الملف هنا يطابق الملف المرفوع على Hugging Face
CHECKPOINT_PATH = 'Tashkeel_model.pt'
checkpoint = torch.load(CHECKPOINT_PATH, map_location=DEVICE)
char_to_idx = checkpoint['char_to_idx']
tashkeel_model = BiLSTMAttentionDiacritizer(
vocab_size=len(char_to_idx),
embed_dim=256,
hidden_dim=256,
num_lstm_layers=3,
num_attn_layers=2,
num_heads=8,
num_classes=NUM_CLASSES,
dropout=0.3
).to(DEVICE)
tashkeel_model.load_state_dict(checkpoint['model_state_dict'])
tashkeel_model.eval()
print("جاري تحميل نموذج AraBART...")
# في حال كان مسار النموذج على Hugging Face مختلفاً، يرجى تعديل السطر التالي
ARABART_MODEL_NAME = "CAMeL-Lab/arabart-qalb15-gec-ged-13"
arabart_tokenizer = AutoTokenizer.from_pretrained(ARABART_MODEL_NAME)
arabart_model = AutoModelForSeq2SeqLM.from_pretrained(ARABART_MODEL_NAME)
gec_pipeline = pipeline("text2text-generation", model=arabart_model, tokenizer=arabart_tokenizer)
print("✅ تمت تهيئة جميع النماذج بنجاح!")
# ==========================================
# 4. دالة التشكيل المساعدة
# ==========================================
def diacritize_text(text, model, char_to_idx, max_len=200):
clean = remove_diacritics(text)
char_ids = [char_to_idx.get(c, 1) for c in clean]
all_preds = [0] * len(char_ids)
counts = [0] * len(char_ids)
stride = max_len - 40
for start in range(0, max(len(char_ids), 1), stride):
chunk = char_ids[start:start+max_len]
actual = len(chunk)
padded = chunk + [0]*(max_len - actual)
mask = [1]*actual + [0]*(max_len - actual)
input_t = torch.LongTensor([padded]).to(DEVICE)
mask_t = torch.LongTensor([mask]).to(DEVICE)
with torch.no_grad():
logits = model(input_t, mask_t)
preds = logits[0, :actual].argmax(dim=-1).cpu().tolist()
for i, p in enumerate(preds):
pos = start + i
if pos < len(all_preds):
if counts[pos] == 0 or p != 0:
all_preds[pos] = p
counts[pos] += 1
if start + max_len >= len(char_ids): break
result = []
for char, pidx in zip(clean, all_preds):
result.append(char)
d = IDX_TO_DIAC.get(pidx, '')
if d: result.append(d)
return ''.join(result)
# ==========================================
# 5. إعداد الخادم ومسارات الـ API
# ==========================================
from fastapi.middleware.cors import CORSMiddleware
app = FastAPI()
# إضافة صلاحيات CORS للسماح لموقعك بالاتصال بالـ API
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # يسمح لجميع النطاقات (مثل Vercel) بالاتصال
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
class TextRequest(BaseModel):
text: str
@app.get("/")
def read_root():
return {
"status": "online",
"message": "مرحباً بك في واجهة برمجة تطبيقات نظام نحو (NAHW) للمعالجة الذكية للنصوص",
"endpoints": ["/tashkeel", "/spell-check", "/grammar-check"]
}
@app.post("/tashkeel")
def process_tashkeel(request: TextRequest):
result = diacritize_text(request.text, tashkeel_model, char_to_idx)
return {"result_text": result}
@app.post("/spell-check")
def spell_check(request: TextRequest):
result = gec_pipeline(request.text)
return {"result_text": result[0]['generated_text']}
@app.post("/grammar-check")
def grammar_check(request: TextRequest):
result = gec_pipeline(request.text)
return {"result_text": result[0]['generated_text']}