diff --git a/emulmps/emulmps.py b/emulmps/emulmps.py index d0e4c21..685f296 100644 --- a/emulmps/emulmps.py +++ b/emulmps/emulmps.py @@ -371,7 +371,7 @@ def calculate(self, state, want_derived=True, **params): emul_params, use_syren=self.use_syren ) - if not np.all(np.isfinite(Pk_lin_mpc)) or np.any(Pk_lin_mpc <= 0): + if not np.all(np.isfinite(Pk_lin_mpc)) or np.any(~(Pk_lin_mpc > 0)): self.log.debug(f"Non-finite or non-positive Pk_lin at params={params} — rejecting point.") return False @@ -380,7 +380,7 @@ def calculate(self, state, want_derived=True, **params): # ------------------------------------------------------------------ _, _, boost = self._emulator.get_boost(emul_params, pk_lin=Pk_lin_mpc, use_syren=self.use_syren) - if not np.all(np.isfinite(boost)) or np.any(boost <= 0): + if not np.all(np.isfinite(boost)) or np.any(~(boost > 0)): self.log.debug( f"Non-finite or non-positive boost at params={params} — " "rejecting point." @@ -389,7 +389,7 @@ def calculate(self, state, want_derived=True, **params): Pk_nl_mpc = (boost * Pk_lin_mpc).astype(np.float32) - if not np.all(np.isfinite(Pk_nl_mpc)) or np.any(Pk_nl_mpc <= 0): + if not np.all(np.isfinite(Pk_nl_mpc)) or np.any(~(Pk_nl_mpc > 0)): self.log.debug( f"Non-finite or non-positive Pk_nl at params={params} — " "rejecting point." diff --git a/emulmps/emulmps_emul/emulmps_w0wa.py b/emulmps/emulmps_emul/emulmps_w0wa.py index 420c8d6..9cacf93 100644 --- a/emulmps/emulmps_emul/emulmps_w0wa.py +++ b/emulmps/emulmps_emul/emulmps_w0wa.py @@ -22,11 +22,21 @@ from pathlib import Path import sys from . import train_utils_pk_emulator as utils -sys.modules['train_utils_pk_emulator'] = utils -sys.modules['train_utils_pk_emulator_v2'] = utils from . train_utils_pk_emulator import CustomActivationLayer, TComponentScaler import tensorflow as tf +class _AliasLoader: + """Redirect any train_utils_pk_emulator_vN import to the canonical module.""" + def find_module(self, name, path=None): + if name.startswith('train_utils_pk_emulator'): + return self + + def load_module(self, name): + if name not in sys.modules: + sys.modules[name] = utils + return sys.modules[name] + +sys.meta_path.append(_AliasLoader()) # --- Custom warning class --- class EmulatorWarning(UserWarning): @@ -607,7 +617,6 @@ def has_nl_model(self) -> bool: """Return True if a nonlinear boost model has been loaded.""" return self._nl_model is not None - # --- Public Module-Level Interface --- _pk_emulator_instance = None