From e73db7e0d7f72dce0a721bc125b7646f4d876952 Mon Sep 17 00:00:00 2001 From: Tim Date: Tue, 25 Feb 2025 18:20:14 +0800 Subject: [PATCH] ignore new numeric type (#638) --- .../coder/data_science/feature/eval_tests/feature_test.txt | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) 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 1aca7903..947fa617 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 @@ -61,8 +61,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([str(dt) for dt in X.dtypes.unique()]) - X_loaded_dtypes_unique_sorted = sorted([str(dt) for dt in X_loaded.dtypes.unique()]) + def normalize_dtype(dtype): + return "numeric" if np.issubdtype(dtype, np.number) else str(dtype) + + X_dtypes_unique_sorted = sorted(set(normalize_dtype(dt) for dt in X.dtypes.unique())) + X_loaded_dtypes_unique_sorted = sorted(set(normalize_dtype(dt) for dt in X_loaded.dtypes.unique())) X_dtypes_unique_sorted_new = [ dt for dt in X_dtypes_unique_sorted if dt not in X_loaded_dtypes_unique_sorted and dt != "object"