import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns

# Load the dataset
file_path = '/mnt/data/healthcare-dataset-stroke-data (1).csv'
df = pd.read_csv(file_path)

# Step 1: Handling Missing Values
# Fill missing 'bmi' values with the median
df['bmi'].fillna(df['bmi'].median(), inplace=True)

# Step 2: Handling Outliers
# Remove outliers using the IQR method for 'avg_glucose_level' and 'bmi'
Q1_glucose = df['avg_glucose_level'].quantile(0.25)
Q3_glucose = df['avg_glucose_level'].quantile(0.75)
IQR_glucose = Q3_glucose - Q1_glucose

Q1_bmi = df['bmi'].quantile(0.25)
Q3_bmi = df['bmi'].quantile(0.75)
IQR_bmi = Q3_bmi - Q1_bmi

# Filtering out rows that have outliers in 'avg_glucose_level' or 'bmi'
df_filtered = df[
    ~((df['avg_glucose_level'] < (Q1_glucose - 1.5 * IQR_glucose)) | (df['avg_glucose_level'] > (Q3_glucose + 1.5 * IQR_glucose)) |
      (df['bmi'] < (Q1_bmi - 1.5 * IQR_bmi)) | (df['bmi'] > (Q3_bmi + 1.5 * IQR_bmi)))
]

# Step 3: Data Visualization
# Distribution of Age
plt.figure(figsize=(10, 6))
sns.histplot(df_filtered['age'], bins=30, kde=True)
plt.title('Distribution of Age')
plt.xlabel('Age')
plt.ylabel('Frequency')
plt.show()

# Count plot of stroke occurrences by gender
plt.figure(figsize=(10, 6))
sns.countplot(data=df_filtered, x='gender', hue='stroke')
plt.title('Stroke Occurrences by Gender')
plt.xlabel('Gender')
plt.ylabel('Count')
plt.show()

# Save the cleaned dataset to a new CSV file
output_file_path = '/mnt/data/healthcare-dataset-stroke-data-cleaned.csv'
df_filtered.to_csv(output_file_path, index=False)
