import gradio as gr import os import torch from transformers import AutoModelForSequenceClassification, AutoTokenizer, AutoModelForSeq2SeqLM print("Booting up PredictiX Inference API (Direct Model Load)...") hf_token = os.environ.get("HF_TOKEN") # 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) except Exception as e: cat_model = None print(f"Failed to load categorizer: {e}") # 2. 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) except Exception as e: ts_model = None print(f"Failed to load ticket summarizer: {e}") # 3. 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) except Exception as e: as_model = None print(f"Failed to load asset summarizer: {e}") # --- 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(): logits = cat_model(**inputs).logits # Get highest score probs = torch.nn.functional.softmax(logits, dim=-1)[0] best_idx = torch.argmax(probs).item() label = cat_model.config.id2label[best_idx] score = probs[best_idx].item() return [{"label": label, "score": score}] 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) summary = ts_tokenizer.decode(outputs[0], skip_special_tokens=True) return {"summary": summary} 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) summary = as_tokenizer.decode(outputs[0], skip_special_tokens=True) return {"summary": summary} # --- Server API Interface --- 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") cat_btn = gr.Button("Categorize") cat_btn.click(categorize, inputs=cat_in, outputs=cat_out, api_name="categorize") with gr.Tab("Ticket Summarization"): ts_in = gr.Textbox(label="Ticket Details") ts_out = gr.JSON(label="Summary") ts_btn = gr.Button("Summarize Ticket") ts_btn.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") as_btn = gr.Button("Summarize Asset") as_btn.click(summarize_asset, inputs=as_in, outputs=as_out, api_name="summarize_asset") if __name__ == "__main__": demo.launch()