Shilpi Kumari commited on
Commit
a1ca7f2
·
1 Parent(s): 286258d

Add application file

Browse files
Files changed (4) hide show
  1. app.py +193 -185
  2. inference.py +70 -0
  3. requirements.txt +9 -0
  4. train.py +191 -0
app.py CHANGED
@@ -1,231 +1,239 @@
1
- # app.py - Main Gradio interface for Hugging Face Space
2
-
3
  import gradio as gr
4
  import torch
5
  from transformers import AutoTokenizer, AutoModelForSequenceClassification
6
- import numpy as np
7
  import pandas as pd
8
- from datasets import load_dataset
9
-
10
- # Load model and tokenizer from Hugging Face Hub
11
- MODEL_NAME = "Rajeshwartiwari/incident-classification-model" # Your uploaded model
12
- CACHE_DIR = "./model_cache"
13
 
14
- # Initialize components
15
- tokenizer = None
16
- model = None
17
- label2id = None
18
- id2label = None
19
-
20
- def load_components():
21
- """Load model and tokenizer"""
22
- global tokenizer, model, label2id, id2label
23
-
24
- print("Loading model and tokenizer...")
25
- tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME, cache_dir=CACHE_DIR)
26
- model = AutoModelForSequenceClassification.from_pretrained(MODEL_NAME, cache_dir=CACHE_DIR)
27
-
28
- # Get label mappings from model config
29
- if hasattr(model.config, 'id2label') and model.config.id2label:
30
- id2label = model.config.id2label
31
- label2id = {v: k for k, v in id2label.items()}
32
- else:
33
- # Fallback: load from dataset
34
- raw_dataset = load_dataset("6StringNinja/synthetic-servicenow-incidents")
35
- labels = sorted(list(set(raw_dataset['train']['category'])))
36
- label2id = {label: i for i, label in enumerate(labels)}
37
- id2label = {i: label for i, label in enumerate(labels)}
38
-
39
- print(f"Model loaded. Available categories: {list(id2label.values())}")
40
- return True
41
 
42
- def preprocess_text(text, max_length=128):
43
- """Preprocess input text"""
44
- return tokenizer(
45
- text,
46
- truncation=True,
47
- padding="max_length",
48
- max_length=max_length,
49
- return_tensors="pt"
50
- )
51
 
52
- def classify_incident(short_description, detailed_description=""):
53
- """Classify incident text"""
54
- if not tokenizer or not model:
55
- return "Error: Model not loaded", {}
56
-
57
- # Combine inputs
58
- full_text = f"{short_description} {detailed_description}".strip()
59
-
60
- if not full_text:
61
- return "Please enter incident description", {}
62
-
63
- # Preprocess
64
- inputs = preprocess_text(full_text)
65
-
66
- # Predict
67
- with torch.no_grad():
68
- outputs = model(**inputs)
69
- logits = outputs.logits
70
- probabilities = torch.nn.functional.softmax(logits, dim=-1)
71
-
72
- # Get predictions
73
- predicted_class_id = logits.argmax().item()
74
- predicted_label = id2label.get(predicted_class_id, "Unknown")
75
-
76
- # Get confidence scores for all classes
77
- confidence_scores = {}
78
- for idx, score in enumerate(probabilities[0]):
79
- label_name = id2label.get(idx, f"Class_{idx}")
80
- confidence_scores[label_name] = float(score) * 100
81
 
82
- # Sort by confidence
83
- sorted_scores = dict(sorted(confidence_scores.items(), key=lambda x: x[1], reverse=True))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
84
 
85
- return predicted_label, sorted_scores
86
-
87
- def batch_classify(file):
88
- """Classify incidents from CSV file"""
89
- try:
90
- df = pd.read_csv(file.name)
91
- results = []
92
-
93
- # Check required columns
94
- if 'short_description' not in df.columns:
95
- return "Error: CSV must contain 'short_description' column"
96
-
97
- for idx, row in df.iterrows():
98
- short_desc = str(row.get('short_description', ''))
99
- detailed_desc = str(row.get('description', ''))
100
 
101
- predicted_label, scores = classify_incident(short_desc, detailed_desc)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
102
 
