Skip to content

Commit 660cdf8

Browse files
authored
Merge branch 'main' into docs/issue-1093
2 parents 2971256 + 7d14682 commit 660cdf8

2 files changed

Lines changed: 64 additions & 34 deletions

File tree

ehrapy/preprocessing/_imputation.py

Lines changed: 29 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -371,6 +371,31 @@ def knn_impute(
371371
return edata if copy else None
372372

373373

374+
@singledispatch
375+
def _knn_impute_function(arr, var_indices: list[int], numerical_indices: list[int], imputer) -> np.ndarray:
376+
raise _raise_array_type_not_implemented(_knn_impute_function, type(arr))
377+
378+
379+
@_knn_impute_function.register(np.ndarray)
380+
@_apply_over_time_axis
381+
def _(arr: np.ndarray, var_indices: list[int], numerical_indices: list[int], imputer) -> np.ndarray:
382+
input_dtype = arr.dtype if np.issubdtype(arr.dtype, np.floating) else np.float64
383+
numerical_arr = arr[:, numerical_indices].astype(input_dtype, copy=True)
384+
385+
# complete columns to be used as anchors
386+
complete_numerical_columns = np.array(numerical_indices)[~np.isnan(numerical_arr).any(axis=0)].tolist()
387+
388+
imputer_data_indices = var_indices + [
389+
column for column in complete_numerical_columns if column not in var_indices
390+
] # columns to impute
391+
imputer_x = arr[:, imputer_data_indices].astype(input_dtype, copy=True)
392+
X_imputed = imputer.fit_transform(imputer_x)
393+
394+
result = arr.copy()
395+
result[:, imputer_data_indices] = X_imputed
396+
return result
397+
398+
374399
def _knn_impute(
375400
edata: EHRData,
376401
var_names: Iterable[str] | None,
@@ -400,43 +425,13 @@ def _knn_impute(
400425
"var_names parameter or perform an encoding of your data."
401426
)
402427
mtx = edata.X if layer is None else edata.layers[layer]
403-
var_indices_original = var_indices
404-
is_3d = False
405-
input_dtype = mtx.dtype if np.issubdtype(mtx.dtype, np.floating) else np.float64
406-
407-
# if input data is 3D, flatten along axis 0 before passing it to the imputer: each timepoint becomes a row
408-
if mtx.ndim == 3:
409-
is_3d = True
410-
n_obs, n_vars, n_t = mtx.shape
411-
mtx = (
412-
mtx[:, var_indices, :]
413-
.astype(input_dtype, copy=True)
414-
.transpose(0, 2, 1)
415-
.reshape(n_obs * n_t, len(var_indices))
416-
)
417-
numerical_indices = list(range(len(var_indices)))
418-
var_indices = numerical_indices
419428

420-
# complete columns to be used as anchors
421-
complete_numerical_columns = np.array(numerical_indices)[~np.isnan(mtx[:, numerical_indices]).any(axis=0)].tolist()
422-
423-
imputer_data_indices = var_indices + [
424-
column for column in complete_numerical_columns if column not in var_indices
425-
] # columns to impute
426-
imputer_x = mtx[:, imputer_data_indices].astype(input_dtype, copy=True)
427-
X_imputed = imputer.fit_transform(imputer_x)
429+
X_imputed = _knn_impute_function(mtx, var_indices, numerical_indices, imputer)
428430

429-
if is_3d:
430-
# slice back to only requested columns and transpose back to n_obs, n_var, n_t
431-
X_imputed = (
432-
X_imputed[:, : len(var_indices_original)].reshape(n_obs, n_t, len(var_indices_original)).transpose(0, 2, 1)
433-
)
434-
edata.layers[layer][:, var_indices_original, :] = X_imputed
431+
if layer is None:
432+
edata.X[:, var_indices] = X_imputed[:, var_indices]
435433
else:
436-
if layer is None:
437-
edata.X[:, imputer_data_indices] = X_imputed
438-
else:
439-
edata.layers[layer][:, imputer_data_indices] = X_imputed
434+
edata.layers[layer][:, var_indices] = X_imputed[:, var_indices]
440435

441436

442437
@singledispatch

tests/preprocessing/test_imputation.py

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -352,6 +352,41 @@ def test_knn_impute_numerical_data(impute_num_edata):
352352
_base_check_imputation(impute_num_edata, edata_imputed)
353353

354354

355+
@pytest.mark.parametrize("array_type", ARRAY_TYPES_NUMERIC)
356+
@pytest.mark.parametrize("backend", ["scikit-learn", "faiss"])
357+
def test_knn_impute_array_types(impute_num_edata, array_type, backend):
358+
impute_num_edata.X = array_type(impute_num_edata.X)
359+
360+
if isinstance(impute_num_edata.X, da.Array | sparse.csr_array | sparse.csc_array):
361+
with pytest.raises(NotImplementedError):
362+
knn_impute(impute_num_edata, backend=backend, copy=True)
363+
else:
364+
edata_imputed = knn_impute(impute_num_edata, backend=backend, copy=True)
365+
366+
_base_check_imputation(impute_num_edata, edata_imputed)
367+
368+
369+
@pytest.mark.parametrize("edata_mini_3D_missing_values", [True], indirect=True)
370+
@pytest.mark.parametrize("array_type", ARRAY_TYPES_NUMERIC_3D_ABLE)
371+
@pytest.mark.parametrize("backend", ["scikit-learn", "faiss"])
372+
def test_knn_impute_3d_array_types(edata_mini_3D_missing_values, array_type, backend):
373+
edata = edata_mini_3D_missing_values.copy()
374+
edata.layers[DEFAULT_TEM_LAYER_NAME] = array_type(edata.layers[DEFAULT_TEM_LAYER_NAME])
375+
376+
if isinstance(edata.layers[DEFAULT_TEM_LAYER_NAME], da.Array):
377+
with pytest.raises(NotImplementedError):
378+
knn_impute(edata, layer=DEFAULT_TEM_LAYER_NAME, backend=backend, copy=True)
379+
else:
380+
edata_imputed = knn_impute(edata, layer=DEFAULT_TEM_LAYER_NAME, backend=backend, copy=True)
381+
382+
_base_check_imputation(
383+
edata_mini_3D_missing_values,
384+
edata_imputed,
385+
before_imputation_layer=DEFAULT_TEM_LAYER_NAME,
386+
after_imputation_layer=DEFAULT_TEM_LAYER_NAME,
387+
)
388+
389+
355390
@pytest.mark.parametrize("edata_mini_3D_missing_values", [True], indirect=True)
356391
def test_missforest_impute_3D_edata(edata_mini_3D_missing_values):
357392
edata = edata_mini_3D_missing_values.copy()

0 commit comments

Comments
 (0)