File size: 9,139 Bytes
c739ae1
 
39da5bf
a17fac3
39da5bf
 
 
a17fac3
39da5bf
 
c739ae1
39da5bf
c739ae1
181f20a
39da5bf
bb780ab
 
181f20a
 
c739ae1
 
a17fac3
 
 
39da5bf
 
c739ae1
a17fac3
c739ae1
 
39da5bf
 
 
 
 
bb780ab
39da5bf
 
 
 
 
 
bb780ab
39da5bf
 
c739ae1
 
a17fac3
 
 
39da5bf
c739ae1
a17fac3
c739ae1
 
39da5bf
c739ae1
a17fac3
 
 
39da5bf
c739ae1
a17fac3
c739ae1
 
39da5bf
bb780ab
 
39da5bf
bb780ab
 
 
 
 
 
 
39da5bf
 
bb780ab
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
39da5bf
 
 
 
 
 
 
 
 
 
 
 
 
 
c739ae1
 
39da5bf
 
a17fac3
 
39da5bf
 
 
 
 
 
 
 
 
 
bb780ab
39da5bf
 
 
 
 
bb780ab
 
 
 
 
 
 
39da5bf
c739ae1
 
39da5bf
 
a17fac3
 
 
39da5bf
 
c739ae1
 
39da5bf
 
a17fac3
 
 
39da5bf
 
c739ae1
39da5bf
c739ae1
39da5bf
 
c739ae1
 
 
39da5bf
 
 
 
 
 
 
c739ae1
 
 
39da5bf
 
c739ae1
 
 
39da5bf
c739ae1
 
39da5bf
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
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
import gradio as gr
import os
import re
import torch
import joblib
import pandas as pd
import nltk
from transformers import AutoModelForSequenceClassification, AutoTokenizer, AutoModelForSeq2SeqLM
from nltk.corpus import stopwords
from nltk.stem import PorterStemmer

print("Booting up PredictiX Inference API...")

hf_token = os.environ.get("HF_TOKEN")
nltk.download("stopwords", quiet=True)
stop_words = set(stopwords.words("english"))
stemmer = PorterStemmer()

# 1. Load Ticket Categorization Model
try:
    cat_path = "./distilbert_category_model"
    cat_id = cat_path if os.path.exists(cat_path) else "Dinusha-Ekanayake/predictix-ticket_categorization_model"
    cat_tokenizer = AutoTokenizer.from_pretrained(cat_id, token=hf_token)
    cat_model = AutoModelForSequenceClassification.from_pretrained(cat_id, token=hf_token)
    cat_model.eval()
    print("Categorization model loaded.")
except Exception as e:
    cat_model = None
    print(f"Failed to load categorizer: {e}")

# 2. Load Ticket Priority Model
try:
    pri_path = "./ticket_priority_classification_model"
    model_file = os.path.join(pri_path, "xgboost_vehicle_priority_model_3class.pkl")
    encoder_file = os.path.join(pri_path, "priority_label_encoder.pkl")
    print(f"Priority files exist: model={os.path.exists(model_file)}, encoder={os.path.exists(encoder_file)}")
    pri_model = joblib.load(model_file)
    pri_encoder = joblib.load(encoder_file)
    print("Priority model loaded.")
except Exception as e:
    pri_model = None
    pri_encoder = None
    print(f"Failed to load priority model: {type(e).__name__}: {e}")

# 3. Load Ticket Summarization Model
try:
    sum_path = "./predictix-ticket_summarization_model"
    ts_id = sum_path if os.path.exists(sum_path) else "Dinusha-Ekanayake/predictix-ticket_summarization_model"
    ts_tokenizer = AutoTokenizer.from_pretrained(ts_id, token=hf_token)
    ts_model = AutoModelForSeq2SeqLM.from_pretrained(ts_id, token=hf_token)
    print("Ticket summarization model loaded.")
except Exception as e:
    ts_model = None
    print(f"Failed to load ticket summarizer: {e}")

# 4. Load Asset Summarization Model
try:
    as_id = "Dinusha-Ekanayake/predictix-asset_summarization_model"
    as_tokenizer = AutoTokenizer.from_pretrained(as_id, token=hf_token)
    as_model = AutoModelForSeq2SeqLM.from_pretrained(as_id, token=hf_token)
    print("Asset summarization model loaded.")
except Exception as e:
    as_model = None
    print(f"Failed to load asset summarizer: {e}")


# --- Priority helpers (exact original logic) ---
def clean_text(text):
    text = str(text).lower()
    text = text.replace("pls", "please")
    text = text.replace("asap", "as soon as possible")
    text = text.replace("veh.", "vehicle")
    text = text.replace("maint.", "maintenance")
    text = text.replace("warn.", "warning")
    text = text.replace("temp", "temperature")
    text = re.sub(r"http\S+|www\S+", " ", text)
    text = re.sub(r"[^a-zA-Z\s]", " ", text)
    text = re.sub(r"\s+", " ", text).strip()
    important_short_terms = {"ac", "abs", "rpm"}
    cleaned_tokens = []
    for token in text.split():
        if token in important_short_terms:
            cleaned_tokens.append(token)
        elif token not in stop_words and len(token) > 2:
            cleaned_tokens.append(stemmer.stem(token))
    return " ".join(cleaned_tokens)