103
- # Get top 3 predictions
104
- top_3 = dict(list(scores.items())[:3])
105
 
106
- results.append({
107
- 'short_description': short_desc[:100] + "..." if len(short_desc) > 100 else short_desc,
108
- 'predicted_category': predicted_label,
109
- 'confidence': f"{max(scores.values()):.1f}%" if scores else "N/A",
110
- 'top_3_predictions': ", ".join([f"{k}: {v:.1f}%" for k, v in top_3.items()])
111
- })
112
-
113
- results_df = pd.DataFrame(results)
114
- return results_df
115
-
116
- except Exception as e:
117
- return f"Error processing file: {str(e)}"
118
 
119
  # Create Gradio interface
120
  def create_interface():
121
- with gr.Blocks(title="Incident Classification System", theme=gr.themes.Soft()) as demo:
122
  gr.Markdown("""
123
- # 🚀 Incident Classification & Routing System
124
- **Automatically classify service tickets and route to the correct team**
125
-
126
- Target: 90%+ accuracy in team identification
127
  """)
128
 
 
129
  with gr.Row():
130
- with gr.Column(scale=1):
131
- gr.Markdown("### 📊 Model Info")
132
- gr.Markdown(f"""
133
- - **Model**: {MODEL_NAME.split('/')[-1]}
134
- - **Categories**: {len(id2label) if id2label else 'Loading...'}
135
- - **Purpose**: Reduce false positives & multiple routing hops
136
- """)
137
-
138
- # Load components button
139
- load_btn = gr.Button("🔄 Load/Refresh Model", variant="primary")
140
- status = gr.Textbox(label="Status", value="Click to load model", interactive=False)
141
-
142
- load_btn.click(
143
- fn=lambda: "✅ Model loaded successfully!" if load_components() else "❌ Failed to load model",
144
- outputs=status
145
  )
146
-
147
- with gr.Column(scale=2):
148
- gr.Markdown("### 🔍 Single Incident Classification")
149
-
150
- with gr.Row():
151
- short_desc = gr.Textbox(
152
- label="Short Description",
153
- placeholder="e.g., Cannot access Oracle database...",
154
- lines=2
155
- )
156
-
157
  detailed_desc = gr.Textbox(
158
  label="Detailed Description (Optional)",
159
- placeholder="Additional details, error messages, steps to reproduce...",
160
- lines=4
161
  )
 
162
 
163
- classify_btn = gr.Button("🎯 Classify Incident", variant="primary")
164
-
165
  with gr.Row():
166
  prediction = gr.Textbox(label="Predicted Category", interactive=False)
167
- confidence = gr.Textbox(label="Confidence", interactive=False)
168
 
169
- # Confidence scores visualization
170
- confidence_plot = gr.Label(
171
- label="Detailed Confidence Scores",
172
  num_top_classes=5
173
  )
174
 
 
175
  with gr.Row():
176
- gr.Markdown("### 📁 Batch Classification (CSV Upload)")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
177
 
178
- file_input = gr.File(
179
- label="Upload CSV file",
180
- file_types=[".csv"],
181
- type="filepath"
182
- )
183
 
184
- batch_btn = gr.Button("📊 Process Batch", variant="secondary")
185
- batch_output = gr.Dataframe(label="Batch Results")
 
 
 
 
 
186
 
187
- # Connect single classification
188
- classify_btn.click(
189
- fn=classify_incident,
190
  inputs=[short_desc, detailed_desc],
191
- outputs=[prediction, confidence_plot]
192
- ).then(
193
- fn=lambda x: f"{x:.1f}%" if isinstance(x, (int, float)) else x,
194
- inputs=gr.State(0), # Placeholder
195
- outputs=confidence
196
  )
197
 
198
- # Connect batch classification
199
  batch_btn.click(
200
- fn=batch_classify,
201
  inputs=file_input,
202
  outputs=batch_output
203
  )
