mintheinwin/Ticket-AI-Powered-Routing
0
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 