BalPOS / app.py
Shah Bakhsh
Update app.py
1c45bbe verified
Raw
History Blame Contribute Delete
10.4 kB
import gradio as gr
import torch
from transformers import AutoTokenizer, AutoModelForTokenClassification
# ============================================================
# BalPOS - Balochi Part-of-Speech Tagger
# CPU Hugging Face Space
# ============================================================
MODEL_NAME = "shahbakhsh/BalPOS"
DEVICE = torch.device("cpu")
print("=" * 60)
print("BalPOS - Balochi Part-of-Speech Tagger")
print("=" * 60)
print(f"Model: {MODEL_NAME}")
print("Device: CPU")
print("Loading model...")
# ============================================================
# LOAD MODEL ONCE
# ============================================================
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
model = AutoModelForTokenClassification.from_pretrained(MODEL_NAME)
model.to(DEVICE)
model.eval()
print("Model loaded successfully.")
print(f"Labels: {model.config.num_labels}")
print("=" * 60)
# ============================================================
# POS COLORS
# ============================================================
LABEL_COLORS = {
"NOUN": "#93c5fd",
"PROPN": "#a5b4fc",
"PRON": "#c4b5fd",
"VERB": "#86efac",
"AUX": "#6ee7b7",
"ADJ": "#fde68a",
"ADV": "#fca5a5",
"ADP": "#fdba74",
"CCONJ": "#f9a8d4",
"SCONJ": "#f0abfc",
"DET": "#67e8f9",
"NUM": "#d9f99d",
"PART": "#e5e7eb",
"PUNCT": "#d1d5db",
"INTJ": "#fecaca",
"X": "#e5e7eb",
}
# ============================================================
# POS TAGGING FUNCTION
# ============================================================
def tag_sentence(sentence):
if not sentence or not sentence.strip():
return (
[],
[],
"⚠️ Please enter a Balochi sentence."
)
sentence = sentence.strip()
words = sentence.split()
# Keep CPU inference practical
if len(words) > 128:
return (
[],
[],
"⚠️ Please enter a shorter sentence. "
"Maximum supported length is 128 words."
)
try:
# Preserve word boundaries
inputs = tokenizer(
words,
is_split_into_words=True,
return_tensors="pt",
truncation=True,
max_length=256
)
word_ids = inputs.word_ids(batch_index=0)
inputs = {
key: value.to(DEVICE)
for key, value in inputs.items()
}
# CPU inference
with torch.inference_mode():
outputs = model(**inputs)
logits = outputs.logits
probabilities = torch.softmax(logits, dim=-1)
predictions = logits.argmax(dim=-1)[0]
confidences = probabilities.max(dim=-1).values[0]
results = []
table_rows = []
seen_words = set()
for token_index, word_id in enumerate(word_ids):
if word_id is None:
continue
if word_id in seen_words:
continue
if word_id >= len(words):
continue
seen_words.add(word_id)
predicted_id = predictions[token_index].item()
label = model.config.id2label.get(
predicted_id,
str(predicted_id)
)
confidence = confidences[token_index].item() * 100
word = words[word_id]
results.append((word, label))
table_rows.append([word, label, f"{confidence:.1f}%"])
if not results:
return (
[],
[],
"⚠️ No predictions were produced."
)
average_confidence = sum(
float(row[2].replace("%", ""))
for row in table_rows
) / len(table_rows)
summary = (
f"**{len(results)} tokens** · "
f"**{len(set(label for _, label in results))} POS categories** · "
f"**{average_confidence:.1f}% average confidence**"
)
return (results, table_rows, summary)
except Exception as error:
print("Prediction error:", repr(error))
return (
[],
[],
"❌ An error occurred while processing the sentence."
)
# ============================================================
# EXAMPLES
# ============================================================
EXAMPLES = [
["چہ وتی فلسفہ ءِ استاد ءَ آ اشکیتگ ات"],
["اے دوئیں دبستانانی گوناپ دیگءَ پرکے ھست۔"],
["اولی جست ءِ پسو نہ اِنت۔"],
["تہ اے پیمیں جست بے نہ بنت۔"],
]
# ============================================================
# CUSTOM CSS
# ============================================================
CUSTOM_CSS = """
.gradio-container {
max-width: 1200px !important;
margin: auto !important;
}
#hero {
text-align: center;
padding: 24px 12px 12px 12px;
}
#hero h1 {
font-size: clamp(2rem, 5vw, 3.2rem);
margin-bottom: 8px;
}
#balochi-input textarea {
direction: rtl !important;
text-align: right !important;
font-size: 1.15rem !important;
line-height: 2 !important;
}
#tag-button {
min-height: 48px !important;
font-weight: 600 !important;
}
.card {
border-radius: 16px !important;
padding: 16px !important;
}
#summary {
text-align: center;
padding: 10px;
}
@media (max-width: 768px) {
.gradio-container {
padding-left: 10px !important;
padding-right: 10px !important;
}
#hero {
padding-top: 12px;
}
#hero h1 {
font-size: 2rem;
}
#balochi-input textarea {
font-size: 1rem !important;
}
}
@media (max-width: 430px) {
#hero h1 {
font-size: 1.7rem;
}
#hero p {
font-size: 0.9rem;
}
}
footer {
display: none !important;
}
"""
# ============================================================
# BUILD GRADIO APP
# ============================================================
with gr.Blocks(title="BalPOS - Balochi POS Tagger") as demo:
# --------------------------------------------------------
# HEADER
# --------------------------------------------------------
with gr.Column(elem_id="hero"):
gr.Markdown(
"""
# 🏔️ BalPOS
## Balochi Part-of-Speech Tagger
"""
)
gr.Markdown(
"""
Automatically predict Universal Part-of-Speech (UPOS)
tags for Balochi text.
"""
)
# --------------------------------------------------------
# MAIN
# --------------------------------------------------------
with gr.Row():
# INPUT
with gr.Column(scale=1, elem_classes=["card"]):
gr.Markdown("### ✍️ Enter Balochi Text")
sentence_input = gr.Textbox(
label="Balochi Sentence",
placeholder="بلوچی جملہ ایتلاپ کن ...",
lines=5,
rtl=True,
text_align="right",
elem_id="balochi-input"
)
with gr.Row():
clear_btn = gr.Button("Clear")
submit_btn = gr.Button(
"🏷️ Tag Sentence",
variant="primary",
elem_id="tag-button"
)
gr.Examples(
examples=EXAMPLES,
inputs=sentence_input,
label="Try an example"
)
# OUTPUT
with gr.Column(scale=1, elem_classes=["card"]):
gr.Markdown("### 🔎 POS Predictions")
highlighted_output = gr.HighlightedText(
label="Tagged Text",
color_map=LABEL_COLORS,
show_legend=True
)
summary_output = gr.Markdown(
"Enter a sentence to begin.",
elem_id="summary"
)
# --------------------------------------------------------
# TABLE
# --------------------------------------------------------
with gr.Column(elem_classes=["card"]):
gr.Markdown("### 📊 Token-Level Results")
table_output = gr.Dataframe(
headers=["Token", "UPOS", "Confidence"],
datatype=["str", "str", "str"],
interactive=False,
wrap=True
)
# --------------------------------------------------------
# MODEL INFORMATION
# --------------------------------------------------------
with gr.Row():
with gr.Column(elem_classes=["card"]):
gr.Markdown(
"""
### 🧠 Model Information
**Model:** `shahbakhsh/BalPOS`
**Task:** Balochi POS Tagging
**Backbone:** BalBERT
**UPOS labels:** 16
### Reported final evaluation
| Metric | Score |
|---|---:|
| Accuracy | **87.26%** |
| Balanced Accuracy | **78.81%** |
| Macro F1 | **78.99%** |
| Weighted F1 | **87.22%** |
| MCC | **0.8526** |
"""
)
with gr.Column(elem_classes=["card"]):
gr.Markdown(
"""
### 🏷️ Supported UPOS Tags
`NOUN` · `PROPN` · `PRON` · `VERB`
`AUX` · `ADJ` · `ADV` · `ADP`
`CCONJ` · `SCONJ` · `DET` · `NUM`
`PART` · `PUNCT` · `INTJ` · `X`
"""
)
# ==========================================================
# IMPORTANT: ALL EVENT HANDLERS MUST STAY INSIDE gr.Blocks
# (this indentation is what the previous version was missing)
# ==========================================================
submit_btn.click(
fn=tag_sentence,
inputs=sentence_input,
outputs=[highlighted_output, table_output, summary_output]
)
sentence_input.submit(
fn=tag_sentence,
inputs=sentence_input,
outputs=[highlighted_output, table_output, summary_output]
)
clear_btn.click(
fn=lambda: ("", [], [], "Enter a sentence to begin."),
inputs=None,
outputs=[sentence_input, highlighted_output, table_output, summary_output]
)
# ============================================================
# LAUNCH
# ============================================================
if __name__ == "__main__":
demo.launch(
css=CUSTOM_CSS,
theme=gr.themes.Soft(
primary_hue="teal",
neutral_hue="slate"
)
)