204
-
205
- # Examples
206
- gr.Markdown("### 📋 Example Incidents")
207
- examples = [
208
- ["I cannot access the Oracle database, getting ORA-12154 error", "Tried connecting via SQL Developer but getting TNS listener error. This started after the recent patch."],
209
- ["Email client not syncing new messages", "Outlook 365 not downloading emails since 2 PM. Tried restarting and repairing Office."],
210
- ["VPN connection drops every 5 minutes", "When connected to corporate VPN, it disconnects randomly. Using Cisco AnyConnect 4.10."],
211
- ["Monitor screen flickering intermittently", "Dell UltraSharp monitor flickers when displaying dark colors. Already tried different cable."]
212
- ]
213
-
214
- gr.Examples(
215
- examples=examples,
216
- inputs=[short_desc, detailed_desc],
217
- outputs=[prediction, confidence_plot],
218
- fn=classify_incident,
219
- cache_examples=True
220
- )
221
 
222
  return demo
223
 
224
- # Initialize components on startup
225
- if load_components():
226
- print("✅ Components loaded successfully!")
227
-
228
- # Launch the interface
229
  if __name__ == "__main__":
 
 
 
 
230
  demo = create_interface()
231
- demo.launch(debug=True)
 
 
 
 
 
 
 
 
 
 
1
  import gradio as gr
2
  import torch
3
  from transformers import AutoTokenizer, AutoModelForSequenceClassification
 
4
  import pandas as pd
5
+ import json
 
 
 
 
6
 
7
+ # Configuration
8
+ MODEL_NAME = "distilbert-base-uncased-finetuned-sst-2-english" # Using a default model for testing
9
+ # Replace with your model: "Rajeshwartiwari/incident-classification-model"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
10
 
11
+ # Default categories (replace with your actual categories)
12
+ DEFAULT_CATEGORIES = [
13
+ "Software", "Hardware", "Network", "Database",
14
+ "Security", "Access", "Email", "VPN", "Application"
15
+ ]
 
 
 
 
16
 
17
+ class IncidentClassifier:
18
+ def __init__(self):
19
+ self.tokenizer = None
20
+ self.model = None
21
+ self.categories = DEFAULT_CATEGORIES
22
+ self.id2label = {i: label for i, label in enumerate(self.categories)}
23
+ self.loaded = False
24
+
25
+ def load_model(self):
26
+ """Load the model - simplified for Hugging Face Spaces"""
27
+ try:
28
+ print("Loading tokenizer and model...")
29
+ self.tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
30
+ self.model = AutoModelForSequenceClassification.from_pretrained(
31
+ MODEL_NAME,
32
+ num_labels=len(self.categories)
33
+ )
34
+ self.loaded = True
35
+ return "✅ Model loaded successfully!"
36
+ except Exception as e:
37
+ return f"❌ Error loading model: {str(e)}"
 
 
 
 
 
 
 
 
38
 
39
+ def predict(self, short_desc, detailed_desc=""):
40
+ """Make prediction"""
41
+ if not self.loaded:
42
+ return "Please load the model first", {}
43
+
44
+ if not short_desc.strip():
45
+ return "Please enter incident description", {}
46
+
47
+ # Combine text
48
+ full_text = f"{short_desc} {detailed_desc}".strip()
49
+
50
+ # Tokenize
51
+ inputs = self.tokenizer(
52
+ full_text,
53
+ return_tensors="pt",
54
+ truncation=True,
55
+ max_length=128
56
+ )
57
+
58
+ # Predict
59
+ with torch.no_grad():
60
+ outputs = self.model(**inputs)
61
+ probabilities = torch.nn.functional.softmax(outputs.logits, dim=-1)
62
+
63
+ # Get predictions
64
+ predicted_id = torch.argmax(probabilities).item()
65
+ predicted_label = self.id2label.get(predicted_id % len(self.categories), "Unknown")
66
+
67
+ # Get all confidence scores
68
+ confidences = {}
69
+ for idx, prob in enumerate(probabilities[0]):
70
+ if idx < len(self.categories):
71
+ label = self.categories[idx]
72
+ confidences[label] = float(prob) * 100
73
+ else:
74
+ break
75
+
76
+ # Sort by confidence
77
+ sorted_confidences = dict(sorted(confidences.items(), key=lambda x: x[1], reverse=True))
78
+
79
+ return predicted_label, sorted_confidences
80
 
