Team Ai
Apppublic

mintheinwin/Ticket-AI-Powered-Routing

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
app.py198 linesDownload Raw Back to root
1import gradio as gr2import torch3import torch.nn as nn4import pickle5import pandas as pd6from transformers import RobertaTokenizerFast, RobertaModel7 8# Implement by MinTheinWin@3907578Y9# Load label mappings10with open("label_mappings.pkl", "rb") as f:11    label_mappings = pickle.load(f)12 13label_to_team = label_mappings.get("label_to_team", {})14label_to_email = label_mappings.get("label_to_email", {})15 16# Load the tokenizer17tokenizer = RobertaTokenizerFast.from_pretrained("roberta-base")18 19# Define RoBERTa Model20class RoBertaClassifier(nn.Module):21    def __init__(self, num_teams, num_emails):22        super(RoBertaClassifier, self).__init__()23        self.roberta = RobertaModel.from_pretrained("roberta-base")24        self.team_classifier = nn.Linear(self.roberta.config.hidden_size, num_teams)25        self.email_classifier = nn.Linear(self.roberta.config.hidden_size, num_emails)26 27    def forward(self, input_ids, attention_mask):28        outputs = self.roberta(input_ids=input_ids, attention_mask=attention_mask)29        cls_output = outputs.last_hidden_state[:, 0, :]30 31        team_logits = self.team_classifier(cls_output)32        email_logits = self.email_classifier(cls_output)33 34        return team_logits, email_logits35 36# Load Model37num_teams = len(label_to_team)38num_emails = len(label_to_email)39model = RoBertaClassifier(num_teams, num_emails)40checkpoint = torch.load("ticket_classification_model.pth", map_location=torch.device("cpu"))41filtered_checkpoint = {k: v for k, v in checkpoint.items() if k in model.state_dict()}42model.load_state_dict(filtered_checkpoint, strict=False)43 44device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")45model.to(device)46model.eval()47 48# Prediction Function49def predict_tickets(ticket_descriptions):50    predictions = []51    csv_data = []52    for idx, description in enumerate(ticket_descriptions, start=1):53        inputs = tokenizer(description, return_tensors="pt", truncation=True, padding="max_length", max_length=128).to(device)54        with torch.no_grad():55            team_logits, email_logits = model(inputs.input_ids, inputs.attention_mask)56        predicted_team_index = team_logits.argmax(dim=-1).cpu().item()57        predicted_email_index = email_logits.argmax(dim=-1).cpu().item()58        predicted_team = label_to_team.get(predicted_team_index, "Unknown Team")59        predicted_email = label_to_email.get(predicted_email_index, "Unknown Email")60        predictions.append(f"**{idx}. {description}**\n   - **Assigned Team:** {predicted_team}\n   - **Team Email:** {predicted_email}\n")61        csv_data.append([idx, description, predicted_team, predicted_email])62 63    df = pd.DataFrame(csv_data, columns=["Index", "Description", "Assigned Team", "Team Email"])64    csv_file = "ticket-predictions.csv"65    df.to_csv(csv_file, index=False)66    return "\n".join(predictions), csv_file67 68# Gradio Functions69def gradio_predict(option, text_input, file_input):70    if option == "Enter Text":71        descriptions = text_input.split("\n")72        descriptions = [desc.strip() for desc in descriptions if desc.strip()]73    elif option == "Upload CSV" and file_input is not None:74        df = pd.read_csv(file_input)75        if "Description" not in df.columns:76            return "⚠️ Error: CSV must contain a 'Description' column.", None77        descriptions = df["Description"].tolist()78    else:79        return "⚠️ Please provide input.", None80 81    results, csv_file = predict_tickets(descriptions)82    return results, csv_file83 84def clear_inputs():85    return "Enter Text", "", None, "", None86 87# Custom CSS for improved UI and fixed input container sizes88custom_css = """89.gradio-container { 90    max-width: 1000px !important; 91    margin: auto !important; 92}93#title { 94    text-align: center; 95    font-size: 26px !important; 96    font-weight: bold; 97}98#predict-button, #clear-button, #download-button { 99    width: 100% !important; 100    height: 55px !important; 101    font-size: 18px !important; 102}103#results-box { 104    height: 350px !important; 105    overflow-y: auto !important; 106    background: #f9f9f9; 107    padding: 15px; 108    border-radius: 10px; 109    font-size: 16px; 110}111/* Reduce vertical padding for the radio component */112#choose_input_method {113    padding-top: 5px !important;114    padding-bottom: 5px !important;115}116/* Force both input components to have the same min-height */117#text_input, #file_input {118    min-height: 200px !important;119    /* Optionally add a consistent border and padding to match styling */120    border: 1px solid #ccc;121    padding: 10px;122}123"""124 125# Gradio App UI126with gr.Blocks(css=custom_css) as app:127    gr.Markdown(128        """129        # AI Solution for Defect Ticket Classification130 131        **Supports:** Multi-line text input & CSV upload.  132        **Output:** Text results & downloadable CSV file.  133        **Model:** Fine-tuned **RoBERTa** for classification.  134 135        Enter ticket Description/Comment/Summary or upload a **CSV file** to predict Assigned Team & Team Email.136        """,137        elem_id="title"138    )139 140    with gr.Row():141        with gr.Column(scale=1):142            # Radio component with elem_id for CSS targeting143            option = gr.Radio(144                ["Enter Text", "Upload CSV"], 145                label="📝 Choose Input Method", 146                value="Enter Text", 147                elem_id="choose_input_method"148            )149 150            # Both inputs are given an element id to force consistent dimensions.151            text_input = gr.Textbox(152                label="Enter Ticket Description/Comment/Summary (One per line)", 153                visible=True,154                lines=6, 155                placeholder="Example:\n - Database performance issue\n - Login fails for admin users...",156                elem_id="text_input"157            )158 159            file_input = gr.File(160                label="📂 Upload CSV (Optional)", 161                type="filepath", 162                visible=False,163                elem_id="file_input"164            )165 166        with gr.Column(scale=1):167            gr.Markdown("## Prediction Results")168            results_output = gr.Markdown(elem_id="results-box", visible=True)169            download_csv = gr.File(label="📥 Download Predictions CSV", interactive=False)170 171    with gr.Row():172        predict_btn = gr.Button("PREDICT", variant="primary")173        clear_btn = gr.Button("CLEAR", variant="secondary")174 175    # Toggle the visibility of input components to ensure consistent sizing176    def toggle_input(selected_option):177        if selected_option == "Enter Text":178            return gr.update(visible=True), gr.update(visible=False)179        else:180            return gr.update(visible=False), gr.update(visible=True)181 182    option.change(fn=toggle_input, inputs=[option], outputs=[text_input, file_input])183    predict_btn.click(fn=gradio_predict, inputs=[option, text_input, file_input], outputs=[results_output, download_csv])184    clear_btn.click(fn=clear_inputs, inputs=[], outputs=[option, text_input, file_input, results_output, download_csv])185 186    # Footer view187    gr.Markdown("---")188    gr.HTML(189        """190        <div style="text-align: center; color: gray; padding-top: 10px;">191            <p>Developed by NYP student @ Min Thein Win: Student ID: 3907578Y</p>192        </div>193        """194    )195 196# Launch App197app.launch(share=True)198