def predict_priority(issue_text):
    cleaned_issue = clean_text(issue_text)
    issue_lower = issue_text.lower()

    if "brake" in issue_lower or "abs" in issue_lower:
        issue_category = "Brake"
        sensor_alert = "Brake system warning"
    elif "battery" in issue_lower or "charging" in issue_lower:
        issue_category = "Battery"
        sensor_alert = "Battery warning"
    elif "transmission" in issue_lower or "gear" in issue_lower or "shift" in issue_lower:
        issue_category = "Transmission"
        sensor_alert = "Transmission warning"
    elif "overheat" in issue_lower or "coolant" in issue_lower or "radiator" in issue_lower:
        issue_category = "Cooling"
        sensor_alert = "Coolant temperature high"
    elif "fuel" in issue_lower or "diesel" in issue_lower:
        issue_category = "Fuel System"
        sensor_alert = "Fuel pressure low"
    elif "steering" in issue_lower:
        issue_category = "Steering"
        sensor_alert = "Steering assist warning"
    elif "hydraulic" in issue_lower or "fork" in issue_lower:
        issue_category = "Hydraulic"
        sensor_alert = "Hydraulic pressure low"
    elif "tire" in issue_lower or "wheel" in issue_lower:
        issue_category = "Tire"
        sensor_alert = "Tire pressure warning"
    elif "ac" in issue_lower or "air condition" in issue_lower:
        issue_category = "AC"
        sensor_alert = "None"
    else:
        issue_category = "Engine"
        sensor_alert = "Check engine light"

    combined_text = cleaned_issue + " " + issue_category.lower() + " " + sensor_alert.lower()

    test_data = pd.DataFrame([{
        "combined_text": combined_text,
        "vehicle_type": "Truck",
        "issue_category": issue_category,
        "sensor_alert": sensor_alert,
        "operating_environment": "Urban",
        "weather_condition": "Normal",
        "vehicle_age": 6,
        "mileage_km": 120000,
        "downtime_hours": 8,
        "maintenance_overdue_days": 20,
        "previous_failures": 1,
    }])

    prediction_encoded = pri_model.predict(test_data)[0]
    return pri_encoder.inverse_transform([prediction_encoded])[0]


_HIGH_KEYWORDS = {
    "fire", "smoke", "explosion", "fuel leak", "brake failure", "no brakes",
    "engine seized", "total failure", "accident", "crash", "rollover",
    "unsafe to drive", "cannot drive", "vehicle stopped", "complete breakdown",
    "coolant leak", "overheating", "electrical fire", "cannot start",
}
_LOW_KEYWORDS = {
    "scratch", "dent", "cosmetic", "minor", "small crack", "paint", "sticker",
    "mirror", "wiper", "next service", "no rush", "low urgency",
    "seat cover", "floor mat", "trim", "logo", "decal", "cleaning",
}


# --- API Functions ---
def categorize(text):
    if not cat_model:
        return {"error": "Categorization model not loaded."}
    inputs = cat_tokenizer(text, return_tensors="pt", truncation=True, max_length=512)
    with torch.no_grad():
        probs = torch.softmax(cat_model(**inputs).logits, dim=-1)[0]
    scores = [
        {"label": cat_model.config.id2label[i], "score": round(float(s), 4)}
        for i, s in enumerate(probs.tolist())
    ]
    scores.sort(key=lambda x: x["score"], reverse=True)
    return scores


def prioritize(text):
    # Rule-based override first — no model needed
    lower = text.lower()
    if any(kw in lower for kw in _HIGH_KEYWORDS):
        return {"priority": "High"}
    if any(kw in lower for kw in _LOW_KEYWORDS):
        return {"priority": "Low"}
    if not pri_model:
        return {"error": "Priority model not loaded."}
    try:
        label = predict_priority(text)
        return {"priority": label}
    except Exception as e:
        return {"error": str(e)}


def summarize_ticket(text):
    if not ts_model:
        return {"error": "Ticket Summarization model not loaded."}
    inputs = ts_tokenizer(text, return_tensors="pt", truncation=True, max_length=1024)
    with torch.no_grad():
        outputs = ts_model.generate(**inputs, min_length=15, max_length=150, num_beams=4, early_stopping=True)
    return {"summary": ts_tokenizer.decode(outputs[0], skip_special_tokens=True)}


def summarize_asset(text):
    if not as_model:
        return {"error": "Asset Summarization model not loaded."}
    inputs = as_tokenizer(text, return_tensors="pt", truncation=True, max_length=1024)
    with torch.no_grad():
        outputs = as_model.generate(**inputs, min_length=20, max_length=150, num_beams=4, early_stopping=True)
    return {"summary": as_tokenizer.decode(outputs[0], skip_special_tokens=True)}


# --- Gradio UI ---
with gr.Blocks(title="PredictiX API") as demo:
    gr.Markdown("# PredictiX Internal Inference Server")

    with gr.Tab("Ticket Categorization"):
        cat_in = gr.Textbox(label="Ticket Title & Description")
        cat_out = gr.JSON(label="Categorization Result")
        gr.Button("Categorize").click(categorize, inputs=cat_in, outputs=cat_out, api_name="categorize")

    with gr.Tab("Ticket Priority"):
        pri_in = gr.Textbox(label="Ticket Title & Description")
        pri_out = gr.JSON(label="Priority Result")
        gr.Button("Prioritize").click(prioritize, inputs=pri_in, outputs=pri_out, api_name="prioritize")

    with gr.Tab("Ticket Summarization"):
        ts_in = gr.Textbox(label="Ticket Details")
        ts_out = gr.JSON(label="Summary")
        gr.Button("Summarize Ticket").click(summarize_ticket, inputs=ts_in, outputs=ts_out, api_name="summarize_ticket")

    with gr.Tab("Asset Summarization"):
        as_in = gr.Textbox(label="Asset Details")
        as_out = gr.JSON(label="Summary")
        gr.Button("Summarize Asset").click(summarize_asset, inputs=as_in, outputs=as_out, api_name="summarize_asset")

if __name__ == "__main__":
    demo.launch()