1+ import pandas as pd
2+ from sklearn .model_selection import train_test_split
3+ from sklearn .linear_model import LogisticRegression
4+ from sklearn .ensemble import RandomForestClassifier
5+ from sklearn .metrics import accuracy_score , classification_report , confusion_matrix
6+ import joblib
7+ import os
8+ from preprocess import preprocess
9+
10+ # Load and preprocess
11+ df = pd .read_csv ('data/titanic.csv' )
12+ df = preprocess (df )
13+
14+ # Split features and target
15+ X = df .drop (columns = ['Survived' ]) # everything except what we're predicting
16+ y = df ['Survived' ] # what we're predicting
17+
18+ # 80% for training, 20% for testing
19+ X_train , X_test , y_train , y_test = train_test_split (X , y , test_size = 0.2 , random_state = 42 )
20+
21+ # --- Model 1: Logistic Regression ---
22+ lr = LogisticRegression (max_iter = 200 )
23+ lr .fit (X_train , y_train )
24+ lr_preds = lr .predict (X_test )
25+ lr_acc = accuracy_score (y_test , lr_preds )
26+
27+ # --- Model 2: Random Forest ---
28+ rf = RandomForestClassifier (n_estimators = 100 , random_state = 42 )
29+ rf .fit (X_train , y_train )
30+ rf_preds = rf .predict (X_test )
31+ rf_acc = accuracy_score (y_test , rf_preds )
32+
33+ # --- Compare ---
34+ print (f"Logistic Regression Accuracy: { lr_acc :.2%} " )
35+ print (f"Random Forest Accuracy: { rf_acc :.2%} " )
36+
37+ # Pick the better model
38+ best_model = rf if rf_acc >= lr_acc else lr
39+ best_name = "Random Forest" if rf_acc >= lr_acc else "Logistic Regression"
40+ print (f"\n Best model: { best_name } " )
41+
42+ # --- Detailed report for best model ---
43+ best_preds = rf_preds if rf_acc >= lr_acc else lr_preds
44+ print ("\n Classification Report:" )
45+ print (classification_report (y_test , best_preds ))
46+ print ("Confusion Matrix:" )
47+ print (confusion_matrix (y_test , best_preds ))
48+
49+ # --- Save the best model ---
50+ os .makedirs ('models' , exist_ok = True )
51+ joblib .dump (best_model , 'models/titanic_model.pkl' )
52+ print ("\n Model saved to models/titanic_model.pkl" )
0 commit comments