"""OpenTextShield demo Space.
Loads the OpenTextShield mBERT model from the Hub and classifies an SMS as
ham (legitimate), spam or phishing, applying the same text normalisation as
the production API so obfuscated messages are handled identically. Styled to
match the OpenTextShield / TelecomsXChange (TCXC) brand.
"""
import gradio as gr
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
from normalizer import normalize_unicode
MODEL_ID = "telecomsxchange/OpenTextShield"
MAX_TOKENS = 96 # matches the production API's truncation length
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
model = AutoModelForSequenceClassification.from_pretrained(MODEL_ID)
model.eval()
DISPLAY = {
"ham": ("Legitimate", "ham", "reads like a normal message"),
"spam": ("Spam", "spam", "unwanted promotional or bulk content"),
"phishing": ("Phishing", "phishing", "an attempt to steal credentials, money or personal data"),
}
EXAMPLES = [
"Running about 15 min late, order me the usual? I'll grab the bill.",
"USPS: Your parcel could not be delivered because of an unpaid customs fee. Settle it within 24h to avoid return: http://usps-redelivery.top/pay",
"BBVA: Hemos detectado un acceso inusual a su cuenta. Verifique su identidad ahora para evitar el bloqueo: http://bbva-seguridad.info/verificar",
"CONGRATULATIONS! Your number was picked for a $1,000 gift card. Reply YES to claim before midnight!",
"Paypal: unusual sign-in detected. Confirm your identity: http://рayрal-id.com/verify",
]
def classify(message: str):
message = (message or "").strip()
if not message:
return "", {}, ""
normalized = normalize_unicode(message)
inputs = tokenizer(
normalized,
return_tensors="pt",
truncation=True,
max_length=MAX_TOKENS,
)
with torch.inference_mode():
probs = torch.softmax(model(**inputs).logits, dim=-1)[0]
top = int(probs.argmax())
raw = model.config.id2label[top]
word, css_class, meaning = DISPLAY[raw]
verdict_html = (
f'
{word}'
f'{probs[top]:.1%}'
f'— {meaning}
'
)
scores = {DISPLAY[model.config.id2label[i]][0]: float(p) for i, p in enumerate(probs)}
note = (
""
if normalized == message
else f"**Obfuscation detected** — classified as the text it imitates:\n\n> {normalized}"
)
return verdict_html, scores, note
MARK_SVG = """"""
HEADER = f"""
{MARK_SVG}
OpenTextShield
Open-source SMS spam & phishing detection, in many languages