Skip to content

Commit 3164d11

Browse files
enhance(data): Minor optimization in how DFs are validated for potential duplicate columns in label_setup. Explicitly del df_labels immediately after splitting into df_labels_train and df_labels_valid.
1 parent be78bf3 commit 3164d11

1 file changed

Lines changed: 10 additions & 10 deletions

File tree

src/eir/data_load/label_setup.py

Lines changed: 10 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -70,10 +70,6 @@ def set_up_train_and_valid_tabular_data(
7070
impute_missing: bool = False,
7171
do_transform_labels: bool = True,
7272
) -> Labels:
73-
"""
74-
Splits and does split based processing (e.g. scaling validation set with training
75-
set for regression) on the labels.
76-
"""
7773
if len(tabular_file_info.con_columns) + len(tabular_file_info.cat_columns) < 1:
7874
raise ValueError(f"No label columns specified in {tabular_file_info}.")
7975

@@ -92,6 +88,8 @@ def set_up_train_and_valid_tabular_data(
9288
train_ids=list(train_ids),
9389
valid_ids=list(valid_ids),
9490
)
91+
del df_labels
92+
9593
pre_check_label_df(df=df_labels_train, name="Training DataFrame")
9694
pre_check_label_df(df=df_labels_valid, name="Validation DataFrame")
9795
check_train_valid_df_sync(
@@ -270,12 +268,14 @@ def get_label_parsing_wrapper(
270268

271269

272270
def _validate_df(df: pl.DataFrame) -> None:
273-
duplicate_counts = (
274-
df.group_by("ID").agg(pl.count().alias("count")).filter(pl.col("count") > 1)
275-
)
276-
277-
if duplicate_counts.height > 0:
278-
duplicated_ids = duplicate_counts.select("ID").limit(10).to_series().to_list()
271+
if df.select(pl.col("ID").is_duplicated()).sum().item():
272+
duplicated_ids = (
273+
df.filter(pl.col("ID").is_duplicated())
274+
.get_column("ID")
275+
.unique()
276+
.head(10)
277+
.to_list()
278+
)
279279
duplicated_indices_str = ", ".join(map(str, duplicated_ids))
280280
raise ValueError(
281281
f"Found duplicated indices in the dataframe. "

0 commit comments

Comments
 (0)