Skip to content

Instantly share code, notes, and snippets.

@eugen-hoppe
Created August 22, 2024 17:27
Show Gist options
  • Select an option

  • Save eugen-hoppe/c053ed449fa2e3ec6a3ce9dd5bcb5c51 to your computer and use it in GitHub Desktop.

Select an option

Save eugen-hoppe/c053ed449fa2e3ec6a3ce9dd5bcb5c51 to your computer and use it in GitHub Desktop.
2024-08-22
@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