Team Ai
Apppublic

giodesi/Multi-class_Classification

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes
app.py1085 linesDownload Raw Back to root
1"""2Multi-class Classification Web App3A generalized Streamlit application for multi-class classification on any CSV dataset4"""5 6import streamlit as st7import pandas as pd8import numpy as np9import matplotlib.pyplot as plt10import seaborn as sns11from sklearn.model_selection import train_test_split12from sklearn.preprocessing import OneHotEncoder, StandardScaler, LabelEncoder13from sklearn.linear_model import LogisticRegression14from sklearn.multiclass import OneVsOneClassifier, OneVsRestClassifier15from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, confusion_matrix, classification_report16import warnings17warnings.filterwarnings('ignore')18 19# Set page configuration20st.set_page_config(21    page_title="Multi-class Classification App",22    page_icon="๐Ÿ“Š",23    layout="wide",24    initial_sidebar_state="expanded"25)26 27# Custom CSS for better styling28st.markdown("""29    <style>30    .main-header {31        text-align: center;32        padding: 1rem 0;33        background: linear-gradient(90deg, #667eea 0%, #764ba2 100%);34        color: white;35        border-radius: 10px;36        margin-bottom: 2rem;37    }38    .metric-card {39        background-color: #f7f9fc;40        padding: 1rem;41        border-radius: 10px;42        border-left: 4px solid #667eea;43    }44    </style>45    """, unsafe_allow_html=True)46 47# Header48st.markdown("<div class='main-header'><h1>๐ŸŽฏ Multi-class Classification Web App</h1></div>", unsafe_allow_html=True)49 50# Initialize session state51if 'data' not in st.session_state:52    st.session_state['data'] = None53if 'preprocessed_data' not in st.session_state:54    st.session_state['preprocessed_data'] = None55if 'model' not in st.session_state:56    st.session_state['model'] = None57 58# Sidebar for configuration59with st.sidebar:60    st.header("๐Ÿ“ Data Upload")61    62    # File uploader63    uploaded_file = st.file_uploader(64        "Choose a CSV file",65        type="csv",66        help="Upload your dataset in CSV format"67    )68    69    if uploaded_file is not None:70        # Load data71        st.session_state['data'] = pd.read_csv(uploaded_file)72        st.success(f"โœ… Data loaded successfully! Shape: {st.session_state['data'].shape}")73        74        # Target column selection75        st.header("๐ŸŽฏ Target Selection")76        columns = st.session_state['data'].columns.tolist()77        target_column = st.selectbox(78            "Select target column",79            columns,80            help="Choose the column you want to predict"81        )82        83        # Feature selection84        st.header("๐Ÿ“Š Feature Selection")85        feature_columns = st.multiselect(86            "Select feature columns",87            [col for col in columns if col != target_column],88            default=[col for col in columns if col != target_column],89            help="Choose the features to use for training"90        )91        92        # Preprocessing options93        st.header("โš™๏ธ Preprocessing Options")94        scale_features = st.checkbox("Standardize numerical features", value=True)95        encode_categorical = st.checkbox("One-hot encode categorical features", value=True)96        97        # Model configuration98        st.header("๐Ÿค– Model Configuration")99        100        # Classification algorithm selection (only Logistic Regression)101        classifier_type = "Logistic Regression"102        st.info("๐Ÿ“Œ Using Logistic Regression classifier")103        104        # Multi-class strategy105        multiclass_strategy = st.selectbox(106            "Select multi-class strategy",107            ["One-vs-Rest (OvR)", "One-vs-One (OvO)", "Auto"],108            help="Choose the strategy for handling multiple classes"109        )110        111        # Test size112        test_size = st.slider(113            "Test set size",114            min_value=0.1,115            max_value=0.5,116            value=0.2,117            step=0.05,118            help="Proportion of data to use for testing"119        )120        121        # Random state122        random_state = st.number_input(123            "Random state",124            value=42,125            help="Set random state for reproducibility"126        )127 128# Main content area129if st.session_state['data'] is not None:130    131    # Create tabs for different sections132    tab1, tab2, tab3, tab4, tab5, tab6 = st.tabs(["๐Ÿ“Š Data Overview", "๐Ÿ” EDA", "โš™๏ธ Preprocessing", "๐ŸŽฏ Training", "๐Ÿ“ˆ Results", "๐Ÿ”ฎ Predict New Data"])133    134    with tab1:135        st.header("Data Overview")136        137        col1, col2, col3 = st.columns(3)138        with col1:139            st.metric("Total Rows", st.session_state['data'].shape[0])140        with col2:141            st.metric("Total Columns", st.session_state['data'].shape[1])142        with col3:143            st.metric("Missing Values", st.session_state['data'].isnull().sum().sum())144        145        # Display first few rows146        st.subheader("First 10 Rows")147        st.dataframe(st.session_state['data'].head(10))148        149        # Data types150        st.subheader("Data Types")151        col1, col2 = st.columns(2)152        with col1:153            st.write("**Column Information:**")154            info_df = pd.DataFrame({155                'Column': st.session_state['data'].columns,156                'Type': st.session_state['data'].dtypes.astype(str),  # Convert dtype to string157                'Non-Null Count': st.session_state['data'].count(),158                'Null Count': st.session_state['data'].isnull().sum()159            })160            st.dataframe(info_df)161        162        with col2:163            st.write("**Basic Statistics:**")164            st.dataframe(st.session_state['data'].describe())165    166    with tab2:167        st.header("Exploratory Data Analysis")168        169        if 'target_column' in locals():170            # Target distribution171            st.subheader("Target Variable Distribution")172            173            col1, col2 = st.columns(2)174            175            with col1:176                fig, ax = plt.subplots(figsize=(8, 6))177                target_counts = st.session_state['data'][target_column].value_counts()178                ax.bar(target_counts.index.astype(str), target_counts.values, color='#667eea')179                ax.set_xlabel(target_column)180                ax.set_ylabel("Count")181                ax.set_title("Target Distribution")182                plt.xticks(rotation=45, ha='right')183                st.pyplot(fig)184            185            with col2:186                fig, ax = plt.subplots(figsize=(8, 6))187                target_counts.plot(kind='pie', ax=ax, autopct='%1.1f%%', startangle=90)188                ax.set_title("Target Distribution (Percentage)")189                ax.set_ylabel("")190                st.pyplot(fig)191            192            # Feature distributions193            st.subheader("Feature Distributions")194            195            if 'feature_columns' in locals() and len(feature_columns) > 0:196                # Numerical features197                numerical_features = st.session_state['data'][feature_columns].select_dtypes(include=[np.number]).columns.tolist()198                199                if len(numerical_features) > 0:200                    st.write("**Numerical Features:**")201                    202                    # Create subplots for numerical features203                    n_cols = min(3, len(numerical_features))204                    n_rows = (len(numerical_features) + n_cols - 1) // n_cols205                    206                    fig, axes = plt.subplots(n_rows, n_cols, figsize=(15, 4*n_rows))207                    axes = axes.flatten() if n_rows * n_cols > 1 else [axes]208                    209                    for i, col in enumerate(numerical_features[:9]):  # Limit to 9 features for display210                        axes[i].hist(st.session_state['data'][col].dropna(), bins=30, color='#764ba2', alpha=0.7)211                        axes[i].set_title(f"Distribution of {col}")212                        axes[i].set_xlabel(col)213                        axes[i].set_ylabel("Frequency")214                    215                    # Hide unused subplots216                    for i in range(len(numerical_features), len(axes)):217                        axes[i].axis('off')218                    219                    plt.tight_layout()220                    st.pyplot(fig)221                222                # Categorical features223                categorical_features = st.session_state['data'][feature_columns].select_dtypes(include=['object']).columns.tolist()224                225                if len(categorical_features) > 0:226                    st.write("**Categorical Features:**")227                    228                    for col in categorical_features[:5]:  # Limit to 5 features for display229                        st.write(f"*{col}*")230                        value_counts = st.session_state['data'][col].value_counts()231                        st.bar_chart(value_counts)232    233    with tab3:234        st.header("Data Preprocessing")235        236        if 'target_column' in locals() and 'feature_columns' in locals():237            238            if st.button("๐Ÿ”„ Apply Preprocessing"):239                with st.spinner("Preprocessing data..."):240                    241                    # Create a copy of the data242                    processed_data = st.session_state['data'][feature_columns + [target_column]].copy()243                    244                    # Handle missing values245                    st.write("**Handling Missing Values:**")246                    missing_counts = processed_data.isnull().sum()247                    248                    # Store training data statistics for later use in predictions249                    training_medians = {}250                    training_modes = {}251                    252                    if missing_counts.sum() > 0:253                        # Fill numerical columns with median254                        numerical_cols = processed_data.select_dtypes(include=[np.number]).columns255                        for col in numerical_cols:256                            if col != target_column:257                                median_val = processed_data[col].median()258                                training_medians[col] = median_val259                                if processed_data[col].isnull().any():260                                    processed_data[col].fillna(median_val, inplace=True)261                        262                        # Fill categorical columns with mode263                        categorical_cols = processed_data.select_dtypes(include=['object']).columns264                        for col in categorical_cols:265                            if col != target_column:266                                mode_val = processed_data[col].mode()[0] if len(processed_data[col].mode()) > 0 else 'unknown'267                                training_modes[col] = mode_val268                                if processed_data[col].isnull().any():269                                    processed_data[col].fillna(mode_val, inplace=True)270                        271                        st.success("โœ… Missing values handled")272                    else:273                        # Still store statistics even if no missing values274                        numerical_cols = processed_data.select_dtypes(include=[np.number]).columns275                        for col in numerical_cols:276                            if col != target_column:277                                training_medians[col] = processed_data[col].median()278                        279                        categorical_cols = processed_data.select_dtypes(include=['object']).columns280                        for col in categorical_cols:281                            if col != target_column and len(processed_data[col].mode()) > 0:282                                training_modes[col] = processed_data[col].mode()[0]283                        284                        st.info("No missing values found")285                    286                    # Store training statistics in session state287                    st.session_state['training_medians'] = training_medians288                    st.session_state['training_modes'] = training_modes289                    290                    # Separate features and target291                    X = processed_data.drop(columns=[target_column])292                    y = processed_data[target_column]293                    294                    # Encode target variable if it's categorical295                    if y.dtype == 'object':296                        label_encoder = LabelEncoder()297                        y = label_encoder.fit_transform(y)298                        st.session_state['label_encoder'] = label_encoder299                        st.success(f"โœ… Target variable encoded: {len(np.unique(y))} classes")300                    301                    # Identify numerical and categorical columns302                    numerical_columns = X.select_dtypes(include=[np.number]).columns.tolist()303                    categorical_columns = X.select_dtypes(include=['object']).columns.tolist()304                    305                    # Standardize numerical features306                    if scale_features and len(numerical_columns) > 0:307                        st.write("**Standardizing Numerical Features:**")308                        scaler = StandardScaler()309                        X[numerical_columns] = scaler.fit_transform(X[numerical_columns])310                        st.session_state['scaler'] = scaler311                        st.success(f"โœ… Standardized {len(numerical_columns)} numerical features")312                    313                    # One-hot encode categorical features314                    if encode_categorical and len(categorical_columns) > 0:315                        st.write("**One-hot Encoding Categorical Features:**")316                        encoder = OneHotEncoder(sparse_output=False, drop='first')317                        encoded_features = encoder.fit_transform(X[categorical_columns])318                        319                        # Create dataframe with encoded features320                        encoded_df = pd.DataFrame(321                            encoded_features,322                            columns=encoder.get_feature_names_out(categorical_columns),323                            index=X.index324                        )325                        326                        # Drop original categorical columns and concatenate encoded ones327                        X = pd.concat([X.drop(columns=categorical_columns), encoded_df], axis=1)328                        st.session_state['encoder'] = encoder329                        st.success(f"โœ… Encoded {len(categorical_columns)} categorical features โ†’ {encoded_df.shape[1]} new features")330                    331                    # Store preprocessed data332                    st.session_state['X'] = X333                    st.session_state['y'] = y334                    st.session_state['preprocessed_data'] = pd.concat([X, pd.Series(y, name=target_column, index=X.index)], axis=1)335                    336                    # Display preprocessed data info337                    st.write("**Preprocessed Data Summary:**")338                    col1, col2, col3 = st.columns(3)339                    with col1:340                        st.metric("Features", X.shape[1])341                    with col2:342                        st.metric("Samples", X.shape[0])343                    with col3:344                        st.metric("Classes", len(np.unique(y)))345                    346                    st.write("**Feature Names after Preprocessing:**")347                    st.write(list(X.columns))348        else:349            st.warning("โš ๏ธ Please select target and feature columns in the sidebar")350    351    with tab4:352        st.header("Model Training")353        354        if 'X' in st.session_state and 'y' in st.session_state:355            356            col1, col2 = st.columns(2)357            358            with col1:359                st.write("**Data Split Configuration:**")360                st.write(f"โ€ข Training set: {100-int(test_size*100)}%")361                st.write(f"โ€ข Test set: {int(test_size*100)}%")362                st.write(f"โ€ข Random state: {random_state}")363            364            with col2:365                st.write("**Model Configuration:**")366                st.write(f"โ€ข Classifier: Logistic Regression")367                st.write(f"โ€ข Strategy: {multiclass_strategy}")368            369            if st.button("๐Ÿš€ Train Model", type="primary"):370                with st.spinner("Training model..."):371                    372                    # Split data373                    X_train, X_test, y_train, y_test = train_test_split(374                        st.session_state['X'], 375                        st.session_state['y'],376                        test_size=test_size,377                        random_state=random_state,378                        stratify=st.session_state['y']379                    )380                    381                    # Store split data382                    st.session_state['X_train'] = X_train383                    st.session_state['X_test'] = X_test384                    st.session_state['y_train'] = y_train385                    st.session_state['y_test'] = y_test386                    387                    # Create base classifier - Only Logistic Regression388                    base_classifier = LogisticRegression(max_iter=1000, random_state=random_state)389                    390                    # Apply multi-class strategy391                    if multiclass_strategy == "One-vs-Rest (OvR)":392                        model = LogisticRegression(multi_class='ovr', max_iter=1000, random_state=random_state)393                    elif multiclass_strategy == "One-vs-One (OvO)":394                        model = OneVsOneClassifier(base_classifier)395                    else:  # Auto396                        model = LogisticRegression(multi_class='auto', max_iter=1000, random_state=random_state)397                    398                    # Train model399                    model.fit(X_train, y_train)400                    st.session_state['model'] = model401                    402                    # Make predictions403                    y_pred_train = model.predict(X_train)404                    y_pred_test = model.predict(X_test)405                    406                    st.session_state['y_pred_train'] = y_pred_train407                    st.session_state['y_pred_test'] = y_pred_test408                    409                    # Calculate metrics410                    train_accuracy = accuracy_score(y_train, y_pred_train)411                    test_accuracy = accuracy_score(y_test, y_pred_test)412                    413                    st.success("โœ… Model trained successfully!")414                    415                    # Display training results416                    col1, col2 = st.columns(2)417                    with col1:418                        st.metric("Training Accuracy", f"{train_accuracy:.3f}")419                    with col2:420                        st.metric("Test Accuracy", f"{test_accuracy:.3f}")421        else:422            st.warning("โš ๏ธ Please preprocess the data first in the 'Preprocessing' tab")423    424    with tab5:425        st.header("Model Results & Analysis")426        427        if 'model' in st.session_state and st.session_state['model'] is not None:428            429            # Performance Metrics430            st.subheader("๐Ÿ“Š Performance Metrics")431            432            y_test = st.session_state['y_test']433            y_pred = st.session_state['y_pred_test']434            435            # Calculate metrics436            accuracy = accuracy_score(y_test, y_pred)437            438            # For multi-class, use weighted average439            precision = precision_score(y_test, y_pred, average='weighted', zero_division=0)440            recall = recall_score(y_test, y_pred, average='weighted', zero_division=0)441            f1 = f1_score(y_test, y_pred, average='weighted', zero_division=0)442            443            # Display metrics in columns444            col1, col2, col3, col4 = st.columns(4)445            with col1:446                st.metric("Accuracy", f"{accuracy:.3f}")447            with col2:448                st.metric("Precision", f"{precision:.3f}")449            with col3:450                st.metric("Recall", f"{recall:.3f}")451            with col4:452                st.metric("F1-Score", f"{f1:.3f}")453            454            # Confusion Matrix455            st.subheader("๐Ÿ”„ Confusion Matrix")456            457            cm = confusion_matrix(y_test, y_pred)458            459            fig, ax = plt.subplots(figsize=(10, 8))460            sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=ax)461            ax.set_xlabel('Predicted')462            ax.set_ylabel('Actual')463            ax.set_title('Confusion Matrix')464            465            # Add class labels if available466            if 'label_encoder' in st.session_state:467                classes = st.session_state['label_encoder'].classes_468                ax.set_xticklabels(classes, rotation=45, ha='right')469                ax.set_yticklabels(classes, rotation=0)470            471            st.pyplot(fig)472            473            # Classification Report474            st.subheader("๐Ÿ“‹ Detailed Classification Report")475            476            report = classification_report(y_test, y_pred, output_dict=True)477            report_df = pd.DataFrame(report).transpose()478            479            # Format the dataframe for better display480            report_df = report_df.round(3)481            st.dataframe(report_df, use_container_width=True)482            483            # Feature Importance (if applicable)484            st.subheader("๐ŸŽฏ Feature Importance Analysis")485            486            feature_names = st.session_state['X'].columns.tolist()487            488            # For logistic regression, use coefficient magnitudes489            if hasattr(st.session_state['model'], 'coef_'):490                # Average absolute coefficients across all classes491                importance = np.mean(np.abs(st.session_state['model'].coef_), axis=0)492                493                # Create feature importance dataframe494                importance_df = pd.DataFrame({495                    'Feature': feature_names,496                    'Importance': importance497                }).sort_values('Importance', ascending=False)498                499                # Plot feature importance500                fig, ax = plt.subplots(figsize=(10, 6))501                top_features = importance_df.head(15)  # Show top 15 features502                ax.barh(top_features['Feature'], top_features['Importance'], color='#764ba2')503                ax.set_xlabel('Average Absolute Coefficient')504                ax.set_title('Top 15 Feature Importances (Coefficient Magnitudes)')505                plt.tight_layout()506                st.pyplot(fig)507            508            # Model Export509            st.subheader("๐Ÿ’พ Export Model & Results")510            511            col1, col2 = st.columns(2)512            513            with col1:514                # Export predictions with all original data515                # Include all original columns from the test set for identification516                predictions_df = pd.DataFrame({517                    'Actual': y_test,518                    'Predicted': y_pred519                })520                521                # Add the original features from test set522                X_test_original = st.session_state['X_test'].copy()523                524                # If we have the original unprocessed data, use it for better readability525                if 'data' in st.session_state and st.session_state['data'] is not None:526                    try:527                        # Get the original indices528                        test_indices = X_test_original.index529                        530                        # Get original data for these indices (all columns except target)531                        original_test_data = st.session_state['data'].loc[test_indices].copy()532                        533                        # Remove target column if it exists534                        if target_column in original_test_data.columns:535                            original_test_data = original_test_data.drop(columns=[target_column])536                        537                        # Combine original data with predictions538                        full_predictions_df = pd.concat([539                            original_test_data.reset_index(drop=True),540                            predictions_df.reset_index(drop=True)541                        ], axis=1)542                    except:543                        # Fallback to processed features if original data retrieval fails544                        full_predictions_df = pd.concat([545                            X_test_original.reset_index(drop=True),546                            predictions_df.reset_index(drop=True)547                        ], axis=1)548                else:549                    # No original data available, use processed features550                    full_predictions_df = pd.concat([551                        X_test_original.reset_index(drop=True),552                        predictions_df.reset_index(drop=True)553                    ], axis=1)554                555                # If label encoder exists, decode the values556                if 'label_encoder' in st.session_state:557                    full_predictions_df['Actual_Label'] = st.session_state['label_encoder'].inverse_transform(y_test)558                    full_predictions_df['Predicted_Label'] = st.session_state['label_encoder'].inverse_transform(y_pred)559                560                csv = full_predictions_df.to_csv(index=False)561                st.download_button(562                    label="๐Ÿ“ฅ Download Predictions",563                    data=csv,564                    file_name="predictions_with_full_data.csv",565                    mime="text/csv"566                )567            568            with col2:569                # Export classification report570                report_csv = report_df.to_csv()571                st.download_button(572                    label="๐Ÿ“ฅ Download Classification Report",573                    data=report_csv,574                    file_name="classification_report.csv",575                    mime="text/csv"576                )577            578            # Additional Analysis Options579            st.subheader("๐Ÿ” Additional Analysis")580            581            with st.expander("Test with Different Parameters"):582                st.write("**Try different test sizes:**")583                584                test_sizes = [0.1, 0.2, 0.3, 0.4]585                results = []586                587                for ts in test_sizes:588                    X_train_temp, X_test_temp, y_train_temp, y_test_temp = train_test_split(589                        st.session_state['X'],590                        st.session_state['y'],591                        test_size=ts,592                        random_state=42,593                        stratify=st.session_state['y']594                    )595                    596                    # Create a temporary model with same configuration (Logistic Regression only)597                    temp_model = LogisticRegression(max_iter=1000, random_state=42)598                    599                    if multiclass_strategy == "One-vs-Rest (OvR)":600                        temp_model = OneVsRestClassifier(temp_model)601                    elif multiclass_strategy == "One-vs-One (OvO)":602                        temp_model = OneVsOneClassifier(temp_model)603                    604                    temp_model.fit(X_train_temp, y_train_temp)605                    y_pred_temp = temp_model.predict(X_test_temp)606                    acc = accuracy_score(y_test_temp, y_pred_temp)607                    608                    results.append({609                        'Test Size': f"{int(ts*100)}%",610                        'Train Samples': len(X_train_temp),611                        'Test Samples': len(X_test_temp),612                        'Accuracy': f"{acc:.3f}"613                    })614                615                results_df = pd.DataFrame(results)616                st.table(results_df)617        618        else:619            st.info("๐Ÿ“Œ Please train a model first in the 'Training' tab")620    621    with tab6:622        st.header("๐Ÿ”ฎ Make Predictions on New Data")623        624        if 'model' in st.session_state and st.session_state['model'] is not None:625            626            st.info("Upload a new dataset or enter data manually to make predictions using your trained model.")627            628            # Get feature columns from training629            if 'feature_columns' in locals():630                expected_features = feature_columns.copy()631            else:632                st.error("โŒ Feature information not available. Please train a model first.")633                st.stop()634            635            # Choose input method636            input_method = st.radio(637                "Select input method:",638                ["๐Ÿ“ Upload CSV File", "โœ๏ธ Enter Data Manually"],639                horizontal=True640            )641            642            if input_method == "๐Ÿ“ Upload CSV File":643                st.subheader("Upload New Data File")644                645                new_file = st.file_uploader(646                    "Choose a CSV file for prediction",647                    type="csv",648                    key="prediction_file",649                    help="Upload a CSV file with the same features as your training data (excluding target column)"650                )651                652                if new_file is not None:653                    # Load new data654                    new_data = pd.read_csv(new_file)655                    st.success(f"โœ… New data loaded! Shape: {new_data.shape}")656                    657                    # Display first few rows658                    st.write("**Preview of new data:**")659                    st.dataframe(new_data.head())660                    661                    # Check if columns match662                    expected_features = feature_columns.copy()663                    missing_cols = set(expected_features) - set(new_data.columns)664                    extra_cols = set(new_data.columns) - set(expected_features)665                    666                    if missing_cols:667                        st.error(f"โŒ Missing columns: {missing_cols}")668                        st.stop()669                    670                    if extra_cols:671                        st.warning(f"โš ๏ธ Extra columns will be ignored: {extra_cols}")672                        new_data = new_data[expected_features]673                    674                    if st.button("๐ŸŽฏ Make Predictions", key="batch_predict"):675                        with st.spinner("Processing and making predictions..."):676                            try:677                                # Preprocess the new data678                                processed_new_data = new_data[expected_features].copy()679                                680                                # Handle missing values681                                numerical_cols = processed_new_data.select_dtypes(include=[np.number]).columns682                                for col in numerical_cols:683                                    if processed_new_data[col].isnull().any():684                                        # Use median from training data if available685                                        if 'training_medians' in st.session_state and col in st.session_state['training_medians']:686                                            processed_new_data[col].fillna(st.session_state['training_medians'][col], inplace=True)687                                        else:688                                            processed_new_data[col].fillna(processed_new_data[col].median(), inplace=True)689                                690                                categorical_cols = processed_new_data.select_dtypes(include=['object']).columns691                                for col in categorical_cols:692                                    if processed_new_data[col].isnull().any():693                                        # Use mode from training data if available694                                        if 'training_modes' in st.session_state and col in st.session_state['training_modes']:695                                            processed_new_data[col].fillna(st.session_state['training_modes'][col], inplace=True)696                                        else:697                                            processed_new_data[col].fillna(processed_new_data[col].mode()[0], inplace=True)698                                699                                # Apply the same preprocessing as training data700                                X_new = processed_new_data.copy()701                                702                                # Identify numerical and categorical columns703                                numerical_columns = X_new.select_dtypes(include=[np.number]).columns.tolist()704                                categorical_columns = X_new.select_dtypes(include=['object']).columns.tolist()705                                706                                # Apply scaling if it was used during training707                                if 'scaler' in st.session_state and len(numerical_columns) > 0:708                                    X_new[numerical_columns] = st.session_state['scaler'].transform(X_new[numerical_columns])709                                710                                # Apply encoding if it was used during training711                                if 'encoder' in st.session_state and len(categorical_columns) > 0:712                                    encoded_features = st.session_state['encoder'].transform(X_new[categorical_columns])713                                    encoded_df = pd.DataFrame(714                                        encoded_features,715                                        columns=st.session_state['encoder'].get_feature_names_out(categorical_columns),716                                        index=X_new.index717                                    )718                                    X_new = pd.concat([X_new.drop(columns=categorical_columns), encoded_df], axis=1)719                                720                                # Ensure columns are in the same order as training data721                                X_new = X_new[st.session_state['X'].columns]722                                723                                # Make predictions724                                predictions = st.session_state['model'].predict(X_new)725                                726                                # Get prediction probabilities if available727                                try:728                                    pred_proba = st.session_state['model'].predict_proba(X_new)729                                    has_proba = True730                                except:731                                    has_proba = False732                                733                                # Create results dataframe734                                results_df = new_data.copy()735                                736                                # Add predictions737                                if 'label_encoder' in st.session_state:738                                    results_df['Predicted_Class'] = st.session_state['label_encoder'].inverse_transform(predictions)739                                    results_df['Predicted_Class_Code'] = predictions740                                    741                                    # Add probability columns if available742                                    if has_proba:743                                        for i, class_name in enumerate(st.session_state['label_encoder'].classes_):744                                            results_df[f'Probability_{class_name}'] = pred_proba[:, i]745                                else:746                                    results_df['Predicted_Class'] = predictions747                                    748                                    # Add probability columns if available749                                    if has_proba:750                                        for i in range(pred_proba.shape[1]):751                                            results_df[f'Probability_Class_{i}'] = pred_proba[:, i]752                                753                                st.success("โœ… Predictions completed!")754                                755                                # Display results756                                st.subheader("Prediction Results")757                                758                                # Show summary statistics759                                col1, col2 = st.columns(2)760                                with col1:761                                    st.metric("Total Predictions", len(results_df))762                                with col2:763                                    if 'label_encoder' in st.session_state:764                                        unique_preds = results_df['Predicted_Class'].value_counts()765                                    else:766                                        unique_preds = pd.Series(predictions).value_counts()767                                    st.metric("Unique Classes Predicted", len(unique_preds))768                                769                                # Display prediction distribution770                                st.write("**Prediction Distribution:**")771                                fig, ax = plt.subplots(figsize=(10, 6))772                                unique_preds.plot(kind='bar', ax=ax, color='#667eea')773                                ax.set_xlabel('Predicted Class')774                                ax.set_ylabel('Count')775                                ax.set_title('Distribution of Predictions')776                                plt.xticks(rotation=45, ha='right')777                                st.pyplot(fig)778                                779                                # Display detailed results780                                st.write("**Detailed Predictions:**")781                                st.dataframe(results_df)782                                783                                # Download predictions784                                csv = results_df.to_csv(index=False)785                                st.download_button(786                                    label="๐Ÿ“ฅ Download Predictions (CSV)",787                                    data=csv,788                                    file_name="new_predictions.csv",789                                    mime="text/csv"790                                )791                                792                            except Exception as e:793                                st.error(f"โŒ Error making predictions: {str(e)}")794                                st.write("Please ensure the new data has the same structure as the training data.")795            796            else:  # Manual entry797                st.subheader("Enter Data Manually")798                799                # Create input fields for each feature800                st.write("Enter values for each feature:")801                802                # Store input values803                input_data = {}804                805                # Get feature information from training data806                if 'X' in st.session_state:807                    # Original features (before encoding)808                    if 'feature_columns' in locals():809                        original_features = feature_columns810                    else:811                        st.error("Feature information not available")812                        st.stop()813                    814                    # Get sample data for reference815                    if st.session_state['data'] is not None and all(f in st.session_state['data'].columns for f in original_features):816                        original_data_sample = st.session_state['data'][original_features]817                    else:818                        # Create dummy sample if data not available819                        original_data_sample = pd.DataFrame(columns=original_features)820                    821                    # Create columns for better layout822                    col1, col2 = st.columns(2)823                    824                    for i, feature in enumerate(original_features):825                        # Determine input type based on original data type826                        if st.session_state['data'] is not None and feature in st.session_state['data'].columns:827                            # We have the original data, use it for reference828                            if feature in st.session_state['data'].select_dtypes(include=[np.number]).columns:829                                # Numerical input830                                with col1 if i % 2 == 0 else col2:831                                    # Get min/max from training data for reference832                                    min_val = float(st.session_state['data'][feature].min())833                                    max_val = float(st.session_state['data'][feature].max())834                                    mean_val = float(st.session_state['data'][feature].mean())835                                    836                                    input_data[feature] = st.number_input(837                                        f"{feature}",838                                        min_value=min_val - abs(min_val),839                                        max_value=max_val + abs(max_val),840                                        value=mean_val,841                                        help=f"Range in training data: [{min_val:.2f}, {max_val:.2f}]"842                                    )843                            else:844                                # Categorical input845                                with col1 if i % 2 == 0 else col2:846                                    unique_values = st.session_state['data'][feature].unique().tolist()847                                    input_data[feature] = st.selectbox(848                                        f"{feature}",849                                        options=unique_values,850                                        help=f"Select one of the {len(unique_values)} categories"851                                    )852                        else:853                            # No data available, provide generic input854                            with col1 if i % 2 == 0 else col2:855                                # Make a guess based on feature name856                                if any(word in feature.lower() for word in ['category', 'type', 'class', 'gender', 'status']):857                                    input_data[feature] = st.text_input(858                                        f"{feature}",859                                        help="Enter value (likely categorical)"860                                    )861                                else:862                                    input_data[feature] = st.number_input(863                                        f"{feature}",864                                        value=0.0,865                                        help="Enter numerical value"866                                    )867                    868                    if st.button("๐ŸŽฏ Make Prediction", key="single_predict"):869                        with st.spinner("Making prediction..."):870                            try:871                                # Create dataframe from input872                                input_df = pd.DataFrame([input_data])873                                874                                # Apply the same preprocessing as training data875                                X_single = input_df.copy()876                                877                                # Identify numerical and categorical columns878                                numerical_columns = X_single.select_dtypes(include=[np.number]).columns.tolist()879                                categorical_columns = X_single.select_dtypes(include=['object']).columns.tolist()880                                881                                # Apply scaling if it was used during training882                                if 'scaler' in st.session_state and len(numerical_columns) > 0:883                                    X_single[numerical_columns] = st.session_state['scaler'].transform(X_single[numerical_columns])884                                885                                # Apply encoding if it was used during training886                                if 'encoder' in st.session_state and len(categorical_columns) > 0:887                                    encoded_features = st.session_state['encoder'].transform(X_single[categorical_columns])888                                    encoded_df = pd.DataFrame(889                                        encoded_features,890                                        columns=st.session_state['encoder'].get_feature_names_out(categorical_columns),891                                        index=X_single.index892                                    )893                                    X_single = pd.concat([X_single.drop(columns=categorical_columns), encoded_df], axis=1)894                                895                                # Ensure columns are in the same order as training data896                                X_single = X_single[st.session_state['X'].columns]897                                898                                # Make prediction899                                prediction = st.session_state['model'].predict(X_single)[0]900                                901                                # Get prediction probability if available902                                try:903                                    pred_proba = st.session_state['model'].predict_proba(X_single)[0]904                                    has_proba = True905                                except:906                                    has_proba = False907                                908                                # Display result909                                st.success("โœ… Prediction completed!")910                                911                                # Show prediction912                                st.subheader("Prediction Result")913                                914                                if 'label_encoder' in st.session_state:915                                    predicted_class = st.session_state['label_encoder'].inverse_transform([prediction])[0]916                                    st.metric("Predicted Class", predicted_class)917                                else:918                                    st.metric("Predicted Class", prediction)919                                920                                # Show probabilities if available921                                if has_proba:922                                    st.write("**Class Probabilities:**")923                                    924                                    if 'label_encoder' in st.session_state:925                                        prob_df = pd.DataFrame({926                                            'Class': st.session_state['label_encoder'].classes_,927                                            'Probability': pred_proba928                                        })929                                    else:930                                        prob_df = pd.DataFrame({931                                            'Class': [f'Class {i}' for i in range(len(pred_proba))],932                                            'Probability': pred_proba933                                        })934                                    935                                    prob_df = prob_df.sort_values('Probability', ascending=False)936                                    937                                    # Display as bar chart938                                    fig, ax = plt.subplots(figsize=(10, 6))939                                    ax.barh(prob_df['Class'], prob_df['Probability'], color='#764ba2')940                                    ax.set_xlabel('Probability')941                                    ax.set_title('Prediction Confidence by Class')942                                    ax.set_xlim(0, 1)943                                    for i, (cls, prob) in enumerate(zip(prob_df['Class'], prob_df['Probability'])):944                                        ax.text(prob + 0.01, i, f'{prob:.3f}', va='center')945                                    st.pyplot(fig)946                                    947                                    # Also show as table948                                    st.write("**Probability Details:**")949                                    prob_df['Probability'] = prob_df['Probability'].apply(lambda x: f"{x:.4f}")950                                    st.table(prob_df)951                                952                                # Show input summary953                                with st.expander("View Input Data"):954                                    st.write("**Your input values:**")955                                    input_summary = pd.DataFrame([input_data])956                                    st.dataframe(input_summary)957                                    958                                    # Add download button for the single prediction959                                    result_data = input_data.copy()960                                    if 'label_encoder' in st.session_state:961                                        result_data['Predicted_Class'] = predicted_class962                                    else:963                                        result_data['Predicted_Class'] = prediction964                                    965                                    if has_proba:966                                        if 'label_encoder' in st.session_state:967                                            for i, cls in enumerate(st.session_state['label_encoder'].classes_):968                                                result_data[f'Prob_{cls}'] = pred_proba[i]969                                        else:970                                            for i in range(len(pred_proba)):971                                                result_data[f'Prob_Class_{i}'] = pred_proba[i]972                                    973                                    result_df = pd.DataFrame([result_data])974                                    csv = result_df.to_csv(index=False)975                                    st.download_button(976                                        label="๐Ÿ“ฅ Download Prediction Result",977                                        data=csv,978                                        file_name="single_prediction.csv",979                                        mime="text/csv"980                                    )981                                982                            except Exception as e:983                                st.error(f"โŒ Error making prediction: {str(e)}")984                                st.write("Please check your input values and try again.")985                else:986                    st.warning("โš ๏ธ Training data information not available. Please train a model first.")987            988            # Add model information989            st.sidebar.divider()990            st.sidebar.subheader("๐Ÿ“Š Model Information")991            if 'model' in st.session_state:992                st.sidebar.write(f"**Classifier:** Logistic Regression")993                994                if 'multiclass_strategy' in locals():995                    st.sidebar.write(f"**Strategy:** {multiclass_strategy}")996                997                if 'label_encoder' in st.session_state:998                    st.sidebar.write(f"**Classes:** {', '.join(map(str, st.session_state['label_encoder'].classes_))}")999                1000                if 'X' in st.session_state:1001                    st.sidebar.write(f"**Features:** {len(st.session_state['X'].columns)}")1002        1003        else:1004            st.warning("โš ๏ธ No trained model found. Please train a model first in the 'Training' tab.")1005            st.info("""1006            **Steps to train a model:**1007            1. Upload your training data1008            2. Select target and feature columns1009            3. Preprocess the data1010            4. Train the model (using Logistic Regression)1011            5. Come back here to make predictions on new data1012            """)1013 1014else:1015    # Landing page when no data is uploaded1016    st.markdown("""1017    ## Welcome to the Multi-class Classification App! ๐Ÿ‘‹1018    1019    This application provides a comprehensive solution for multi-class classification on any CSV dataset.1020    1021    ### ๐Ÿš€ Getting Started1022    1023    1. **Upload your CSV file** using the sidebar1024    2. **Select your target column** - the variable you want to predict1025    3. **Choose feature columns** - the variables to use for prediction1026    4. **Configure preprocessing options** - standardization and encoding1027    5. **Select model and strategy** - choose classifier and multi-class approach1028    6. **Train your model** and analyze results1029    1030    ### โœจ Features1031    1032    - **Automatic Feature Detection**: Identifies numerical and categorical columns1033    - **Data Preprocessing**: Handles missing values, scaling, and encoding1034    - **Multiple Classifiers**: Logistic Regression, SVM, Random Forest1035    - **Multi-class Strategies**: One-vs-Rest (OvR) and One-vs-One (OvO)1036    - **Comprehensive Analysis**: Confusion matrix, classification report, feature importance1037    - **Interactive Visualizations**: Explore data distributions and model performance1038    - **Export Results**: Download predictions and reports1039    1040    ### ๐Ÿ“Š Supported Classification Types1041    1042    - Binary Classification1043    - Multi-class Classification1044    - Imbalanced Datasets (with stratified splitting)1045    1046    ### ๐ŸŽฏ Use Cases1047    1048    - Customer segmentation1049    - Disease diagnosis1050    - Product categorization1051    - Risk assessment1052    - Quality control1053    - And many more...1054    1055    ---1056    1057    **Ready to start?** Upload your dataset using the sidebar! ๐Ÿ“1058    """)1059    1060    # Add sample dataset info1061    with st.expander("๐Ÿ“ Sample Dataset Format"):1062        st.markdown("""1063        Your CSV file should have:1064        - **Features**: Columns containing predictor variables (numerical or categorical)1065        - **Target**: A column with the classes/categories you want to predict1066        1067        Example structure:1068        """)1069        1070        sample_data = pd.DataFrame({1071            'Feature1': [1.2, 2.3, 3.1, 4.5, 5.0],1072            'Feature2': ['A', 'B', 'A', 'C', 'B'],1073            'Feature3': [10, 20, 15, 30, 25],1074            'Target': ['Class1', 'Class2', 'Class1', 'Class3', 'Class2']1075        })1076        st.dataframe(sample_data)1077 1078# Footer1079st.markdown("---")1080st.markdown("""1081<div style='text-align: center'>1082    <p>Built with โค๏ธ using Streamlit | Multi-class Classification Web App</p>1083</div>1084""", unsafe_allow_html=True)1085