diff --git a/tabfm/src/classifier_and_regressor.py b/tabfm/src/classifier_and_regressor.py index 4a0ebee..4cc74d1 100644 --- a/tabfm/src/classifier_and_regressor.py +++ b/tabfm/src/classifier_and_regressor.py @@ -208,6 +208,7 @@ def transform(self, X: Any) -> np.ndarray: for i in range(n_features): col = X[:, i] + missing_mask = pd.isna(col) cats = self.categories_[i] cat_to_idx = {c: idx for idx, c in enumerate(cats)} mapped = pd.Series(col).map(cat_to_idx) @@ -216,7 +217,9 @@ def transform(self, X: Any) -> np.ndarray: if mapped.dtype.name == "category": mapped = mapped.astype(object) - # Fill NaNs/Unknowns with unknown_value + mapped.loc[missing_mask] = self.encoded_missing_value + + # Fill unknowns with unknown_value X_out[:, i] = mapped.fillna(self.unknown_value).values.astype(self.dtype) return X_out @@ -334,7 +337,10 @@ def fit(self, X: Any, y: Any = None) -> "DatetimeTransformer": series = pd.to_datetime( X.iloc[:, pos], utc=True, errors="coerce", format="mixed" ) - self._fillna_map[pos] = series.mean() + fillna_value = series.mean() + if pd.isna(fillna_value): + fillna_value = pd.Timestamp(0, tz="UTC") + self._fillna_map[pos] = fillna_value return self def transform(self, X: Any) -> np.ndarray: diff --git a/tabfm/src/classifier_and_regressor_test.py b/tabfm/src/classifier_and_regressor_test.py index 546fbee..6937b7c 100644 --- a/tabfm/src/classifier_and_regressor_test.py +++ b/tabfm/src/classifier_and_regressor_test.py @@ -27,6 +27,8 @@ except ImportError: HAS_JAX = False from tabfm.src.classifier_and_regressor import _looks_like_datetime +from tabfm.src.classifier_and_regressor import CategoricalOrdinalEncoder +from tabfm.src.classifier_and_regressor import DatetimeTransformer from tabfm.src.classifier_and_regressor import EnsembleGenerator from tabfm.src.classifier_and_regressor import TabFMClassifier from tabfm.src.classifier_and_regressor import TabFMRegressor @@ -1165,6 +1167,32 @@ def test_mostly_non_date_column_not_detected(self): self.assertFalse(_looks_like_datetime(pd.Series(vals, dtype="string"))) +class CategoricalOrdinalEncoderTest(absltest.TestCase): + + def test_missing_and_unknown_values_are_encoded_separately(self): + encoder = CategoricalOrdinalEncoder( + unknown_value=-1, encoded_missing_value=-99 + ) + encoder.fit(np.array([["known"], ["other"]], dtype=object)) + + encoded = encoder.transform( + np.array([[np.nan], ["unseen"], ["known"]], dtype=object) + ) + + np.testing.assert_array_equal(encoded[:, 0], [-99, -1, 0]) + + +class DatetimeTransformerTest(absltest.TestCase): + + def test_all_missing_column_uses_epoch_fill_value(self): + X = pd.DataFrame({"ts": pd.to_datetime([None, None, None])}) + + transformed = DatetimeTransformer().fit_transform(X) + + self.assertEqual(transformed.shape, (3, 5)) + self.assertTrue(np.isfinite(transformed).all()) + + class ColumnNameRobustnessTest(absltest.TestCase): def test_duplicate_column_names_raise_a_clear_error(self):