diff --git a/python/src/spark_rapids_ml/umap.py b/python/src/spark_rapids_ml/umap.py index 03cce339..e311cd46 100644 --- a/python/src/spark_rapids_ml/umap.py +++ b/python/src/spark_rapids_ml/umap.py @@ -1256,7 +1256,7 @@ def _train_udf(pdf_iter: Iterable[pd.DataFrame]) -> Iterable[pd.DataFrame]: indices = csr_chunk.indices indptr = csr_chunk.indptr data = csr_chunk.data - if cuda_managed_mem_enabled: + if cuda_managed_mem_enabled or cuda_system_mem_enabled: yield pd.DataFrame( data=[ { @@ -1281,7 +1281,7 @@ def _train_udf(pdf_iter: Iterable[pd.DataFrame]) -> Iterable[pd.DataFrame]: ] ) else: - if cuda_managed_mem_enabled: + if cuda_managed_mem_enabled or cuda_system_mem_enabled: yield pd.DataFrame( { "embedding_": list(embedding[start:end].get()),