From a0e274889b678eea3a2e3d4f3d5f6bafbe80dd2c Mon Sep 17 00:00:00 2001 From: Tim Date: Sat, 8 Feb 2025 15:21:11 +0800 Subject: [PATCH] support multiple types in feature engineering (#562) --- .../data_science/ensemble/eval_tests/ensemble_test.txt | 3 +++ .../coder/data_science/feature/eval_tests/feature_test.txt | 6 +++--- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/rdagent/components/coder/data_science/ensemble/eval_tests/ensemble_test.txt b/rdagent/components/coder/data_science/ensemble/eval_tests/ensemble_test.txt index 7875558f..396f2c46 100644 --- a/rdagent/components/coder/data_science/ensemble/eval_tests/ensemble_test.txt +++ b/rdagent/components/coder/data_science/ensemble/eval_tests/ensemble_test.txt @@ -20,6 +20,9 @@ X, y, test_X, test_ids = load_data() X, y, test_X = feat_eng(X, y, test_X) train_X, val_X, train_y, val_y = train_test_split(X, y, test_size=0.2, random_state=42) +# Print the types of train_y and val_y +print(f"train_y type: {type(train_y)}, val_y type: {type(val_y)}") + test_preds_dict = {} val_preds_dict = {} {% for mn in model_names %} diff --git a/rdagent/components/coder/data_science/feature/eval_tests/feature_test.txt b/rdagent/components/coder/data_science/feature/eval_tests/feature_test.txt index cc5a5f15..726c258a 100644 --- a/rdagent/components/coder/data_science/feature/eval_tests/feature_test.txt +++ b/rdagent/components/coder/data_science/feature/eval_tests/feature_test.txt @@ -57,11 +57,11 @@ if isinstance(X, pd.DataFrame) and isinstance(X_test, pd.DataFrame): assert get_column_list(X) == get_column_list(X_test), "Mismatch in column names of training and test data." if isinstance(X, pd.DataFrame): - X_dtypes_unique_sorted = sorted(X.dtypes.unique().tolist()) - X_loaded_dtypes_unique_sorted = sorted(X_loaded.dtypes.unique().tolist()) + X_dtypes_unique_sorted = sorted([str(dt) for dt in X.dtypes.unique()]) + X_loaded_dtypes_unique_sorted = sorted([str(dt) for dt in X_loaded.dtypes.unique()]) assert ( len(X_loaded_dtypes_unique_sorted) == 1 - and (X_loaded_dtypes_unique_sorted[0] == np.float64 or X_loaded_dtypes_unique_sorted[0] == np.float32) + and X_loaded_dtypes_unique_sorted[0] in {np.float64, np.float32} ) or ( X_dtypes_unique_sorted == X_loaded_dtypes_unique_sorted ), f"feature engineering has produced new data types which is not allowed, data loader data types are {X_loaded_dtypes_unique_sorted} and feature engineering data types are {X_dtypes_unique_sorted}"