diff --git a/scripts/encoding_func.py b/scripts/encoding_func.py new file mode 100644 index 0000000..5506145 --- /dev/null +++ b/scripts/encoding_func.py @@ -0,0 +1,32 @@ +from sklearn.preprocessing import OneHotEncoder +import numpy as np +import pandas as pd + + +def one_hot_encode_column(df, column_name): + """ + One-hot encodes a specified column in the DataFrame. + + Parameters: + df (pd.DataFrame): The DataFrame containing the column to encode. + column_name (str): The name of the column to one-hot encode. + + Returns: + pd.DataFrame: The DataFrame with the one-hot encoded columns added. + """ + # Replace NaN values with 'unknown' for consistency + df[[column_name]] = df[[column_name]].replace('', np.nan) + df[[column_name]]= df[[column_name]].fillna('unknown') + + enc = OneHotEncoder(handle_unknown='ignore') + enc.fit(df[[column_name]]) + encoded_array = enc.transform(df[[column_name]]).toarray() + encoded_df = pd.DataFrame(encoded_array, columns=enc.get_feature_names_out([column_name])) + + return pd.concat([df, encoded_df], axis=1) + + +# Example of usage +filtered_df = one_hot_encode_column(filtered_df, 'pnns_groups_1') +filtered_df.head() +# filtered_df = filtered_df.drop(columns=['pnns_groups_1'])