PauloVBM commited on
Commit
f57e388
·
verified ·
1 Parent(s): 9786513

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +186 -40
app.py CHANGED
@@ -1,6 +1,12 @@
1
  import torch
2
- from transformers import AutoTokenizer, AutoModelForSequenceClassification
3
  import gradio as gr
 
 
 
 
 
 
4
 
5
  REPO_ID = "engsoftexperimental/modelos"
6
  MODEL_NAME = "modelo_final_distilbert2"
@@ -11,50 +17,190 @@ model = AutoModelForSequenceClassification.from_pretrained(REPO_ID, subfolder=MO
11
  model.config.id2label = {0: "hate_speech", 1: "offensive", 2: "neither", 3: "spam"}
12
  model.config.label2id = {"hate_speech": 0, "offensive": 1, "neither": 2, "spam": 3}
13
 
14
- def classify_text(texto):
15
- inputs = tokenizer(
16
- texto,
17
- return_tensors="pt",
18
- truncation=True,
19
- padding="max_length",
20
- max_length=128
21
- ).to(model.device)
22
 
 
 
 
 
 
23
  with torch.no_grad():
24
  logits = model(**inputs).logits
25
- probs = torch.softmax(logits, dim=1)[0].cpu().numpy()
26
- pred = int(torch.argmax(logits, dim=1))
27
-
28
- return {
29
- "classe_prevista": model.config.id2label[pred],
30
- "prob_hate_speech": float(probs[0]),
31
- "prob_offensive": float(probs[1]),
32
- "prob_neither": float(probs[2]),
33
- "prob_spam": float(probs[3])
34
- }
35
-
36
- def gradio_predict(text):
37
- result = classify_text(text)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
38
  return (
39
- result["classe_prevista"],
40
- f"{result['prob_hate_speech']:.4f}",
41
- f"{result['prob_offensive']:.4f}",
42
- f"{result['prob_neither']:.4f}",
43
- f"{result['prob_spam']:.4f}"
 
 
 
44
  )
45
 
46
- demo = gr.Interface(
47
- fn=gradio_predict,
48
- inputs=gr.Textbox(lines=5, label="Digite ou cole um texto"),
49
- outputs=[
50
- gr.Label(label="Classificação Prevista"),
51
- gr.Textbox(label="Probabilidade (Hate Speech)"),
52
- gr.Textbox(label="Probabilidade (Offensive)"),
53
- gr.Textbox(label="Probabilidade (Neither)"),
54
- gr.Textbox(label="Probabilidade (Spam)")
55
- ],
56
- title="Protótipo de Moderador de Conteúdo (HateBR + BERT)",
57
- description="Ferramenta experimental utilizada na pesquisa para classificação de conteúdo em três classes: Hate Speech, Offensive, Neither e Spam."
58
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
59
 
60
  demo.launch()
 
1
  import torch
2
+ import numpy as np
3
  import gradio as gr
4
+ from transformers import AutoTokenizer, AutoModelForSequenceClassification, pipeline
5
+ from deep_translator import GoogleTranslator
6
+ import lime
7
+ from lime.lime_text import LimeTextExplainer
8
+ import shap
9
+ import matplotlib.pyplot as plt
10
 
11
  REPO_ID = "engsoftexperimental/modelos"
12
  MODEL_NAME = "modelo_final_distilbert2"
 
17
  model.config.id2label = {0: "hate_speech", 1: "offensive", 2: "neither", 3: "spam"}
18
  model.config.label2id = {"hate_speech": 0, "offensive": 1, "neither": 2, "spam": 3}
19
 
20
+ lime_explainer = LimeTextExplainer(class_names=labels)
 
 
 
 
 
 
 
21
 
22
+ pipe = pipeline("text-classification", model=model, tokenizer=tokenizer, top_k=None)
23
+ shap_explainer = shap.Explainer(pipe)
24
+
25
+ def predictor_for_lime(texts):
26
+ inputs = tokenizer(texts, return_tensors="pt", truncation=True, padding=True, max_length=128).to(model.device)
27
  with torch.no_grad():
28
  logits = model(**inputs).logits
29
+ probs = torch.softmax(logits, dim=1).cpu().numpy()
30
+ return probs
31
+
32
+ def translate_text(text, source_lang_name, languages_dict):
33
+ if not text: return ""
34
+ if source_lang_name == "Detectar Automaticamente":
35
+ src = 'auto'
36
+ else:
37
+ src = languages_dict.get(source_lang_name, 'auto')
38
+ try:
39
+ translator = GoogleTranslator(source=src, target='en')
40
+ return translator.translate(text)
41
+ except Exception as e:
42
+ return f"Erro: {str(e)}"
43
+
44
+ def generate_lime_plot(text_en, pred_idx):
45
+ try:
46
+ exp = lime_explainer.explain_instance(
47
+ text_en,
48
+ predictor_for_lime,
49
+ num_features=6,
50
+ num_samples=100,
51
+ labels=[pred_idx]
52
+ )
53
+ fig = exp.as_pyplot_figure(label=pred_idx)
54
+ plt.tight_layout()
55
+ return fig
56
+ except Exception as e:
57
+ print(f"Erro LIME: {e}")
58
+ return None
59
+
60
+ def generate_shap_plot(text_en, pred_idx):
61
+ try:
62
+ shap_values = shap_explainer([text_en])
63
+
64
+ vals = shap_values.values
65
+
66
+ if len(vals.shape) == 3:
67
+ class_values = vals[0, :, pred_idx]
68
+ elif len(vals.shape) == 2:
69
+ class_values = vals[0, :]
70
+ else:
71
+ raise ValueError(f"Shape inesperado do SHAP: {vals.shape}")
72
+
73
+ raw_data = shap_values.data[0]
74
+
75
+ if hasattr(raw_data, 'tolist'):
76
+ tokens = raw_data.tolist()
77
+ else:
78
+ tokens = list(raw_data)
79
+
80
+ tokens = [str(t).strip() for t in tokens]
81
+
82
+ if len(tokens) != len(class_values):
83
+ min_len = min(len(tokens), len(class_values))
84
+ tokens = tokens[:min_len]
85
+ class_values = class_values[:min_len]
86
+
87
+ explanation = shap.Explanation(
88
+ values=class_values,
89
+ feature_names=tokens
90
+ )
91
+
92
+ plt.close('all')
93
+ plt.figure()
94
+
95
+ shap.plots.bar(explanation, max_display=10, show=False)
96
+
97
+ plt.tight_layout()
98
+ return plt.gcf()
99
+
100
+ except Exception as e:
101
+ print(f"Erro CRÍTICO no SHAP: {e}")
102
+ try:
103
+ print(f"Debug Shape: {shap_values.values.shape}")
104
+ print(f"Debug Tokens: {shap_values.data[0]}")
105
+ except:
106
+ pass
107
+
108
+ plt.close('all')
109
+ fig, ax = plt.subplots(figsize=(6, 2))
110
+ ax.text(0.5, 0.5, "Não foi possível gerar a explicação SHAP\npara este texto específico.",
111
+ ha='center', va='center', fontsize=10, color='red')
112
+ ax.axis('off')
113
+ return fig
114
+
115
+ def process_pipeline(source_lang, text):
116
+ try:
117
+ LANGUAGES_DICT = GoogleTranslator().get_supported_languages(as_dict=True)
118
+ except:
119
+ LANGUAGES_DICT = {'portuguese': 'pt', 'english': 'en'}
120
+
121
+ text_en = translate_text(text, source_lang, LANGUAGES_DICT)
122
+
123
+ probs = predictor_for_lime([text_en])[0]
124
+ pred_idx = int(np.argmax(probs))
125
+ pred_label = labels[pred_idx]
126
+
127
+ lime_fig = generate_lime_plot(text_en, pred_idx)
128
+ shap_fig = generate_shap_plot(text_en, pred_idx)
129
+
130
  return (
131
+ text_en,
132
+ pred_label,
133
+ f"{probs[0]:.4f}",
134
+ f"{probs[1]:.4f}",
135
+ f"{probs[3]:.4f}",
136
+ f"{probs[2]:.4f}",
137
+ lime_fig,
138
+ shap_fig
139
  )
140
 
141
+ try:
142
+ LANGUAGES_DICT = GoogleTranslator().get_supported_languages(as_dict=True)
143
+ except:
144
+ LANGUAGES_DICT = {'portuguese': 'pt', 'english': 'en'}
145
+ language_options = ["Detectar Automaticamente"] + list(LANGUAGES_DICT.keys())
146
+
147
+ with gr.Blocks(title="Moderador AI + XAI") as demo:
148
+
149
+ gr.Markdown("# Moderador de Conteúdo com IA Explicável (Modelo DistilBERT com Hate Speech Dataset + LIME/SHAP)")
150
+ gr.Markdown("Ferramenta experimental utilizada na pesquisa para classificação de conteúdo em três classes: Hate Speech, Offensive, Spam e Neither.<br>O modelo foi treinado com dados no idioma inglês. Por isso, textos em outros idiomas serão traduzidos antes da classificação para garantir a precisão da IA.")
151
+ gr.Markdown("Digite um texto. O sistema realiza a tradução automática, classifica e gera gráficos de explicabilidade.")
152
+
153
+ with gr.Row():
154
+ with gr.Column(scale=1):
155
+ dd_lang = gr.Dropdown(choices=language_options, value="Detectar Automaticamente", label="Idioma de Origem")
156
+ txt_input = gr.Textbox(lines=5, label="Texto Original", placeholder="Digite aqui...")
157
+ btn_submit = gr.Button("Analisar e Explicar", variant="primary")
158
+
159
+ gr.Markdown("### Tradução Utilizada")
160
+ txt_translated = gr.Textbox(lines=3, label="Texto em Inglês (Base da Análise)", interactive=False)
161
+
162
+ with gr.Column(scale=1):
163
+ lbl_class = gr.Label(label="Classificação Final")
164
+ with gr.Group():
165
+ with gr.Row():
166
+ txt_p_hate = gr.Textbox(label="Prob. Hate Speech")
167
+ txt_p_offensive = gr.Textbox(label="Prob. Offensive")
168
+ with gr.Row():
169
+ txt_p_spam = gr.Textbox(label="Prob. Spam")
170
+ txt_p_neither = gr.Textbox(label="Prob. Neither")
171
+
172
+ gr.Markdown("---")
173
+ gr.Markdown("## Explicabilidade (XAI)")
174
+
175
+ with gr.Tabs():
176
+ with gr.TabItem("LIME Analysis"):
177
+ gr.Markdown("""
178
+ Mostra o peso das palavras para a classe vencedora. O LIME (Local Interpretable Model-agnostic Explanations) tenta explicar a decisão simulando pequenas variações no texto. Ele adota uma abordagem mais próxima da leitura humana, separando os termos por espaços em branco. Por isso, ele geralmente analisa palavras completas. O gráfico mostra as 6 palavras mais impactantes na decisão.
179
+ """)
180
+
181
+ gr.Markdown("""
182
+ * <span style='color: green'>VERDE:</span> Palavras que aumentaram a probabilidade da classe escolhida.
183
+ * <span style='color: red'>VERMELHO:</span> Palavras que diminuíram a confiança na escolha ou que apontaram para outra classe.
184
+ """)
185
+
186
+ plot_lime = gr.Plot(label="Gráfico LIME")
187
+
188
+ with gr.TabItem("SHAP Analysis"):
189
+ gr.Markdown("""
190
+ Mostra a magnitude do impacto (Positivo/Negativo) de cada token para a classe prevista. O SHAP (SHapley Additive exPlanations) baseia-se na Teoria dos Jogos para calcular a contribuição matemática exata de cada fragmento. Ele utiliza o vocabulário fixo do modelo (DistilBERT). Isso significa que ele identifica o texto exatamente como a rede neural: palavras desconhecidas podem ser divididas em pedaços (ex: "Jackman" vira "Jack" + "man") e sinais de pontuação (aspas, vírgulas) são avaliados isoladamente.
191
+ """)
192
+
193
+ gr.Markdown("""
194
+ * <span style='color: #ff0051'>ROSA (Positivo):</span> Tokens que empurraram a probabilidade da classe vencedora para cima. São as evidências fortes que confirmam a decisão final.
195
+ * <span style='color: #008bfb'>AZUL (Negativo):</span> Tokens que empurraram a probabilidade para baixo. São elementos que fizeram o modelo duvidar ou considerar as outras categorias.
196
+ """)
197
+
198
+ plot_shap = gr.Plot(label="Gráfico SHAP")
199
+
200
+ btn_submit.click(
201
+ fn=process_pipeline,
202
+ inputs=[dd_lang, txt_input],
203
+ outputs=[txt_translated, lbl_class, txt_p_hate, txt_p_offensive, txt_p_spam, txt_p_neither, plot_lime, plot_shap]
204
+ )
205
 
206
  demo.launch()