| import argparse |
| import pandas as pd |
| import json |
| import sys |
|
|
| from src.model import ChurnClassifier |
|
|
| def predict_single( |
| age: int, |
| monthly_charges: float, |
| contract_length: int, |
| support_calls: int, |
| tech_support: str, |
| model_path: str = "models/churn_model.pkl" |
| ): |
| |
| try: |
| classifier = ChurnClassifier.load(model_path) |
| except FileNotFoundError: |
| print(f"Error: Model not found at path {model_path}. Please run train.py first.", file=sys.stderr) |
| sys.exit(1) |
| |
| input_df = pd.DataFrame([{ |
| 'age': age, |
| 'monthly_charges': monthly_charges, |
| 'contract_length': contract_length, |
| 'support_calls': support_calls, |
| 'tech_support': tech_support |
| }]) |
| |
| |
| preprocessor = classifier.preprocessor |
| X_processed = preprocessor.transform(input_df) |
| |
| |
| churn_pred = classifier.predict(X_processed)[0] |
| churn_proba = classifier.predict_proba(X_processed)[0, 1] |
| |
| result = { |
| "prediction": int(churn_pred), |
| "churn_probability": float(churn_proba), |
| "status": "Churn" if churn_pred == 1 else "No Churn" |
| } |
| return result |
|
|
| if __name__ == "__main__": |
| parser = argparse.ArgumentParser(description="Predict customer churn probability.") |
| parser.add_argument("--model_path", type=str, default="models/churn_model.pkl", help="Path to saved model") |
| parser.add_argument("--age", type=int, required=True, help="Customer age") |
| parser.add_argument("--monthly_charges", type=float, required=True, help="Monthly charges in USD") |
| parser.add_argument("--contract_length", type=int, choices=[1, 12, 24], required=True, help="Contract length in months") |
| parser.add_argument("--support_calls", type=int, required=True, help="Number of customer support calls") |
| parser.add_argument("--tech_support", type=str, choices=["yes", "no"], required=True, help="Whether tech support is active") |
| |
| args = parser.parse_args() |
| |
| res = predict_single( |
| age=args.age, |
| monthly_charges=args.monthly_charges, |
| contract_length=args.contract_length, |
| support_calls=args.support_calls, |
| tech_support=args.tech_support, |
| model_path=args.model_path |
| ) |
| |
| print(json.dumps(res, indent=2)) |
|
|