| 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") |
|
|
| |
| 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}") |
|
|
| |
| 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}") |
|
|
| |
| 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}") |
|
|
| |
| 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 |
| |
| 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} |
|
|
| |
| 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() |