Created
August 22, 2024 17:27
-
-
Save eugen-hoppe/c053ed449fa2e3ec6a3ce9dd5bcb5c51 to your computer and use it in GitHub Desktop.
2024-08-22
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| @dataclass | |
| class Cat: | |
| ignore_steps: bool = False | |
| step: int = 0 | |
| labels: pd.DataFrame | None = None | |
| categories: pd.DataFrame | None = None | |
| sorted_labels: pd.DataFrame | None = None | |
| merged_final: pd.DataFrame | None = None | |
| category_prefixes: list[str] = field( | |
| default_factory=lambda: ['c1', 'c2'] | |
| ) | |
| def validate_step(self, step: int) -> None: | |
| if self.ignore_steps: | |
| return None | |
| if self.step != step - 1: | |
| raise ValueError( | |
| f'Step {step} cannot proceed because the current step is {self.step}.' | |
| ) | |
| self.step = step | |
| def create_labels(self, df_: pd.DataFrame) -> pd.DataFrame: | |
| self.validate_step(1) | |
| exploded_df = df_['category_list'].str.split('|').explode() | |
| label_dict = {'company_id': exploded_df.index, 'label': exploded_df} | |
| labels = pd.DataFrame(label_dict).reset_index(drop=True) | |
| self.labels = labels[labels['label'].notna()] | |
| return self.labels | |
| def create_categories(self, na_cat: str = 'NA') -> pd.DataFrame: | |
| self.validate_step(2) | |
| categories = pd.DataFrame(self.labels['label'].unique(), columns=['category']) | |
| categories.sort_values(by='category', ignore_index=True, inplace=True) | |
| categories_na = pd.DataFrame([na_cat], columns=['category']) | |
| categories = pd.concat([categories_na, categories], ignore_index=True) | |
| categories['category_id'] = categories.index | |
| self.categories = categories | |
| return self.categories | |
| def merge_with_categories_to_assign_category_ids(self) -> None: | |
| self.validate_step(3) | |
| labels = pd.merge( | |
| self.labels, | |
| self.categories, | |
| left_on='label', | |
| right_on='category', | |
| how='inner' | |
| ) | |
| self.labels = labels[['category_id', 'company_id']] | |
| def count_id_occurrences_and_merge_with_categories(self) -> None: | |
| self.validate_step(4) | |
| category_counts = self.labels['category_id'].value_counts().reset_index() | |
| category_counts.columns = ['category_id', 'assignments'] | |
| categories = pd.merge( | |
| self.categories, | |
| category_counts, | |
| on='category_id', | |
| how='left' | |
| ) | |
| categories.rename(columns={'category': 'name'}, inplace=True) | |
| self.categories = categories | |
| def labels_with_assignments(self) -> pd.DataFrame: | |
| self.validate_step(5) | |
| labels_with_assignments = pd.merge( | |
| self.labels, | |
| self.categories[['category_id', 'assignments']], | |
| on='category_id', | |
| how='left' | |
| ) | |
| self.sorted_labels = labels_with_assignments.sort_values( | |
| by=['company_id', 'assignments'], ascending=[True, False] | |
| ) | |
| return self.sorted_labels | |
| def get_top_categories(self) -> list[pd.DataFrame]: | |
| self.validate_step(6) | |
| top_categories = self.sorted_labels.groupby('company_id').head( | |
| len(self.category_prefixes) | |
| ) | |
| category_dfs = [] | |
| for i, prefix in enumerate(self.category_prefixes): | |
| nth_category = top_categories.groupby('company_id').nth(i).reset_index() | |
| nth_category.rename( | |
| columns={'category_id': f'{prefix}_category_id'}, inplace=True | |
| ) | |
| category_dfs.append(nth_category) | |
| return category_dfs | |
| def merge_categories_with_original(self, df_: pd.DataFrame) -> pd.DataFrame: | |
| self.validate_step(7) | |
| origin_df = df_.reset_index().rename(columns={'index': 'company_id'}) | |
| top_category_dfs = self.get_top_categories() | |
| merged = origin_df | |
| for i, df in enumerate(top_category_dfs): | |
| merged = pd.merge( | |
| merged, | |
| df[['company_id', f'{self.category_prefixes[i]}_category_id']], | |
| on='company_id', | |
| how='left' | |
| ) | |
| category_cols = [f'{prefix}_category_id' for prefix in self.category_prefixes] | |
| merged[category_cols] = merged[category_cols].fillna(0).astype(int) | |
| self.merged_final = merged | |
| return self.merged_final | |
| def map_category_names(self) -> pd.DataFrame: | |
| self.validate_step(8) | |
| category_name_map = self.categories.set_index('category_id')['name'].to_dict() | |
| for prefix in self.category_prefixes: | |
| self.merged_final[f'{prefix}_cat'] = ( | |
| self.merged_final[f'{prefix}_category_id'].map(category_name_map) | |
| ) | |
| return self.merged_final | |
| def drop_unused_columns(self, columns_to_drop: list[str] = []) -> pd.DataFrame: | |
| self.validate_step(9) | |
| if not len(columns_to_drop): | |
| columns_to_drop = [ | |
| 'company_id', | |
| 'category_list' | |
| ] | |
| columns_to_drop.extend( | |
| [f'{prefix}_category_id' for prefix in self.category_prefixes] | |
| ) | |
| self.merged_final.drop(columns=columns_to_drop, inplace=True) | |
| return self.merged_final | |
| def generate_columns(self, df: pd.DataFrame) -> pd.DataFrame: | |
| '''Generates Ranked Category Columns | |
| The three categories primary, secondary, and tertiary (Default) are created | |
| from the category_list column using a multi-stage process in the Cat class. | |
| First, the category_list column is split to extract individual categories and | |
| store them in a DataFrame called labels. These categories are then counted and | |
| sorted, whereby the frequency of occurrences per company_id is decisive. The | |
| method get_top_categories groups the labels by company_id and selects the most | |
| frequent categories for each company_id. The most frequent category is labeled | |
| as primary, the second most frequent as secondary, and the third most frequent | |
| as tertiary based on their order of occurrence. | |
| ''' | |
| self.ignore_steps = True | |
| self.create_labels(df) | |
| self.create_categories() | |
| self.merge_with_categories_to_assign_category_ids() | |
| self.count_id_occurrences_and_merge_with_categories() | |
| self.labels_with_assignments() | |
| self.merge_categories_with_original(df) | |
| self.map_category_names() | |
| final_df = self.drop_unused_columns() | |
| return final_df | |
| def columns(self) -> list[str]: | |
| return [cat + '_cat' for cat in self.category_prefixes] | |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment