diff --git a/src/autoencodix/base/_base_pipeline.py b/src/autoencodix/base/_base_pipeline.py index fa0d446f..83c49ab2 100644 --- a/src/autoencodix/base/_base_pipeline.py +++ b/src/autoencodix/base/_base_pipeline.py @@ -106,7 +106,8 @@ def __init__( TypeError: If inputs have incorrect types. """ if not hasattr(self, "_default_config"): - raise ValueError(""" + raise ValueError( + """ The _default_config attribute has not been specified in your pipeline class. Example: @@ -116,7 +117,8 @@ def __init__( _default_config in its corresponding pipeline class. For more details, please refer to the 'how to add a new architecture' section in our documentation. - """) + """ + ) self.model_map = kwargs.pop("model_map", None) self._validate_config(config=config) self._validate_user_input(data=data) diff --git a/src/autoencodix/visualize/_general_visualizer.py b/src/autoencodix/visualize/_general_visualizer.py index 2483463f..96400617 100644 --- a/src/autoencodix/visualize/_general_visualizer.py +++ b/src/autoencodix/visualize/_general_visualizer.py @@ -473,18 +473,23 @@ def _plot_2D( print( "The provided label column is numeric and converted to categories." ) - labels = [ - float("nan") if not isinstance(x, float) else x for x in labels - ] - labels = ( - pd.qcut( - x=pd.Series(labels), - q=4, - labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], + # Try to convert all labels to float, if fails, convert to nan + for i in range(len(labels)): + try: + labels[i] = float(labels[i]) + except ValueError: + labels[i] = float("nan") + # Check if all labels are NaN, convert to string + if all(np.isnan(labels)): + labels = [str(x) for x in labels] + else: + labels = list( + pd.qcut( + x=pd.Series(labels), + q=4, + labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], + ).astype(str) ) - .astype(str) - .to_list() - ) else: center = False ## Disable centering for numeric params numeric = True @@ -692,15 +697,23 @@ def _plot_latent_ridge( # print(labels[0]) if not isinstance(labels[0], str): if len(np.unique(labels)) > 3: - # Change all non-float labels to NaN - labels = [x if isinstance(x, float) else float("nan") for x in labels] - labels = list( - pd.qcut( - x=pd.Series(labels), - q=4, - labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], - ).astype(str) - ) + # Try to convert all labels to float, if fails, convert to nan + for i in range(len(labels)): + try: + labels[i] = float(labels[i]) + except ValueError: + labels[i] = float("nan") + # Check if all labels are NaN, convert to string + if all(np.isnan(labels)): + labels = [str(x) for x in labels] + else: + labels = list( + pd.qcut( + x=pd.Series(labels), + q=4, + labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], + ).astype(str) + ) else: labels = [str(x) for x in labels] diff --git a/src/autoencodix/visualize/_xmodal_visualizer.py b/src/autoencodix/visualize/_xmodal_visualizer.py index e7c2d837..d9287775 100644 --- a/src/autoencodix/visualize/_xmodal_visualizer.py +++ b/src/autoencodix/visualize/_xmodal_visualizer.py @@ -695,13 +695,24 @@ def _plot_latent_ridge_multi( # print(labels[0]) if not isinstance(labels[0], str): if len(np.unique(labels)) > 3: - # Change all non-float labels to NaN - labels = [x if isinstance(x, float) else float("nan") for x in labels] - labels = pd.qcut( - x=pd.Series(labels), - q=4, - labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], - ).astype(str) + # Try to convert all labels to float, if fails, convert to nan + for i in range(len(labels)): + try: + labels[i] = float(labels[i]) + except ValueError: + labels[i] = float("nan") + # Check if all labels are NaN, convert to string + if all(np.isnan(labels)): + labels = [str(x) for x in labels] + else: + labels = list( + pd.qcut( + x=pd.Series(labels), + q=4, + labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], + ).astype(str) + ) + else: labels = [str(x) for x in labels] diff --git a/src/autoencodix/visualize/visualize.py b/src/autoencodix/visualize/visualize.py index 9f09d4bf..3e8e7251 100644 --- a/src/autoencodix/visualize/visualize.py +++ b/src/autoencodix/visualize/visualize.py @@ -704,19 +704,23 @@ def plot_2D( print( "The provided label column is numeric and converted to categories." ) - # Change non-float labels to NaN - labels = [ - x if isinstance(x, float) else float("nan") for x in labels - ] - labels = ( - pd.qcut( - x=pd.Series(labels), - q=4, - labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], + # Try to convert all labels to float, if fails, convert to nan + for i in range(len(labels)): + try: + labels[i] = float(labels[i]) + except ValueError: + labels[i] = float("nan") + # Check if all labels are NaN, convert to string + if all(np.isnan(labels)): + labels = [str(x) for x in labels] + else: + labels = list( + pd.qcut( + x=pd.Series(labels), + q=4, + labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], + ).astype(str) ) - .astype(str) - .to_list() - ) else: center = False ## Disable centering for numeric params numeric = True @@ -834,13 +838,23 @@ def plot_latent_ridge( # print(labels[0]) if not isinstance(labels[0], str): if len(np.unique(labels)) > 3: - # Change non-float labels to NaN - labels = [x if isinstance(x, float) else float("nan") for x in labels] - labels = pd.qcut( - x=pd.Series(labels), - q=4, - labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], - ).astype(str) + # Try to convert all labels to float, if fails, convert to nan + for i in range(len(labels)): + try: + labels[i] = float(labels[i]) + except ValueError: + labels[i] = float("nan") + # Check if all labels are NaN, convert to string + if all(np.isnan(labels)): + labels = [str(x) for x in labels] + else: + labels = list( + pd.qcut( + x=pd.Series(labels), + q=4, + labels=["1stQ", "2ndQ", "3rdQ", "4thQ"], + ).astype(str) + ) else: labels = [str(x) for x in labels]