File size: 2,588 Bytes
741e25f
 
 
 
 
 
 
 
 
4d44920
 
 
 
741e25f
 
 
 
 
 
 
 
 
 
4d44920
 
741e25f
4d44920
 
 
 
 
741e25f
 
4d44920
 
741e25f
4d44920
741e25f
4d44920
741e25f
 
 
 
 
4d44920
741e25f
 
 
4e37478
d048b52
 
4e37478
 
 
 
741e25f
 
 
 
 
4d44920
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
import gradio as gr
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer

# Load the model and tokenizer from Hugging Face
model_name = "TextLabRUET/xlm-r_based_bangla_sentence_classifier"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForSequenceClassification.from_pretrained(model_name)

# Set device (GPU if available, otherwise CPU)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)

# Mapping predicted class to Bangla sentence types
class_mapping = {
    0: "Assertive Sentence (বর্ণনামূলক বাক্য)",
    1: "Interrogative Sentence (প্রশ্নবোধক বাক্য)",
    2: "Imperative Sentence (অনুজ্ঞাসূচক বাক্য)",
    3: "Optative Sentence (প্রার্থনা সূচক বাক্য)",
    4: "Exclamatory Sentence (বিস্ময়সূচক বাক্য)"
}

# Function for prediction
def predict_bangla_sentence(sentence):
    # Tokenize the input sentence
    inputs = tokenizer(sentence, return_tensors="pt", truncation=True, padding=True, max_length=128)
    
    # Move input tensors to the same device as the model
    inputs = {key: val.to(device) for key, val in inputs.items()}
    
    # Perform inference
    with torch.no_grad():
        outputs = model(**inputs)
    
    # Get the predicted class
    logits = outputs.logits
    predicted_class = torch.argmax(logits, dim=-1).item()
    
    # Return the predicted sentence type
    sentence_type = class_mapping.get(predicted_class, "Unknown Sentence Type")
    return f"Predicted Class: {sentence_type}"

# Create Gradio UI
iface = gr.Interface(
    fn=predict_bangla_sentence,
    inputs=gr.Textbox(lines=2, placeholder="Enter a Bangla sentence..."),
    outputs="text",
    title="Bangla Sentence Classifier",
    description = (
        "This model was trained using a curated Bangla dataset by **TextLab RUET**. "
        "It classifies Bangla sentences into five distinct categories: Assertive, Interrogative, Imperative, Optative, and Exclamatory "
        "using the **XLM-R** (a multilingual transformer model based on the BERT architecture).\n"
        "Enter a Bangla sentence below to see how our model interprets it!\n\n"
        "**Note**: While we strive for accuracy, the model may occasionally misclassify sentences due to dataset limitations. "
        "We apologize for any errors and appreciate your understanding."
    ),
    theme="compact"
)

# Launch the Gradio app
iface.launch(share=True)