81
+ def batch_predict(self, csv_file):
82
+ """Process batch of incidents from CSV"""
83
+ if not self.loaded:
84
+ return pd.DataFrame({"Error": ["Model not loaded"]})
85
+
86
+ try:
87
+ # Read CSV
88
+ df = pd.read_csv(csv_file.name)
 
 
 
 
 
 
 
89
 
90
+ results = []
91
+ for idx, row in df.iterrows():
92
+ short_desc = str(row.get('short_description', row.get('description', '')))
93
+ if len(short_desc) > 100:
94
+ short_display = short_desc[:100] + "..."
95
+ else:
96
+ short_display = short_desc
97
+
98
+ predicted, confidences = self.predict(short_desc)
99
+ top_conf = max(confidences.values()) if confidences else 0
100
+
101
+ results.append({
102
+ "Incident": short_display,
103
+ "Predicted Category": predicted,
104
+ "Confidence": f"{top_conf:.1f}%",
105
+ "Top 3 Predictions": ", ".join([
106
+ f"{k}: {v:.1f}%"
107
+ for k, v in list(confidences.items())[:3]
108
+ ])
109
+ })
110
 
111
+ return pd.DataFrame(results)
 
112
 
113
+ except Exception as e:
114
+ return pd.DataFrame({"Error": [f"Failed to process CSV: {str(e)}"]})
115
+
116
+ # Initialize classifier
117
+ classifier = IncidentClassifier()
 
 
 
 
 
 
 
118
 
119
  # Create Gradio interface
120
  def create_interface():
121
+ with gr.Blocks(title="Incident Classifier", theme=gr.themes.Soft()) as demo:
122
  gr.Markdown("""
123
+ # 🔧 Incident Classification System
124
+ Automatically classify and route IT incidents to the correct team
 
 
125
  """)
126
 
127
+ # Model Status Section
128
  with gr.Row():
129
+ status_box = gr.Textbox(
130
+ label="Model Status",
131
+ value="⚠️ Model not loaded",
132
+ interactive=False
133
+ )
134
+ load_btn = gr.Button("🔄 Load Model", variant="primary")
135
+
136
+ # Single Incident Classification
137
+ with gr.Row():
138
+ with gr.Column():
139
+ gr.Markdown("### 📝 Single Incident")
140
+ short_desc = gr.Textbox(
141
+ label="Short Description*",
142
+ placeholder="Brief description of the issue...",
143
+ lines=2
144
  )
 
 
 
 
 
 
 
 
 
 
 
145
  detailed_desc = gr.Textbox(
146
  label="Detailed Description (Optional)",
147
+ placeholder="Additional details, error messages...",
148
+ lines=3
149
  )
150
+ predict_btn = gr.Button("🔍 Classify", variant="primary")
151
 
152
+ # Results
 
153
  with gr.Row():
154
  prediction = gr.Textbox(label="Predicted Category", interactive=False)
155
+ confidence = gr.Textbox(label="Top Confidence", interactive=False)
156
 
157
+ confidences_chart = gr.Label(
158
+ label="Confidence Scores",
 
159
  num_top_classes=5
160
  )
161
 
162
+ # Batch Processing
163
  with gr.Row():
164
+ with gr.Column():
165
+ gr.Markdown("### 📁 Batch Processing")
166
+ gr.Markdown("Upload CSV with 'short_description' column")
167
+
168
+ file_input = gr.File(
169
+ label="Upload CSV",
170
+ file_types=[".csv"],
171
+ type="filepath"
172
+ )
173
+ batch_btn = gr.Button("📊 Process Batch", variant="secondary")
174
+ batch_output = gr.Dataframe(label="Results")
175
+
176
+ # Examples
177
+ gr.Markdown("### 💡 Example Incidents")
178
+ examples = gr.Examples(
179
+ examples=[
180
+ ["Oracle database connection error ORA-12154", "Users cannot connect to production Oracle DB"],
181
+ ["Outlook not syncing emails", "Email client stopped receiving new messages"],
182
+ ["VPN keeps disconnecting", "Cisco AnyConnect drops connection every 5 minutes"],
183
+ ["Monitor screen flickering", "Display flickers with dark backgrounds"]
184
+ ],
185
+ inputs=[short_desc, detailed_desc],
186
+ label="Try these examples:"
187
+ )
188
+
189
+ # Event Handlers
190
+ def on_load_model():
191
+ message = classifier.load_model()
192
+ if classifier.loaded:
193
+ return f"✅ Model loaded! {len(classifier.categories)} categories available"
194
+ return message
195
+
196
+ def on_predict(short, detailed):
197
+ if not classifier.loaded:
198
+ return "Please load model first", {}, "0%"
199
 
