abersbail's picture
Deploy Customer Churn ML Predictor & Demo Video to Hugging Face
d0bac84 verified
Raw
History Blame Contribute Delete
2.33 kB
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"
):
# Load model and preprocessor
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
}])
# Preprocess
preprocessor = classifier.preprocessor
X_processed = preprocessor.transform(input_df)
# Predict
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))