AI_Detector / app.py
jay123jay's picture
also print the possibility of which model is used
e8ebc4c verified
Raw
History Blame
6.17 kB
import gradio as gr
from transformers import AutoTokenizer, AutoModelForSequenceClassification
import torch
import re
from tokenizers.normalizers import Sequence, Replace, Strip
from tokenizers import Regex
import matplotlib.pyplot as plt
# --- Setup and Model Loading ---
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# Model paths and URLs
model1_path = "modernbert.bin"
model2_url = "https://huggingface.co/mihalykiss/modernbert_2/resolve/main/Model_groups_3class_seed12"
model3_url = "https://huggingface.co/mihalykiss/modernbert_2/resolve/main/Model_groups_3class_seed22"
tokenizer = AutoTokenizer.from_pretrained("answerdotai/ModernBERT-base")
def load_model(path_or_url, num_labels=41):
model = AutoModelForSequenceClassification.from_pretrained("answerdotai/ModernBERT-base", num_labels=num_labels)
if path_or_url.startswith("http"):
# Load from Hugging Face URL
state_dict = torch.hub.load_state_dict_from_url(path_or_url, map_location=device)
else:
# Load from local file
state_dict = torch.load(path_or_url, map_location=device)
model.load_state_dict(state_dict)
return model.to(device).eval()
# Initializing the ensemble
model_1 = load_model(model1_path)
model_2 = load_model(model2_url)
model_3 = load_model(model3_url)
label_mapping = {
0: '13B', 1: '30B', 2: '65B', 3: '7B', 4: 'GLM130B', 5: 'bloom_7b',
6: 'bloomz', 7: 'cohere', 8: 'davinci', 9: 'dolly', 10: 'dolly-v2-12b',
11: 'flan_t5_base', 12: 'flan_t5_large', 13: 'flan_t5_small',
14: 'flan_t5_xl', 15: 'flan_t5_xxl', 16: 'gemma-7b-it', 17: 'gemma2-9b-it',
18: 'gpt-3.5-turbo', 19: 'gpt-35', 20: 'gpt4', 21: 'gpt4o',
22: 'gpt_j', 23: 'gpt_neox', 24: 'human', 25: 'llama3-70b', 26: 'llama3-8b',
27: 'mixtral-8x7b', 28: 'opt_1.3b', 29: 'opt_125m', 30: 'opt_13b',
31: 'opt_2.7b', 32: 'opt_30b', 33: 'opt_350m', 34: 'opt_6.7b',
35: 'opt_iml_30b', 36: 'opt_iml_max_1.3b', 37: 't0_11b', 38: 't0_3b',
39: 'text-davinci-002', 40: 'text-davinci-003'
}
# --- Text Preprocessing ---
def clean_text(text: str) -> str:
text = re.sub(r'\s{2,}', ' ', text)
text = re.sub(r'\s+([,.;:?!])', r'\1', text)
return text
# Custom Tokenizer Normalization
newline_to_space = Replace(Regex(r'\s*\n\s*'), " ")
join_hyphen_break = Replace(Regex(r'(\w+)[--]\s*\n\s*(\w+)'), r"\1\2")
tokenizer.backend_tokenizer.normalizer = Sequence([
tokenizer.backend_tokenizer.normalizer,
join_hyphen_break,
newline_to_space,
Strip()
])
# --- Core Classification Logic ---
def classify_text(text):
cleaned_text = clean_text(text)
if not cleaned_text.strip():
return "Please enter text to analyze.", None
inputs = tokenizer(cleaned_text, return_tensors="pt", truncation=True, padding=True).to(device)
with torch.no_grad():
logits_1 = model_1(**inputs).logits
logits_2 = model_2(**inputs).logits
logits_3 = model_3(**inputs).logits
# Soft voting ensemble (averaging probabilities)
s1, s2, s3 = torch.softmax(logits_1, 1), torch.softmax(logits_2, 1), torch.softmax(logits_3, 1)
avg_probs = (s1 + s2 + s3) / 3
probs = avg_probs[0]
# Calculate probabilities for the 2 main categories (Human vs AI)
human_prob = probs[24].item()
# To find the specific LLM, we ignore the 'human' index (24)
ai_probs_only = probs.clone()
ai_probs_only[24] = 0
ai_total_prob = ai_probs_only.sum().item()
# Normalize percentages
total_sum = human_prob + ai_total_prob
human_pct = (human_prob / total_sum) * 100
ai_pct = (ai_total_prob / total_sum) * 100
# Identify the specific AI model with the highest sub-probability
top_ai_idx = torch.argmax(ai_probs_only).item()
predicted_llm = label_mapping[top_ai_idx]
# Construct the result display
if human_pct > ai_pct:
result_message = (
f"### Result: <span class='highlight-human'>**{human_pct:.2f}% Human written**</span>\n\n"
"The content matches human writing patterns."
)
else:
result_message = (
f"### Result: <span class='highlight-ai'>**{ai_pct:.2f}% AI generated**</span>\n\n"
f"**Specific Model Identified:** `{predicted_llm}`\n\n"
"The structure and syntax are highly characteristic of this LLM."
)
# Visualization
fig, ax = plt.subplots(figsize=(8, 4))
bars = ax.bar(['Human', 'AI'], [human_pct, ai_pct], color=['#4CAF50', '#FF5733'], alpha=0.8)
ax.set_ylabel('Probability (%)')
ax.set_title('Detection Probability')
ax.set_ylim(0, 110) # Room for text labels
for bar in bars:
height = bar.get_height()
ax.text(bar.get_x() + bar.get_width()/2., height + 2, f'{height:.1f}%', ha='center', fontweight='bold')
plt.tight_layout()
return result_message, fig
# --- Gradio UI Layout ---
with gr.Blocks(css="""
.highlight-human { color: #4CAF50; font-weight: bold; background: rgba(76, 175, 80, 0.1); padding: 5px; border-radius: 5px; }
.highlight-ai { color: #FF5733; font-weight: bold; background: rgba(255, 87, 51, 0.1); padding: 5px; border-radius: 5px; }
#output-container { text-align: center; padding: 20px; }
""") as iface:
gr.Markdown("# AI Text Detector & LLM Identifier")
gr.Markdown("This tool uses an ensemble of **ModernBERT** models to predict if text is human or AI, and specifies the likely source model.")
with gr.Row():
with gr.Column(scale=2):
text_input = gr.Textbox(
label="Input Text",
placeholder="Paste your English text here...",
lines=10
)
with gr.Column(scale=1):
result_output = gr.Markdown("Analysis results will appear here.", elem_id="output-container")
plot_output = gr.Plot()
# Trigger classification on text change
text_input.change(classify_text, inputs=text_input, outputs=[result_output, plot_output])
gr.Markdown("---")
gr.Markdown("**Developed by SzegedAI**")
if __name__ == "__main__":
# share=True creates a public URL
iface.launch(share=True)