I would like to find a way to use sklearn.preprocessing inside tf.data.Dataset.map() in TF2.
Let's say I have a dataset generated from
import tensorflow as tf
ds = tf.data.Dataset.from_tensor_slices((tf.random.uniform((3, 3))))
ds = ds.batch(1)
x = tf.concat(list(ds.as_numpy_iterator()), axis=0)
print(x)
# tf.Tensor(
# [[0.51869464 0.9198195 0.87195873]
# [0.5842893 0.5363847 0.93642473]
# [0.0109899 0.7908174 0.25996208]], shape=(3, 3), dtype=float32)
Then calculate the QuantileTransformer
from sklearn.preprocessing import QuantileTransformer
qt = QuantileTransformer(n_quantiles=2, random_state=0)
qt.fit_transform(x)
print(qt.quantiles_)
# [[0.0109899 0.5363847 0.25996208]
# [0.58428931 0.91981947 0.93642473]]
However, I was not able to use QuantileTransformer in tf.data.Dataset.map. For example,
ds.map(lambda x: qt.transform(x))
gives error
TypeError: in user code:
<ipython-input-106-867a262b9b69>:13 None *
lambda x: qt.transform(x)
/lib/python3.8/site-packages/sklearn/preprocessing/_data.py:2769 transform *
X = self._check_inputs(X, in_fit=False, copy=self.copy)
/lib/python3.8/site-packages/sklearn/preprocessing/_data.py:2699 _check_inputs *
X = self._validate_data(X, reset=in_fit,
/lib/python3.8/site-packages/sklearn/base.py:420 _validate_data *
X = check_array(X, **check_params)
/lib/python3.8/site-packages/sklearn/utils/validation.py:981 inner_f *
return f(*args, **kwargs)
/lib/python3.8/site-packages/sklearn/utils/validation.py:616 check_array *
array = np.asarray(array, order=order, dtype=dtype)
/lib/python3.8/site-packages/numpy/core/_asarray.py:83 asarray **
return array(a, dtype, copy=False, order=order)
TypeError: __array__() takes 1 positional argument but 2 were given```