Yxxi4 commited on
Commit
b74a4fd
·
verified ·
1 Parent(s): edd4911

Create main.py

Browse files
Files changed (1) hide show
  1. main.py +195 -0
main.py ADDED
@@ -0,0 +1,195 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ import math
4
+ import re
5
+ from fastapi import FastAPI
6
+ from pydantic import BaseModel
7
+
8
+ # ==========================================
9
+ # 1. الثوابت (Constants)
10
+ # ==========================================
11
+ FATHA = '\u064E'
12
+ DAMMA = '\u064F'
13
+ KASRA = '\u0650'
14
+ SUKUN = '\u0652'
15
+ SHADDA = '\u0651'
16
+ FATHATAN = '\u064B'
17
+ DAMMATAN = '\u064C'
18
+ KASRATAN = '\u064D'
19
+
20
+ ALL_DIACRITICS = set([FATHA, DAMMA, KASRA, SUKUN, SHADDA, FATHATAN, DAMMATAN, KASRATAN])
21
+ DIACRITIC_CLASSES = [
22
+ '', FATHA, DAMMA, KASRA, SUKUN, SHADDA,
23
+ SHADDA + FATHA, SHADDA + DAMMA, SHADDA + KASRA,
24
+ FATHATAN, DAMMATAN, KASRATAN,
25
+ SHADDA + FATHATAN, SHADDA + DAMMATAN, SHADDA + KASRATAN,
26
+ ]
27
+ IDX_TO_DIAC = {i: d for i, d in enumerate(DIACRITIC_CLASSES)}
28
+ NUM_CLASSES = 15
29
+ DEVICE = torch.device('cpu') # إجبار العمل على المعالج العادي للسيرفر المجاني
30
+
31
+ def remove_diacritics(text):
32
+ return ''.join(c for c in text if c not in ALL_DIACRITICS)
33
+
34
+ # ==========================================
35
+ # 2. بنية النموذج (Architecture Classes)
36
+ # ==========================================
37
+ class MultiHeadSelfAttention(nn.Module):
38
+ def __init__(self, hidden_dim, num_heads=8, dropout=0.1):
39
+ super().__init__()
40
+ assert hidden_dim % num_heads == 0
41
+ self.num_heads = num_heads
42
+ self.head_dim = hidden_dim // num_heads
43
+ self.q_proj = nn.Linear(hidden_dim, hidden_dim)
44
+ self.k_proj = nn.Linear(hidden_dim, hidden_dim)
45
+ self.v_proj = nn.Linear(hidden_dim, hidden_dim)
46
+ self.out_proj = nn.Linear(hidden_dim, hidden_dim)
47
+ self.dropout = nn.Dropout(dropout)
48
+ self.scale = math.sqrt(self.head_dim)
49
+
50
+ def forward(self, x, mask=None):
51
+ B, T, C = x.shape
52
+ Q = self.q_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
53
+ K = self.k_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
54
+ V = self.v_proj(x).view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
55
+ attn = torch.matmul(Q, K.transpose(-2, -1)) / self.scale
56
+ if mask is not None:
57
+ attn_mask = mask.unsqueeze(1).unsqueeze(2)
58
+ attn = attn.masked_fill(attn_mask == 0, float('-inf'))
59
+ attn = torch.softmax(attn, dim=-1)
60
+ attn = self.dropout(attn)
61
+ out = torch.matmul(attn, V)
62
+ out = out.transpose(1, 2).contiguous().view(B, T, C)
63
+ return self.out_proj(out)
64
+
65
+ class AttentionBlock(nn.Module):
66
+ def __init__(self, hidden_dim, num_heads=8, dropout=0.1):
67
+ super().__init__()
68
+ self.norm1 = nn.LayerNorm(hidden_dim)
69
+ self.attn = MultiHeadSelfAttention(hidden_dim, num_heads, dropout)
70
+ self.norm2 = nn.LayerNorm(hidden_dim)
71
+ self.ff = nn.Sequential(
72
+ nn.Linear(hidden_dim, hidden_dim * 2),
73
+ nn.GELU(),
74
+ nn.Dropout(dropout),
75
+ nn.Linear(hidden_dim * 2, hidden_dim),
76
+ nn.Dropout(dropout)
77
+ )
78
+
79
+ def forward(self, x, mask=None):
80
+ x = x + self.attn(self.norm1(x), mask)
81
+ x = x + self.ff(self.norm2(x))
82
+ return x
83
+
84
+ class BiLSTMAttentionDiacritizer(nn.Module):
85
+ def __init__(self, vocab_size, embed_dim=256, hidden_dim=256,
86
+ num_lstm_layers=3, num_attn_layers=2, num_heads=8,
87
+ num_classes=15, dropout=0.3):
88
+ super().__init__()
89
+ self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
90
+ self.embed_dropout = nn.Dropout(dropout)
91
+ self.lstm = nn.LSTM(
92
+ embed_dim, hidden_dim, num_lstm_layers,
93
+ batch_first=True, bidirectional=True,
94
+ dropout=dropout if num_lstm_layers > 1 else 0
95
+ )
96
+ lstm_out_dim = hidden_dim * 2
97
+ self.attn_layers = nn.ModuleList([
98
+ AttentionBlock(lstm_out_dim, num_heads, dropout)
99
+ for _ in range(num_attn_layers)
100
+ ])
101
+ self.final_norm = nn.LayerNorm(lstm_out_dim)
102
+ self.dropout = nn.Dropout(dropout)
103
+ self.classifier = nn.Sequential(
104
+ nn.Linear(lstm_out_dim, hidden_dim),
105
+ nn.GELU(),
106
+ nn.Dropout(dropout),
107
+ nn.Linear(hidden_dim, num_classes)
108
+ )
109
+
110
+ def forward(self, input_ids, attention_mask=None):
111
+ x = self.embed_dropout(self.embedding(input_ids))
112
+ lstm_out, _ = self.lstm(x)
113
+ for attn_layer in self.attn_layers:
114
+ lstm_out = attn_layer(lstm_out, attention_mask)
115
+ lstm_out = self.dropout(self.final_norm(lstm_out))
116
+ return self.classifier(lstm_out)
117
+
118
+ # ==========================================
119
+ # 3. تحميل النموذج والأوزان
120
+ # ==========================================
121
+ # تحميل الملف المرجعي وتوجيهه للـ CPU
122
+ checkpoint = torch.load('Tashkeel_model.pt', map_location=DEVICE)
123
+
124
+ # استخراج القاموس لمعرفة حجم المفردات
125
+ char_to_idx = checkpoint['char_to_idx']
126
+ VOCAB_SIZE = len(char_to_idx)
127
+
128
+ # تهيئة النموذج بنفس الإعدادات التي تدرب عليها
129
+ tashkeel_model = BiLSTMAttentionDiacritizer(
130
+ vocab_size=VOCAB_SIZE,
131
+ embed_dim=256,
132
+ hidden_dim=256,
133
+ num_lstm_layers=3,
134
+ num_attn_layers=2,
135
+ num_heads=8,
136
+ num_classes=NUM_CLASSES,
137
+ dropout=0.3
138
+ ).to(DEVICE)
139
+
140
+ # تركيب الأوزان
141
+ tashkeel_model.load_state_dict(checkpoint['model_state_dict'])
142
+ tashkeel_model.eval() # تفعيل وضع الاستخدام
143
+
144
+ # ==========================================
145
+ # 4. دوال التشكيل والـ API
146
+ # ==========================================
147
+ def diacritize_text(text, model, char_to_idx, max_len=200):
148
+ clean = remove_diacritics(text)
149
+ char_ids = [char_to_idx.get(c, 1) for c in clean]
150
+ all_preds = [0] * len(char_ids)
151
+ counts = [0] * len(char_ids)
152
+ stride = max_len - 40
153
+
154
+ for start in range(0, max(len(char_ids), 1), stride):
155
+ chunk = char_ids[start:start+max_len]
156
+ actual = len(chunk)
157
+ padded = chunk + [0]*(max_len - actual)
158
+ mask = [1]*actual + [0]*(max_len - actual)
159
+
160
+ input_t = torch.LongTensor([padded]).to(DEVICE)
161
+ mask_t = torch.LongTensor([mask]).to(DEVICE)
162
+
163
+ with torch.no_grad():
164
+ logits = model(input_t, mask_t)
165
+ preds = logits[0, :actual].argmax(dim=-1).cpu().tolist()
166
+
167
+ for i, p in enumerate(preds):
168
+ pos = start + i
169
+ if pos < len(all_preds):
170
+ if counts[pos] == 0 or p != 0:
171
+ all_preds[pos] = p
172
+ counts[pos] += 1
173
+
174
+ if start + max_len >= len(char_ids): break
175
+
176
+ result = []
177
+ for char, pidx in zip(clean, all_preds):
178
+ result.append(char)
179
+ d = IDX_TO_DIAC.get(pidx, '')
180
+ if d: result.append(d)
181
+ return ''.join(result)
182
+
183
+ app = FastAPI()
184
+
185
+ class TextRequest(BaseModel):
186
+ text: str
187
+
188
+ @app.post("/tashkeel")
189
+ def process_tashkeel(request: TextRequest):
190
+ result = diacritize_text(request.text, tashkeel_model, char_to_idx)
191
+ return {"diacritized_text": result}
192
+
193
+ @app.get("/")
194
+ def read_root():
195
+ return {"message": "API نحو للتشكيل الآلي يعمل بنجاح!"}