200
+ predicted, confidences = classifier.predict(short, detailed)
201
+ top_conf = max(confidences.values()) if confidences else 0
 
 
 
202
 
203
+ return predicted, confidences, f"{top_conf:.1f}%"
204
+
205
+ # Connect events
206
+ load_btn.click(
207
+ fn=on_load_model,
208
+ outputs=status_box
209
+ )
210
 
211
+ predict_btn.click(
212
+ fn=on_predict,
 
213
  inputs=[short_desc, detailed_desc],
214
+ outputs=[prediction, confidences_chart, confidence]
 
 
 
 
215
  )
216
 
 
217
  batch_btn.click(
218
+ fn=classifier.batch_predict,
219
  inputs=file_input,
220
  outputs=batch_output
221
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
222
 
223
  return demo
224
 
225
+ # Launch the app
 
 
 
 
226
  if __name__ == "__main__":
227
+ # Try to load model on startup
228
+ print("Initializing Incident Classifier...")
229
+
230
+ # Create and launch interface
231
  demo = create_interface()
232
+
233
+ # For Hugging Face Spaces, use share=False
234
+ demo.launch(
235
+ server_name="0.0.0.0",
236
+ server_port=7860,
237
+ share=False,
238
+ debug=True
239
+ )
inference.py ADDED
@@ -0,0 +1,70 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # inference.py - Simple inference script
2
+
3
+ import torch
4
+ from transformers import AutoTokenizer, AutoModelForSequenceClassification
5
+ import gradio as gr
6
+
7
+ class IncidentClassifier:
8
+ def __init__(self, model_name="Rajeshwartiwari/incident-classification-model"):
9
+ self.model_name = model_name
10
+ self.tokenizer = None
11
+ self.model = None
12
+ self.id2label = None
13
+
14
+ def load(self):
15
+ """Load the model and tokenizer"""
16
+ print(f"Loading model: {self.model_name}")
17
+ self.tokenizer = AutoTokenizer.from_pretrained(self.model_name)
18
+ self.model = AutoModelForSequenceClassification.from_pretrained(self.model_name)
19
+ self.id2label = self.model.config.id2label
20
+ print(f"Model loaded with {len(self.id2label)} categories")
21
+ return self
22
+
23
+ def predict(self, text):
24
+ """Make prediction on input text"""
25
+ if not text.strip():
26
+ return "Please enter text", {}
27
+
28
+ # Tokenize
29
+ inputs = self.tokenizer(
30
+ text,
31
+ return_tensors="pt",
32
+ truncation=True,
33
+ padding=True,
34
+ max_length=128
35
+ )
36
+
37
+ # Predict
38
+ with torch.no_grad():
39
+ outputs = self.model(**inputs)
40
+ probabilities = torch.nn.functional.softmax(outputs.logits, dim=-1)
41
+
42
+ # Get top prediction
43
+ predicted_id = outputs.logits.argmax().item()
44
+ predicted_label = self.id2label.get(predicted_id, "Unknown")
45
+
46
+ # Get all confidences
47
+ confidences = {}
48
+ for idx, prob in enumerate(probabilities[0]):
49
+ label = self.id2label.get(idx, f"Class_{idx}")
50
+ confidences[label] = float(prob) * 100
51
+
52
+ return predicted_label, confidences
53
+
54
+ # Quick test
55
+ if __name__ == "__main__":
56
+ classifier = IncidentClassifier().load()
57
+
58
+ test_cases = [
59
+ "Oracle database connection error ORA-12154",
60
+ "Email not syncing in Outlook",
61
+ "VPN keeps disconnecting every few minutes",
62
+ "Monitor screen flickering issues"
63
+ ]
64
+
65
+ for test in test_cases:
66
+ label, confidences = classifier.predict(test)
67
+ top_3 = dict(sorted(confidences.items(), key=lambda x: x[1], reverse=True)[:3])
68
+ print(f"\n📝 Input: {test}")
69
+ print(f" 🎯 Predicted: {label}")
70
+ print(f" 📊 Top 3: {top_3}")
requirements.txt ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ torch>=2.0.0
2
+ transformers>=4.36.0
3
+ gradio>=4.0.0
4
+ datasets>=2.16.0
5
+ pandas>=2.0.0
6
+ scikit-learn>=1.3.0
7
+ numpy>=1.24.0
8
+ accelerate>=0.24.0
9
+ evaluate>=0.4.0
train.py ADDED
@@ -0,0 +1,191 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # train.py - Training script for Hugging Face
2
+
3
+ import os
4
+ import torch
5
+ import pandas as pd
6
+ import numpy as np
7
+ from datasets import load_dataset, DatasetDict
8
+ from transformers import (
9
+ AutoTokenizer,
10
+ AutoModelForSequenceClassification,
11
+ TrainingArguments,
12
+ Trainer,
13
+ DataCollatorWithPadding
14
+ )
15
+ import evaluate
16
+ from sklearn.model_selection import train_test_split
17
+ import json
18
+
19
+ # Configuration
20
+ CONFIG = {
21
+ "model_name": "distilbert-base-uncased",
22
+ "dataset_name": "6StringNinja/synthetic-servicenow-incidents",
23
+ "output_dir": "./incident-classifier",
24
+ "test_size": 0.2,
25
+ "random_state": 42,
26
+ "max_length": 128,
27
+ "batch_size": 16,
28
+ "learning_rate": 2e-5,
29
+ "num_epochs": 3,
30
+ "push_to_hub": True,
31
+ "hub_model_id": "Rajeshwartiwari/incident-classification-model"
32
+ }
33
+
34
+ def load_and_prepare_data():
35
+ """Load and split dataset"""
36
+ print("📂 Loading dataset...")
37
+ dataset = load_dataset(CONFIG["dataset_name"])
38
+
39
+ # Convert to pandas for splitting
40
+ df = pd.DataFrame(dataset["train"])
41
+
42
+ # Train/test split
43
+ train_df, test_df = train_test_split(
44
+ df,
45
+ test_size=CONFIG["test_size"],
46
+ random_state=CONFIG["random_state"],
47
+ stratify=df["category"]
48
+ )
49
+
50
+ # Convert back to Hugging Face datasets
51
+ train_dataset = DatasetDict({"train": dataset["train"].new_from_pandas(train_df)})
52
+ test_dataset = DatasetDict({"test": dataset["train"].new_from_pandas(test_df)})
53
+
54
+ # Get labels
55
+ labels = sorted(list(set(df["category"])))
56
+ label2id = {label: i for i, label in enumerate(labels)}
57
+ id2label = {i: label for i, label in enumerate(labels)}
58
+
59
+ print(f"✅ Dataset loaded. Categories: {labels}")
60
+ print(f" Train samples: {len(train_df)}, Test samples: {len(test_df)}")
61
+
62
+ return train_dataset["train"], test_dataset["test"], label2id, id2label
63
+
64
+ def tokenize_function(examples, tokenizer):
65
+ """Tokenize the examples"""
66
+ # Combine short description and description
67
+ texts = [
68
+ f"{sd} {d}" if d else sd
69
+ for sd, d in zip(examples['short_description'], examples['description'])
70
+ ]
71
+
72
+ # Tokenize
73
+ tokenized = tokenizer(
74
+ texts,
75
+ truncation=True,
76
+ padding=True,
77
+ max_length=CONFIG["max_length"]
78
+ )
79
+
80
+ # Add labels
81
+ tokenized["labels"] = [label2id[l] for l in examples["category"]]
82
+ return tokenized
83
+
84
+ def compute_metrics(eval_pred):
85
+ """Compute evaluation metrics"""
86
+ metric = evaluate.load("accuracy")
87
+ logits, labels = eval_pred
88
+ predictions = np.argmax(logits, axis=-1)
89
+
90
+ # Calculate accuracy
91
+ accuracy = metric.compute(predictions=predictions, references=labels)
92
+
93
+ # You can add more metrics here
94
+ return accuracy
95
+
96
+ def main():
97
+ """Main training function"""
98
+ print("🚀 Starting Incident Classification Model Training")
99
+
100
+ # Load data
101
+ train_dataset, test_dataset, label2id, id2label = load_and_prepare_data()
102
+
103
+ # Initialize tokenizer and model
104
+ print("🔧 Initializing tokenizer and model...")
105
+ tokenizer = AutoTokenizer.from_pretrained(CONFIG["model_name"])
106
+ model = AutoModelForSequenceClassification.from_pretrained(
107
+ CONFIG["model_name"],
108
+ num_labels=len(label2id),
109
+ id2label=id2label,
110
+ label2id=label2id
111
+ )
112
+
113
+ # Tokenize datasets
114
+ print("🔠 Tokenizing datasets...")
115
+ tokenized_train = train_dataset.map(
116
+ lambda x: tokenize_function(x, tokenizer),
117
+ batched=True,
118
+ remove_columns=train_dataset.column_names
119
+ )
120
+
121
+ tokenized_test = test_dataset.map(
122
+ lambda x: tokenize_function(x, tokenizer),
123
+ batched=True,
124
+ remove_columns=test_dataset.column_names
125
+ )
126
+
127
+ # Data collator
128
+ data_collator = DataCollatorWithPadding(tokenizer=tokenizer)
129
+
130
+ # Training arguments
131
+ training_args = TrainingArguments(
132
+ output_dir=CONFIG["output_dir"],
133
+ learning_rate=CONFIG["learning_rate"],
134
+ per_device_train_batch_size=CONFIG["batch_size"],
135
+ per_device_eval_batch_size=CONFIG["batch_size"],
136
+ num_train_epochs=CONFIG["num_epochs"],
137
+ weight_decay=0.01,
138
+ evaluation_strategy="epoch",
139
+ save_strategy="epoch",
140
+ load_best_model_at_end=True,
141
+ metric_for_best_model="accuracy",
142
+ report_to="none", # Disable WandB by default
143
+ push_to_hub=CONFIG["push_to_hub"],
144
+ hub_model_id=CONFIG["hub_model_id"],
145
+ hub_strategy="end",
146
+ save_total_limit=2,
147
+ )
148
+
149
+ # Initialize trainer
150
+ trainer = Trainer(
151
+ model=model,
152
+ args=training_args,
153
+ train_dataset=tokenized_train,
154
+ eval_dataset=tokenized_test,
155
+ tokenizer=tokenizer,
156
+ data_collator=data_collator,
157
+ compute_metrics=compute_metrics,
158
+ )
159
+
160
+ # Train
161
+ print("🎯 Starting training...")
162
+ trainer.train()
163
+
164
+ # Evaluate
165
+ print("📊 Evaluating model...")
166
+ eval_results = trainer.evaluate()
167
+ print(f"✅ Evaluation results: {eval_results}")
168
+
169
+ # Save everything locally
170
+ print("💾 Saving model locally...")
171
+ trainer.save_model(CONFIG["output_dir"])
172
+ tokenizer.save_pretrained(CONFIG["output_dir"])
173
+
174
+ # Save label mappings
175
+ with open(os.path.join(CONFIG["output_dir"], "label_mappings.json"), "w") as f:
176
+ json.dump({"label2id": label2id, "id2label": id2label}, f, indent=2)
177
+
178
+ # Save configuration
179
+ with open(os.path.join(CONFIG["output_dir"], "config.json"), "w") as f:
180
+ json.dump(CONFIG, f, indent=2)
181
+
182
+ print(f"🎉 Training complete! Model saved to {CONFIG['output_dir']}")
183
+
184
+ if CONFIG["push_to_hub"]:
185
+ print("☁️ Pushing to Hugging Face Hub...")
186
+ trainer.push_to_hub()
187
+ tokenizer.push_to_hub(CONFIG["hub_model_id"])
188
+ print(f"✅ Model pushed to: https://huggingface.co/{CONFIG['hub_model_id']}")
189
+
190
+ if __name__ == "__main__":
191
+ main()