abinet2018 commited on
Commit
734d25b
ยท
verified ยท
1 Parent(s): f5d1fcb

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +64 -62
app.py CHANGED
@@ -1,35 +1,36 @@
1
  # app.py
2
- # Amharic Poetry Joint Theme Classifier - Multi-Model Comparison
 
3
  # Compatible with Gradio 6.20.0
4
 
5
  import gradio as gr
6
  import matplotlib.pyplot as plt
7
  import torch
8
  import torch.nn as nn
9
- import torch.nn.functional as F
10
  import os
11
  import json
12
  import numpy as np
13
- from transformers import AutoTokenizer, AutoModel, RobertaConfig
14
 
15
  print("=" * 80)
16
  print("๐Ÿš€ Amharic Poetry Joint Theme Classifier")
17
  print("=" * 80)
18
  print("๐Ÿ“Š Model: Rasyosef-RoBERTa (Multi-Task)")
 
19
  print("=" * 80)
20
 
21
  # ============================================
22
  # CLASS NAMES (English and Amharic)
23
- # MUST MATCH THE TRAINING ORDER
24
  # ============================================
25
 
26
  CLASS_NAMES_EN = [
27
- 'Religious', # LABEL_0
28
- 'Ethical', # LABEL_1
29
- 'Political', # LABEL_2
30
- 'Philosophical', # LABEL_3
31
- 'Historical', # LABEL_4
32
- 'Love' # LABEL_5
33
  ]
34
 
