diff --git a/chebILP/ilp_problem_builder.py b/chebILP/ilp_problem_builder.py index e89fe4c..3cff2b8 100644 --- a/chebILP/ilp_problem_builder.py +++ b/chebILP/ilp_problem_builder.py @@ -263,16 +263,16 @@ def gather_samples_for_chebi_cls(self, target_id: str, min_pos_samples=25, max_p sibling_neg_ids = set(sibling_neg_ids) samples_by_split = dict() - pos_train_samples = df_pos[df_pos.index.astype(str).isin(self.splits[self.splits["split"] == "train"])] + pos_train_samples = df_pos[df_pos.index.astype(str).isin(self.splits.loc[self.splits["split"] == "train", "id"])] samples_by_split[("pos", "train")] = pos_train_samples.sample(min(max_pos_samples, len(pos_train_samples)), random_state=42) # if there are more positives than max_pos_samples, sample randomly - neg_train_samples = df_neg[df_neg.index.astype(str).isin(self.splits[self.splits["split"] == "train"])] + neg_train_samples = df_neg[df_neg.index.astype(str).isin(self.splits.loc[self.splits["split"] == "train", "id"])] samples_by_split[("neg", "train")] = self.build_negative_mix(neg_train_samples, sibling_neg_ids, max_neg_samples) - samples_by_split[("pos", "validation")] = df_pos[df_pos.index.astype(str).isin(self.splits[self.splits["split"] == "validation"]) & df_pos.index.astype(str).isin(pos_ids)] - neg_val_samples = df_neg[df_neg.index.astype(str).isin(self.splits[self.splits["split"] == "validation"])] + samples_by_split[("pos", "validation")] = df_pos[df_pos.index.astype(str).isin(self.splits.loc[self.splits["split"] == "validation", "id"]) & df_pos.index.astype(str).isin(pos_ids)] + neg_val_samples = df_neg[df_neg.index.astype(str).isin(self.splits.loc[self.splits["split"] == "validation", "id"])] samples_by_split[("neg", "validation")] = self.build_negative_mix(neg_val_samples, sibling_neg_ids, max_neg_samples) - samples_by_split[("pos", "test")] = df_pos[df_pos.index.astype(str).isin(self.splits[self.splits["split"] == "test"]) & df_pos.index.astype(str).isin(pos_ids)] - neg_test_samples = df_neg[df_neg.index.astype(str).isin(self.splits[self.splits["split"] == "test"])] + samples_by_split[("pos", "test")] = df_pos[df_pos.index.astype(str).isin(self.splits.loc[self.splits["split"] == "test", "id"]) & df_pos.index.astype(str).isin(pos_ids)] + neg_test_samples = df_neg[df_neg.index.astype(str).isin(self.splits.loc[self.splits["split"] == "test", "id"])] samples_by_split[("neg", "test")] = self.build_negative_mix(neg_test_samples, sibling_neg_ids, max_neg_samples) for (posneg, split), df in samples_by_split.items(): diff --git a/chebILP/molecule_processing/data_preparation.py b/chebILP/molecule_processing/data_preparation.py index 74dc963..e21ea01 100644 --- a/chebILP/molecule_processing/data_preparation.py +++ b/chebILP/molecule_processing/data_preparation.py @@ -190,7 +190,7 @@ def load_splits_from_csv(self) -> pd.DataFrame: if not os.path.exists(splits_path): raise FileNotFoundError(f"Splits file not found: {splits_path}. " f"Run `python -m chebILP prepare_dataset` to create it.") - splits_df = pd.read_csv(splits_path) + splits_df = pd.read_csv(splits_path, dtype={"id": str}) if "id" not in splits_df.columns or "split" not in splits_df.columns: raise ValueError(f"Splits CSV must contain 'id' and 'split' columns: {splits_path}") if not all(s in splits_df["split"].unique() for s in ["train", "validation", "test"]): diff --git a/pyproject.toml b/pyproject.toml index 9c87406..4f4cea4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -3,7 +3,7 @@ name = "chebilp" version = "1.1" description = "An Inductive Logic Programming framework for classifying chemical compounds into ChEBI classes." readme = "README.md" -requires-python = ">=3.10" +requires-python = ">=3.11" dependencies = [ "chebi-utils>=0.2.1", "clingo>=5.8.0",