File size: 8,664 Bytes
aa53522 f57e388 aa53522 f57e388 aa53522 f0c2d4f aa53522 f0c2d4f aa53522 4570282 aa53522 f57e388 aa53522 f57e388 aa53522 f57e388 aa53522 f57e388 aa53522 f57e388 14376d9 f57e388 aa53522 7d411c6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 | import torch
import numpy as np
import gradio as gr
from transformers import AutoTokenizer, AutoModelForSequenceClassification, pipeline
from deep_translator import GoogleTranslator
import lime
from lime.lime_text import LimeTextExplainer
import shap
import matplotlib.pyplot as plt
REPO_ID = "engsoftexperimental/modelos"
MODEL_NAME = "modelo_final_distilbert2"
tokenizer = AutoTokenizer.from_pretrained(REPO_ID, subfolder=MODEL_NAME)
model = AutoModelForSequenceClassification.from_pretrained(REPO_ID, subfolder=MODEL_NAME)
labels = ["hate_speech", "offensive", "neither", "spam"]
model.config.id2label = {0: labels[0], 1: labels[1], 2: labels[2], 3: labels[3]}
model.config.label2id = {labels[0]: 0, labels[1]: 1, labels[2]: 2, labels[3]: 3}
lime_explainer = LimeTextExplainer(class_names=labels)
pipe = pipeline("text-classification", model=model, tokenizer=tokenizer, top_k=None)
shap_explainer = shap.Explainer(pipe)
def predictor_for_lime(texts):
inputs = tokenizer(texts, return_tensors="pt", truncation=True, padding=True, max_length=128).to(model.device)
with torch.no_grad():
logits = model(**inputs).logits
probs = torch.softmax(logits, dim=1).cpu().numpy()
return probs
def translate_text(text, source_lang_name, languages_dict):
if not text: return ""
if source_lang_name == "Detectar Automaticamente":
src = 'auto'
else:
src = languages_dict.get(source_lang_name, 'auto')
try:
translator = GoogleTranslator(source=src, target='en')
return translator.translate(text)
except Exception as e:
return f"Erro: {str(e)}"
def generate_lime_plot(text_en, pred_idx):
try:
exp = lime_explainer.explain_instance(
text_en,
predictor_for_lime,
num_features=6,
num_samples=100,
labels=[pred_idx]
)
fig = exp.as_pyplot_figure(label=pred_idx)
plt.tight_layout()
return fig
except Exception as e:
print(f"Erro LIME: {e}")
return None
def generate_shap_plot(text_en, pred_idx):
try:
shap_values = shap_explainer([text_en])
vals = shap_values.values
if len(vals.shape) == 3:
class_values = vals[0, :, pred_idx]
elif len(vals.shape) == 2:
class_values = vals[0, :]
else:
raise ValueError(f"Shape inesperado do SHAP: {vals.shape}")
raw_data = shap_values.data[0]
if hasattr(raw_data, 'tolist'):
tokens = raw_data.tolist()
else:
tokens = list(raw_data)
tokens = [str(t).strip() for t in tokens]
if len(tokens) != len(class_values):
min_len = min(len(tokens), len(class_values))
tokens = tokens[:min_len]
class_values = class_values[:min_len]
explanation = shap.Explanation(
values=class_values,
feature_names=tokens
)
plt.close('all')
plt.figure()
shap.plots.bar(explanation, max_display=10, show=False)
plt.tight_layout()
return plt.gcf()
except Exception as e:
print(f"Erro CRÍTICO no SHAP: {e}")
try:
print(f"Debug Shape: {shap_values.values.shape}")
print(f"Debug Tokens: {shap_values.data[0]}")
except:
pass
plt.close('all')
fig, ax = plt.subplots(figsize=(6, 2))
ax.text(0.5, 0.5, "Não foi possível gerar a explicação SHAP\npara este texto específico.",
ha='center', va='center', fontsize=10, color='red')
ax.axis('off')
return fig
def process_pipeline(source_lang, text):
try:
LANGUAGES_DICT = GoogleTranslator().get_supported_languages(as_dict=True)
except:
LANGUAGES_DICT = {'portuguese': 'pt', 'english': 'en'}
text_en = translate_text(text, source_lang, LANGUAGES_DICT)
probs = predictor_for_lime([text_en])[0]
pred_idx = int(np.argmax(probs))
pred_label = labels[pred_idx]
lime_fig = generate_lime_plot(text_en, pred_idx)
shap_fig = generate_shap_plot(text_en, pred_idx)
return (
text_en,
pred_label,
f"{probs[0]:.4f}",
f"{probs[1]:.4f}",
f"{probs[3]:.4f}",
f"{probs[2]:.4f}",
lime_fig,
shap_fig
)
try:
LANGUAGES_DICT = GoogleTranslator().get_supported_languages(as_dict=True)
except:
LANGUAGES_DICT = {'portuguese': 'pt', 'english': 'en'}
language_options = ["Detectar Automaticamente"] + list(LANGUAGES_DICT.keys())
with gr.Blocks(title="Moderador AI + XAI") as demo:
gr.Markdown("# Moderador de Conteúdo com IA Explicável (Modelo DistilBERT com Hate Speech Dataset + LIME/SHAP)")
gr.Markdown("Ferramenta experimental utilizada na pesquisa para classificação de conteúdo em quatro 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.")
gr.Markdown("Digite um texto. O sistema realiza a tradução automática, classifica e gera gráficos de explicabilidade.")
with gr.Row():
with gr.Column(scale=1):
dd_lang = gr.Dropdown(choices=language_options, value="Detectar Automaticamente", label="Idioma de Origem")
txt_input = gr.Textbox(lines=5, label="Texto Original", placeholder="Digite aqui...")
btn_submit = gr.Button("Analisar e Explicar", variant="primary")
gr.Markdown("### Tradução Utilizada")
txt_translated = gr.Textbox(lines=3, label="Texto em Inglês (Base da Análise)", interactive=False)
with gr.Column(scale=1):
lbl_class = gr.Label(label="Classificação Final")
with gr.Group():
with gr.Row():
txt_p_hate = gr.Textbox(label="Prob. Hate Speech")
txt_p_offensive = gr.Textbox(label="Prob. Offensive")
with gr.Row():
txt_p_spam = gr.Textbox(label="Prob. Spam")
txt_p_neither = gr.Textbox(label="Prob. Neither")
gr.Markdown("---")
gr.Markdown("## Explicabilidade (XAI)")
with gr.Tabs():
with gr.TabItem("LIME Analysis"):
gr.Markdown("""
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.
""")
gr.Markdown("""
* <span style='color: green'>VERDE:</span> Palavras que aumentaram a probabilidade da classe escolhida.
* <span style='color: red'>VERMELHO:</span> Palavras que diminuíram a confiança na escolha ou que apontaram para outra classe.
""")
plot_lime = gr.Plot(label="Gráfico LIME")
with gr.TabItem("SHAP Analysis"):
gr.Markdown("""
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.
""")
gr.Markdown("""
* <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.
* <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.
""")
plot_shap = gr.Plot(label="Gráfico SHAP")
btn_submit.click(
fn=process_pipeline,
inputs=[dd_lang, txt_input],
outputs=[txt_translated, lbl_class, txt_p_hate, txt_p_offensive, txt_p_spam, txt_p_neither, plot_lime, plot_shap]
)
demo.launch() |