mintheinwin/Ticket-AI-Powered-Routing-streamlitApp
0
1import streamlit as st2import torch3import torch.nn as nn4import pickle5import pandas as pd6from transformers import RobertaTokenizerFast, RobertaModel7 8# Load label mappings9with open("label_mappings.pkl", "rb") as f:10 label_mappings = pickle.load(f)11 12label_to_team = label_mappings.get("label_to_team", {})13label_to_email = label_mappings.get("label_to_email", {})14 15 16# Load tokenizer17tokenizer = RobertaTokenizerFast.from_pretrained("roberta-base")18 19# Define RoBERTa Model for multi-task classification20class 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 team_logits = self.team_classifier(cls_output)31 email_logits = self.email_classifier(cls_output)32 return team_logits, email_logits33 34# Initialize model and load checkpoint35num_teams = len(label_to_team)36num_emails = len(label_to_email)37model = RoBertaClassifier(num_teams, num_emails)38 39checkpoint = torch.load("ticket_classification_model.pth", map_location=torch.device("cpu"))40filtered_checkpoint = {k: v for k, v in checkpoint.items() if k in model.state_dict()}41model.load_state_dict(filtered_checkpoint, strict=False)42 43device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")44model.to(device)45model.eval()46 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(54 description, 55 return_tensors="pt", 56 truncation=True, 57 padding="max_length", 58 max_length=12859 ).to(device)60 with torch.no_grad():61 team_logits, email_logits = model(inputs.input_ids, inputs.attention_mask)62 predicted_team_index = team_logits.argmax(dim=-1).cpu().item()63 predicted_email_index = email_logits.argmax(dim=-1).cpu().item()64 predicted_team = label_to_team.get(predicted_team_index, "Unknown Team")65 predicted_email = label_to_email.get(predicted_email_index, "Unknown Email")66 predictions.append(67 f"{idx}. {description}\n - Assigned Team: {predicted_team}\n - Team Email: {predicted_email}\n"68 )69 csv_data.append([idx, description, predicted_team, predicted_email])70 df = pd.DataFrame(csv_data, columns=["Index", "Description", "Assigned Team", "Team Email"])71 return "\n".join(predictions), df72 73 74# Streamlit UI75st.markdown("<h2 style='text-align: center; font-size:22px;'>AI Solution for Defect Ticket Classification</h2>", unsafe_allow_html=True)76 77st.markdown("""78<p style='text-align: center; font-size:16px;'><strong>Supports:</strong> Multi-line text input & CSV upload.</p> 79<p style='text-align: center; font-size:16px;'><strong>Output:</strong> Text results & downloadable CSV file.</p> 80<p style='text-align: center; font-size:16px;'><strong>Model:</strong> Fine-tuned <strong>RoBERTa</strong> for classification.</p>81""", unsafe_allow_html=True)82 83st.markdown("<h3 style='font-size:16px;'>Enter ticket Description/Comment/Summary or upload a CSV file to predict Assigned Team & Team Email.</h3>", unsafe_allow_html=True)84 85 86# Choose input method87option = st.radio("📝 Choose Input Method", ["Enter Text", "Upload CSV"])88 89descriptions = []90if option == "Enter Text":91 text_input = st.text_area(92 "Enter Ticket Description/Comment/Summary (One per line)", 93 placeholder="Example:\n - Database performance issue\n - Login fails for admin users..."94 )95 descriptions = [line.strip() for line in text_input.split("\n") if line.strip()]96else:97 file_input = st.file_uploader("Upload CSV", type=["csv"])98 if file_input is not None:99 df_input = pd.read_csv(file_input)100 if "Description" not in df_input.columns:101 st.error("⚠️ Error: CSV must contain a 'Description' column.")102 else:103 descriptions = df_input["Description"].dropna().tolist()104 105 106# Store prediction results in session state so they persist107if "prediction_results" not in st.session_state:108 st.session_state.prediction_results = None109if "df_results" not in st.session_state:110 st.session_state.df_results = None111 112# Create a horizontal layout for the buttons113col1, col2 = st.columns([1, 1])114 115with col1:116 if st.button("PREDICT"):117 if not descriptions:118 st.error("⚠️ Please provide valid input.")119 else:120 with st.spinner("Predicting..."):121 results, df_results = predict_tickets(descriptions)122 st.session_state.prediction_results = results123 st.session_state.df_results = df_results124 125# Display prediction results if available126if st.session_state.prediction_results:127 st.markdown("<h3 style='font-size:16px;'>Prediction Results</h3>", unsafe_allow_html=True)128 st.text(st.session_state.prediction_results)129 csv_data = st.session_state.df_results.to_csv(index=False).encode('utf-8')130 st.download_button(131 label="📥 Download Predictions CSV",132 data=csv_data,133 file_name="ticket-predictions.csv",134 mime="text/csv"135 )136 137with col2:138 if st.button("CLEAR"):139 # Clear the prediction results from session state140 st.session_state.prediction_results = None141 st.session_state.df_results = None142 st.rerun()143 144st.markdown("---")145st.markdown(146 "<p style='text-align: center;color: gray; font-size:14px;'>Developed by NYP student @ Min Thein Win: Student ID: 3907578Y</p>", 147 unsafe_allow_html=True148)