Team Ai
Apppublic

mintheinwin/Ticket-AI-Powered-Routing-streamlitApp

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
0likes
app.py148 linesDownload Raw Back to root
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)