35
  CLASS_NAMES_AM = [
@@ -47,29 +48,35 @@ LABEL_TO_ID = {name: i for i, name in enumerate(CLASS_NAMES_EN)}
47
 
48
  # ============================================
49
  # MULTI-TASK MODEL ARCHITECTURE
50
- # MATCHES THE TRAINED MODEL
51
  # ============================================
52
 
53
  class MultiTaskAmharicPoetryModel(nn.Module):
54
- def __init__(self, model_name, num_labels, dropout=0.3, projection_dim=128):
 
 
 
55
  super().__init__()
56
  self.num_labels = num_labels
57
  self.projection_dim = projection_dim
58
 
59
- # Shared encoder
60
  self.encoder = AutoModel.from_pretrained(model_name)
61
  hidden_size = self.encoder.config.hidden_size
62
 
63
- # Dropout
64
  self.dropout = nn.Dropout(dropout)
65
 
66
  # Multi-label head (Sigmoid, 6 outputs)
 
67
  self.multi_label_head = nn.Linear(hidden_size, num_labels)
68
 
69
  # Main-class head (Softmax, 6 classes)
 
70
  self.main_class_head = nn.Linear(hidden_size, num_labels)
71
 
72
  # Contrastive learning projection head
 
73
  self.projection_head = nn.Sequential(
74
  nn.Linear(hidden_size, hidden_size),
75
  nn.ReLU(),
@@ -83,13 +90,13 @@ class MultiTaskAmharicPoetryModel(nn.Module):
83
  input_ids=input_ids,
84
  attention_mask=attention_mask
85
  )
86
- # Use CLS token representation
87
  pooled_output = outputs.last_hidden_state[:, 0, :]
88
  pooled_output = self.dropout(pooled_output)
89
 
90
  # Task-specific outputs
91
- multi_logits = self.multi_label_head(pooled_output)
92
- main_logits = self.main_class_head(pooled_output)
93
 
94
  if return_projection:
95
  projection = self.projection_head(pooled_output)
@@ -112,28 +119,25 @@ def load_model():
112
  if model is not None:
113
  return model, tokenizer, device
114
 
115
- device = torch.device(
116
- "cuda" if torch.cuda.is_available() else "cpu"
117
- )
118
-
119
  print(f"๐Ÿ“ฑ Device: {device}")
120
 
121
  if not os.path.exists("best_model.bin"):
122
  print("โŒ best_model.bin not found!")
123
- print(os.listdir("."))
124
  return None, None, device
125
 
126
  try:
127
- print("๐Ÿ“ฅ Loading tokenizer from current directory...")
128
  tokenizer = AutoTokenizer.from_pretrained(".")
129
-
130
  # Ensure pad token exists
131
  if tokenizer.pad_token is None:
132
  tokenizer.pad_token = tokenizer.eos_token
133
 
134
  print("๐Ÿ—๏ธ Building Multi-Task Model architecture...")
135
 
136
- # Build the multi-task model
137
  model = MultiTaskAmharicPoetryModel(
138
  "rasyosef/roberta-base-amharic",
139
  num_labels=6,
@@ -143,33 +147,26 @@ def load_model():
143
 
144
  print("๐Ÿ“ฆ Loading best_model.bin weights...")
145
 
146
- state_dict = torch.load(
147
- "best_model.bin",
148
- map_location=device
149
- )
150
 
151
- # Load state dict
152
- missing, unexpected = model.load_state_dict(
153
- state_dict,
154
- strict=False
155
- )
156
 
157
  if missing:
158
- print(f"โš ๏ธ Missing keys: {len(missing)} keys")
159
  if len(missing) <= 10:
160
- print("Missing keys:", missing)
161
  if unexpected:
162
- print(f"โš ๏ธ Unexpected keys: {len(unexpected)} keys")
163
  if len(unexpected) <= 10:
164
- print("Unexpected keys:", unexpected)
165
 
166
  model.to(device)
167
  model.eval()
168
 
169
- print("โœ… best_model.bin loaded successfully!")
170
- print(
171
- f"๐Ÿ“Š Parameters: {sum(p.numel() for p in model.parameters()):,}"
172
- )
173
 
174
  return model, tokenizer, device
175
 
@@ -188,15 +185,14 @@ try:
188
  print("โŒ Model loading failed!")
189
  except Exception as e:
190
  print(f"โŒ Model loading failed: {e}")
191
- import traceback
192
- traceback.print_exc()
193
 
194
  # ============================================
195
  # PREDICTION FUNCTION
196
  # ============================================
197
 
198
  def gradio_predict(text, threshold=0.5):
199
- """Gradio prediction function."""
 
200
 
201
  if not text or not text.strip():
202
  return "โš ๏ธ Please enter a poem.", None
@@ -205,7 +201,7 @@ def gradio_predict(text, threshold=0.5):
205
  return "โŒ Model not loaded. Please check the logs.", None
206
 
207
  try:
208
- # Tokenize with max_length=128 (from config)
209
  inputs = tokenizer(
210
  text,
211
  return_tensors='pt',
@@ -219,17 +215,22 @@ def gradio_predict(text, threshold=0.5):
219
  # Get predictions
220
  with torch.no_grad():
221
  multi_logits, main_logits = model(input_ids, attention_mask)
222
- probs = torch.sigmoid(multi_logits) # Multi-label uses sigmoid
223
- probs_np = probs.cpu().numpy()[0]
 
 
 
 
 
224
 
225
- # Main class (highest probability from multi-label)
226
- main_idx = np.argmax(probs_np)
227
  main_class = ID_TO_LABEL[main_idx]
228
- main_confidence = float(probs_np[main_idx])
229
 
230
- # Multi-label (all classes above threshold, excluding main)
231
  multi_labels = []
232
- for i, prob in enumerate(probs_np):
233
  if prob >= threshold and i != main_idx:
234
  multi_labels.append({
235
  'label': ID_TO_LABEL[i],
@@ -239,10 +240,10 @@ def gradio_predict(text, threshold=0.5):
239
 
240
  # All probabilities
241
  all_probs = {}
242
- for i, prob in enumerate(probs_np):
243
  all_probs[ID_TO_LABEL[i]] = float(prob)
244
 
245
- # Build output
246
  output = f"## ๐ŸŽฏ Predicted Theme Class of the Poem\n(แ‹จแŒแŒฅแˆ™ แ‹‹แŠ“ แ‹จแŒญแ‰ฅแŒฅ แˆแ‹ตแ‰ฅ)\n\n"
247
 
248
  # Display as "Love poetry (แ‹จแแ‰…แˆญ แŒแŒฅแˆ)"
@@ -258,7 +259,7 @@ def gradio_predict(text, threshold=0.5):
258
  else:
259
  output += "*No additional themes above threshold.*\n"
260
 
261
- # Create plot
262
  labels = list(all_probs.keys())
263
  probs = list(all_probs.values())
264
 
@@ -274,8 +275,9 @@ def gradio_predict(text, threshold=0.5):
274
  ax.set_ylim(0, 1)
275
  ax.set_ylabel("Probability", fontsize=12, fontweight='bold')
276
  ax.set_title("Class Probabilities", fontsize=14, fontweight='bold')
277
- ax.axhline(y=threshold, linestyle="--", alpha=0.7, color='#e74c3c', linewidth=2)
278
  ax.grid(True, alpha=0.2, axis='y')
 
279
 
280
  for bar, p in zip(bars, probs):
281
  if p > 0.01:
@@ -292,15 +294,15 @@ def gradio_predict(text, threshold=0.5):
292
  return f"โŒ Error: {str(e)}", None
293
 
294
  # ============================================
295
- # SAMPLE POEMS
296
  # ============================================
297
 
298
  sample_poems = [
299
- # Ethical poem
300
  """แˆฐแ‹ แŠจแŒŽแˆจแ‰คแ‰ฑ แŠ แ‰ฅแˆฎ แˆˆแˆ˜แŠ–แˆญแฃ
301
  แ‰ แˆแ‰ก แ‹ญแŠ‘แˆจแ‹ แ‹ฐแˆตแ‰ณแŠ“ แแ‰…แˆญแกแก""",
302
 
303
- # Historical poem
304
  """แˆ˜แ‰…แ‹ฐแˆ‹ แŠ แ‹แ‰ แˆ‹แ‹ญ แŒฉแŠ‹แ‰ต แ‰ แˆจแŠจแ‰ฐ
305
  แ‹จแˆดแ‰ฑแŠ• แŠ แŠ“แ‹‰แ‰…แˆ แ‹ˆแŠ•แ‹ต แŠ แŠ•แ‹ต แˆฐแ‹‰ แˆžแ‰ฐ
306
  แ‹จแˆฐแˆœแŠ‘แŠ• แŠ•แŒ‰แˆต แˆฒแŠ•แ‰ แˆฒแŠ•แ‰
@@ -336,7 +338,7 @@ sample_poems = [
336
  ]
337
 
338
  # ============================================
339
- # GRADIO UI
340
  # ============================================
341
 
342
  # Custom CSS
@@ -420,7 +422,7 @@ with gr.Blocks(
420
  outputs=[text_input, threshold_slider, output_text, output_plot]
421
  )
422
 
423
- # Footer
424
  gr.Markdown("""
425
  ---
426
  <div style="text-align: center; font-size: 13px; color: #5d6d7e;">
 
1
  # app.py
2
+ # Amharic Poetry Joint Theme Classifier
3
+ # Based on: Transformer-Based Joint Thematic Classification for Amharic Poetry
4
  # Compatible with Gradio 6.20.0
5
 
6
  import gradio as gr
7
  import matplotlib.pyplot as plt
8
  import torch
9
  import torch.nn as nn
 
10
  import os
11
  import json
12
  import numpy as np
13
+ from transformers import AutoTokenizer, AutoModel
14
 
15
  print("=" * 80)
16
  print("๐Ÿš€ Amharic Poetry Joint Theme Classifier")
17
  print("=" * 80)
18
  print("๐Ÿ“Š Model: Rasyosef-RoBERTa (Multi-Task)")
19
+ print("๐Ÿ“„ Paper: Transformer-Based Joint Thematic Classification for Amharic Poetry")
20
  print("=" * 80)
21
 
22
  # ============================================
23
  # CLASS NAMES (English and Amharic)
24
+ # MATCHES THE PAPER'S 6 THEMATIC CATEGORIES
25
  # ============================================
26
 
27
  CLASS_NAMES_EN = [
28
+ 'Religious', # LABEL_0 - แˆƒแ‹ญแˆ›แŠ–แ‰ณแ‹Š แŒแŒฅแˆ
29
+ 'Ethical', # LABEL_1 - แˆฅแА-แˆแŒแ‰ฃแˆซแ‹Š แŒแŒฅแˆ
30
+ 'Political', # LABEL_2 - แ–แˆˆแ‰ฒแŠซแ‹Š แŒแŒฅแˆ
31
+ 'Philosophical', # LABEL_3 - แแˆแˆตแแŠ“แ‹Š แŒแŒฅแˆ
32
+ 'Historical', # LABEL_4 - แ‰ณแˆชแŠซแ‹Š แŒแŒฅแˆ
33
+ 'Love' # LABEL_5 - แ‹จแแ‰…แˆญ แŒแŒฅแˆ
34
  ]
35
 
36
  CLASS_NAMES_AM = [
 
48
 
49
  # ============================================
50
  # MULTI-TASK MODEL ARCHITECTURE
51
+ # MATCHES THE PAPER'S JOINT LEARNING FRAMEWORK
52
  # ============================================
53
 
54
  class MultiTaskAmharicPoetryModel(nn.Module):
55
+ """Multi-task model for joint dominant-theme and multi-label classification.
56
+ This architecture matches the paper's joint learning framework."""
57
+
58
+ def __init__(self, model_name, num_labels=6, dropout=0.3, projection_dim=128):
59
  super().__init__()
60
  self.num_labels = num_labels
61
  self.projection_dim = projection_dim
62
 
63
+ # Shared encoder (Rasyosef-RoBERTa)
64
  self.encoder = AutoModel.from_pretrained(model_name)
65
  hidden_size = self.encoder.config.hidden_size
66
 
67
+ # Dropout for regularization
68
  self.dropout = nn.Dropout(dropout)
69
 
70
  # Multi-label head (Sigmoid, 6 outputs)
71
+ # Used for multi-label thematic classification
72
  self.multi_label_head = nn.Linear(hidden_size, num_labels)
73
 
74
  # Main-class head (Softmax, 6 classes)
75
+ # Used for dominant-theme prediction
76
  self.main_class_head = nn.Linear(hidden_size, num_labels)
77
 
78
  # Contrastive learning projection head
79
+ # Used for discriminative feature learning
80
  self.projection_head = nn.Sequential(
81
  nn.Linear(hidden_size, hidden_size),
82
  nn.ReLU(),
 
90
  input_ids=input_ids,
91
  attention_mask=attention_mask
92
  )
93
+ # Use CLS token representation for document-level semantics
94
  pooled_output = outputs.last_hidden_state[:, 0, :]
95
  pooled_output = self.dropout(pooled_output)
96
 
97
  # Task-specific outputs
98
+ multi_logits = self.multi_label_head(pooled_output) # Sigmoid for multi-label
99
+ main_logits = self.main_class_head(pooled_output) # Softmax for dominant theme
100
 
101
  if return_projection:
102
  projection = self.projection_head(pooled_output)
 
119
  if model is not None:
120
  return model, tokenizer, device
121
 
122
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
 
 
 
123
  print(f"๐Ÿ“ฑ Device: {device}")
124
 
125
  if not os.path.exists("best_model.bin"):
126
  print("โŒ best_model.bin not found!")
127
+ print("๐Ÿ“ Files:", os.listdir("."))
128
  return None, None, device
129
 
130
  try:
131
+ print("๐Ÿ“ฅ Loading tokenizer...")
132
  tokenizer = AutoTokenizer.from_pretrained(".")
133
+
134
  # Ensure pad token exists
135
  if tokenizer.pad_token is None:
136
  tokenizer.pad_token = tokenizer.eos_token
137
 
138
  print("๐Ÿ—๏ธ Building Multi-Task Model architecture...")
139
 
140
+ # Build the multi-task model with Rasyosef-RoBERTa
141
  model = MultiTaskAmharicPoetryModel(
142
  "rasyosef/roberta-base-amharic",
143
  num_labels=6,
 
147
 
148
  print("๐Ÿ“ฆ Loading best_model.bin weights...")
149
 
150
+ state_dict = torch.load("best_model.bin", map_location=device)
 
 
 
151
 
152
+ # Load state dict with strict=False to handle mismatches
153
+ missing, unexpected = model.load_state_dict(state_dict, strict=False)
 
 
 
154
 
155
  if missing:
156
+ print(f"โš ๏ธ Missing keys: {len(missing)}")
157
  if len(missing) <= 10:
158
+ print("Missing:", missing)
159
  if unexpected:
160
+ print(f"โš ๏ธ Unexpected keys: {len(unexpected)}")
161
  if len(unexpected) <= 10:
162
+ print("Unexpected:", unexpected)
163
 
164
  model.to(device)
165
  model.eval()
166
 
167
+ print(f"โœ… Model loaded successfully!")
168
+ print(f"๐Ÿ“Š Parameters: {sum(p.numel() for p in model.parameters()):,}")
169
+ print(f"๐Ÿ“Š Hidden size: {model.encoder.config.hidden_size}")
 
170
 
171
  return model, tokenizer, device
172
 
 
185
  print("โŒ Model loading failed!")
186
  except Exception as e:
187
  print(f"โŒ Model loading failed: {e}")
 
 
188
 
189
  # ============================================
190
  # PREDICTION FUNCTION
191
  # ============================================
192
 
193
  def gradio_predict(text, threshold=0.5):
194
+ """Gradio prediction function.
195
+ Uses sigmoid for multi-label and softmax for dominant theme."""
196
 
197
  if not text or not text.strip():
198
  return "โš ๏ธ Please enter a poem.", None
 
201
  return "โŒ Model not loaded. Please check the logs.", None
202
 
203
  try:
204
+ # Tokenize with max_length=128 (from paper's config)
205
  inputs = tokenizer(
206
  text,
207
  return_tensors='pt',
 
215
  # Get predictions
216
  with torch.no_grad():
217
  multi_logits, main_logits = model(input_ids, attention_mask)
218
+ # Multi-label uses sigmoid (as in paper)
219
+ multi_probs = torch.sigmoid(multi_logits)
220
+ # Dominant theme uses softmax (as in paper)
221
+ main_probs = torch.softmax(main_logits, dim=1)
222
+
223
+ multi_probs_np = multi_probs.cpu().numpy()[0]
224
+ main_probs_np = main_probs.cpu().numpy()[0]
225
 
226
+ # Dominant theme: highest probability from multi-label
227
+ main_idx = np.argmax(multi_probs_np)
228
  main_class = ID_TO_LABEL[main_idx]
229
+ main_confidence = float(multi_probs_np[main_idx])
230
 
231
+ # Multi-label: all classes above threshold, excluding main
232
  multi_labels = []
233
+ for i, prob in enumerate(multi_probs_np):
234
  if prob >= threshold and i != main_idx:
235
  multi_labels.append({
236
  'label': ID_TO_LABEL[i],
 
240
 
241
  # All probabilities
242
  all_probs = {}
243
+ for i, prob in enumerate(multi_probs_np):
244
  all_probs[ID_TO_LABEL[i]] = float(prob)
245
 
246
+ # Build output (matching UI images)
247
  output = f"## ๐ŸŽฏ Predicted Theme Class of the Poem\n(แ‹จแŒแŒฅแˆ™ แ‹‹แŠ“ แ‹จแŒญแ‰ฅแŒฅ แˆแ‹ตแ‰ฅ)\n\n"
248
 
249
  # Display as "Love poetry (แ‹จแแ‰…แˆญ แŒแŒฅแˆ)"
 
259
  else:
260
  output += "*No additional themes above threshold.*\n"
261
 
262
+ # Create plot (matching UI images)
263
  labels = list(all_probs.keys())
264
  probs = list(all_probs.values())
265
 
 
275
  ax.set_ylim(0, 1)
276
  ax.set_ylabel("Probability", fontsize=12, fontweight='bold')
277
  ax.set_title("Class Probabilities", fontsize=14, fontweight='bold')
278
+ ax.axhline(y=threshold, linestyle="--", alpha=0.7, color='#e74c3c', linewidth=2, label=f'Threshold ({threshold:.0%})')
279
  ax.grid(True, alpha=0.2, axis='y')
280
+ ax.legend(loc='upper right')
281
 
282
  for bar, p in zip(bars, probs):
283
  if p > 0.01:
 
294
  return f"โŒ Error: {str(e)}", None
295
 
296
  # ============================================
297
+ # SAMPLE POEMS (From UI Images and Paper)
298
  # ============================================
299
 
300
  sample_poems = [
301
+ # Ethical poem (from UI image - Figure 9)
302
  """แˆฐแ‹ แŠจแŒŽแˆจแ‰คแ‰ฑ แŠ แ‰ฅแˆฎ แˆˆแˆ˜แŠ–แˆญแฃ
303
  แ‰ แˆแ‰ก แ‹ญแŠ‘แˆจแ‹ แ‹ฐแˆตแ‰ณแŠ“ แแ‰…แˆญแกแก""",
304
 
305
+ # Historical poem (from UI image - Figure 10)
306
  """แˆ˜แ‰…แ‹ฐแˆ‹ แŠ แ‹แ‰ แˆ‹แ‹ญ แŒฉแŠ‹แ‰ต แ‰ แˆจแŠจแ‰ฐ
307
  แ‹จแˆดแ‰ฑแŠ• แŠ แŠ“แ‹‰แ‰…แˆ แ‹ˆแŠ•แ‹ต แŠ แŠ•แ‹ต แˆฐแ‹‰ แˆžแ‰ฐ
308
  แ‹จแˆฐแˆœแŠ‘แŠ• แŠ•แŒ‰แˆต แˆฒแŠ•แ‰ แˆฒแŠ•แ‰
 
338
  ]
339
 
340
  # ============================================
341
+ # GRADIO UI (Matches UI Images)
342
  # ============================================
343
 
344
  # Custom CSS
 
422
  outputs=[text_input, threshold_slider, output_text, output_plot]
423
  )
424
 
425
+ # Footer with paper citation
426
  gr.Markdown("""
427
  ---
428
  <div style="text-align: center; font-size: 13px; color: #5d6d7e;">