diff --git a/docs/changes/79.bugfix.md b/docs/changes/79.bugfix.md new file mode 100644 index 0000000..bf44b91 --- /dev/null +++ b/docs/changes/79.bugfix.md @@ -0,0 +1 @@ +Harden gamma/hadron classification with robust features, inactive-telescope masking, source-aware validation splits, and stricter input validation. diff --git a/src/eventdisplay_ml/config.py b/src/eventdisplay_ml/config.py index c22611b..ee8411e 100644 --- a/src/eventdisplay_ml/config.py +++ b/src/eventdisplay_ml/config.py @@ -104,6 +104,15 @@ def configure_training(analysis_type): action="store_true", help="Remove ze_bin from gamma/hadron training features.", ) + parser.add_argument( + "--feature_profile", + choices=("robust", "extended"), + default="robust", + help=( + "Classification feature set. 'robust' uses stable image/stereo variables " + "available in existing data; 'extended' retains the historical feature set." + ), + ) parser.add_argument( "--max_cores", type=int, @@ -194,6 +203,7 @@ def configure_training(analysis_type): f"Balance class zenith weights: {model_configs.get('balance_class_zenith_weights')}" ) _logger.info(f"Ignore ze_bin feature: {model_configs.get('ignore_ze_bin')}") + _logger.info(f"Classification feature profile: {model_configs.get('feature_profile')}") model_configs["models"] = hyper_parameters( analysis_type, model_configs.get("hyperparameter_config") diff --git a/src/eventdisplay_ml/data_processing.py b/src/eventdisplay_ml/data_processing.py index 313a073..789c529 100644 --- a/src/eventdisplay_ml/data_processing.py +++ b/src/eventdisplay_ml/data_processing.py @@ -79,6 +79,21 @@ def read_telescope_config(root_file): } +def _telescope_configs_match(first, second): + """Compare telescope configurations, tolerating float serialization noise.""" + for key in ("tel_ids", "mirror_area", "tel_x", "tel_y"): + first_values = np.asarray(first[key]) + second_values = np.asarray(second[key]) + if first_values.shape != second_values.shape: + return False + if key == "tel_ids": + if not np.array_equal(first_values, second_values): + return False + elif not np.allclose(first_values, second_values, rtol=1e-7, atol=1e-7, equal_nan=True): + return False + return True + + def _resolve_branch_aliases(tree, branch_list): """ Resolve branch name aliases (e.g. R_core vs R) and drop missing optional branches. @@ -468,6 +483,7 @@ def flatten_telescope_data_vectorized( Flattened DataFrame with per-telescope columns suffixed by ``_{i}`` """ flat_features = {} + classification_mode = analysis_type == "classification" tel_list_matrix = _to_dense_array(df["DispTelList_T"]) n_evt = len(df) max_tel_id = tel_config["max_tel_id"] if tel_config else (n_tel - 1) @@ -490,6 +506,11 @@ def flatten_telescope_data_vectorized( size_data = _normalize_telescope_variable_to_tel_id_space( _to_dense_array(df["size"]), index_list_for_remapping, max_tel_id, n_evt ) + # A telescope absent from DispTelList_T is not a zero-sized image. Keep it + # explicitly missing so sorting and the XGBoost missing-value path cannot + # learn a detector-slot/observing-condition proxy. + if classification_mode: + size_data = np.where(active_mask, size_data, np.nan) size_data = _clip_size_array(size_data) core_x, core_y = _get_core_arrays(df) @@ -550,6 +571,9 @@ def flatten_telescope_data_vectorized( data, index_list_for_remapping, max_tel_id, n_evt ) + if classification_mode and var != "tel_active": + data_normalized = np.where(active_mask, data_normalized, np.nan) + # All variables are now in telescope-ID space; apply sorting and flatten uniformly data_normalized = data_normalized[np.arange(n_evt)[:, np.newaxis], sort_indices] @@ -947,6 +971,7 @@ def load_training_data(model_configs, file_list, analysis_type): pandas.DataFrame Flattened DataFrame ready for training. """ + classification_mode = analysis_type == "classification" max_events = model_configs.get("max_events", None) random_state = model_configs.get("random_state", None) memory_profile = model_configs.get("memory_profile", False) @@ -962,6 +987,8 @@ def load_training_data(model_configs, file_list, analysis_type): _logger.info(f"Adding zenith binning: {model_configs.get('zenith_bins_deg', [])}") input_files = utils.read_input_file_list(file_list) + if classification_mode and not input_files: + raise ValueError(f"Input file list is empty: {file_list}") tmva_style = model_configs.get("tmva_style", False) if tmva_style and analysis_type == "classification": @@ -979,9 +1006,15 @@ def load_training_data(model_configs, file_list, analysis_type): max_events_per_file = max_events // len(input_files) else: max_events_per_file = None + if classification_mode and max_events is not None and max_events > 0: + # Integer floor division can turn a small cap into zero, which means + # unlimited sampling. Classification applies an exact final cap below. + max_events_per_file = max(1, int(np.ceil(max_events / len(input_files)))) _logger.info(f"Max events per file: {max_events_per_file}") tel_config = None # Will be read from first file + if classification_mode: + tel_config = model_configs.get("tel_config") dfs = [] executor = ThreadPoolExecutor(max_workers=model_configs.get("max_cores", 1)) total_files = len(input_files) @@ -998,7 +1031,13 @@ def load_training_data(model_configs, file_list, analysis_type): else: # Check if current file has a larger max_tel_id and update if needed current_tel_config = read_telescope_config(root_file) - if current_tel_config["max_tel_id"] > tel_config["max_tel_id"]: + if classification_mode: + if not _telescope_configs_match(current_tel_config, tel_config): + raise ValueError( + "Classification/training input files have incompatible " + f"telescope configurations: {f}." + ) + elif current_tel_config["max_tel_id"] > tel_config["max_tel_id"]: _logger.info( f"Updating telescope configuration: max_tel_id from " f"{tel_config['max_tel_id']} to {current_tel_config['max_tel_id']} " @@ -1105,6 +1144,9 @@ def load_training_data(model_configs, file_list, analysis_type): if file_df is None or file_df.empty: continue + if analysis_type == "classification": + file_df["__source_file"] = str(f) + _logger.info( f"Number of events before / after event cut: {n_before} / " f"{n_after_event_cut} (fraction retained: {n_after_event_cut / n_before:.4f})" @@ -1122,9 +1164,25 @@ def load_training_data(model_configs, file_list, analysis_type): enabled=memory_profile, ) except Exception as e: + if classification_mode and isinstance(e, (KeyError, ValueError)): + raise raise FileNotFoundError(f"Error opening or reading file {f}: {e}") from e + if classification_mode and not dfs: + raise ValueError("No data loaded from input files.") df_final = pd.concat(dfs, ignore_index=True) + if ( + analysis_type == "classification" + and max_events is not None + and max_events > 0 + and len(df_final) > max_events + ): + df_final = df_final.sample( + n=max_events, + random_state=random_state, + ignore_index=True, + ) + _logger.info("Applied global classification event cap: %d", max_events) del dfs utils.log_memory_checkpoint("after final pandas concat", df_final, enabled=memory_profile) all_nan_columns = [col for col in df_final.columns if df_final[col].isna().all()] @@ -1443,11 +1501,35 @@ def extra_columns(df, analysis_type, training, index, tel_config=None, observato def zenith_in_bins(zenith_angles, bins): - """Apply zenith binning based on zenith angles and given bin edges.""" + """Apply zenith binning, marking out-of-range angles with ``-1``.""" + if bins is None or len(bins) == 0: + raise ValueError("Zenith-bin definitions must not be empty.") if isinstance(bins[0], dict): - bins = [b["Ze_min"] for b in bins] + [bins[-1]["Ze_max"]] + if not all(isinstance(value, dict) for value in bins): + raise ValueError("Zenith-bin definitions must use one format.") + try: + edges = [float(bins[0]["Ze_min"]), *(float(b["Ze_max"]) for b in bins)] + starts = np.asarray([float(b["Ze_min"]) for b in bins[1:]]) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError("Zenith-bin dictionaries require numeric Ze_min and Ze_max.") from exc + if not np.allclose(starts, edges[1:-1]): + raise ValueError("Zenith-bin dictionaries must be ordered and contiguous.") + bins = edges + elif len(bins) < 2: + raise ValueError("At least two zenith-bin edges are required.") bins = np.asarray(bins, dtype=float) - idx = np.clip(np.digitize(zenith_angles, bins) - 1, 0, len(bins) - 2) + if bins.ndim != 1 or len(bins) < 2 or not np.all(np.isfinite(bins)): + raise ValueError("Zenith-bin edges must be a finite one-dimensional sequence.") + if np.any(np.diff(bins) <= 0): + raise ValueError("Zenith-bin edges must be strictly increasing.") + zenith_angles = np.asarray(zenith_angles, dtype=float) + idx = np.full(zenith_angles.shape, -1, dtype=np.int32) + valid = np.isfinite(zenith_angles) & (zenith_angles >= bins[0]) & (zenith_angles <= bins[-1]) + if np.any(valid): + idx[valid] = np.minimum( + np.digitize(zenith_angles[valid], bins) - 1, + len(bins) - 2, + ) return idx.astype(np.int32) diff --git a/src/eventdisplay_ml/features.py b/src/eventdisplay_ml/features.py index d4b277d..edbdc9e 100644 --- a/src/eventdisplay_ml/features.py +++ b/src/eventdisplay_ml/features.py @@ -47,6 +47,49 @@ def target_features(analysis_type): raise ValueError(f"Unknown analysis type: {analysis_type}") +def classification_feature_columns(columns, profile="extended", ignore_ze_bin=False): + """Select classification features from a flattened data frame.""" + if profile not in {"robust", "extended"}: + raise ValueError("classification feature profile must be 'robust' or 'extended'") + + selected = [ + name + for name in columns + if name not in {"label", "Erec", "ErecS"} and not name.startswith("__") + ] + if profile == "robust": + event_features = { + "MSCW", + "MSCL", + "EChi2S", + "EmissionHeight", + "EmissionHeightChi2", + "Core_Distance", + "ze_bin", + } + telescope_features = ( + "cosphi_", + "sinphi_", + "loss_", + "dist_", + "width_", + "length_", + "asym_", + "tgrad_x_", + ) + selected = [ + name + for name in selected + if name in event_features or name.startswith(telescope_features) + ] + + if ignore_ze_bin: + selected = [name for name in selected if name != "ze_bin"] + if not selected: + raise ValueError(f"No usable classification features for profile '{profile}'.") + return selected + + def excluded_features(analysis_type, ntel): """ Features not to be used for training/prediction. diff --git a/src/eventdisplay_ml/models.py b/src/eventdisplay_ml/models.py index e37af42..4b8f12d 100644 --- a/src/eventdisplay_ml/models.py +++ b/src/eventdisplay_ml/models.py @@ -243,6 +243,8 @@ def _calculate_classification_thresholds(efficiency, min_efficiency=0.2, steps=5 dict[int, float] Mapping from efficiency (percent) to classification threshold. """ + if efficiency is None or len(efficiency) == 0: + raise ValueError("Classification efficiency diagnostics are missing from the model file.") df = efficiency.copy() df = df.sort_values("signal_efficiency") eff_targets = np.arange(min_efficiency * 100, 100, steps) / 100.0 @@ -298,7 +300,20 @@ def _validate_energy_bin_metadata(energy_bin, model_file): f"missing required key(s): {missing}." ) - return energy_bin + try: + e_min = float(energy_bin["E_min"]) + e_max = float(energy_bin["E_max"]) + except (TypeError, ValueError) as exc: + raise ValueError( + "Classification model file " + f"'{model_file}' has non-numeric energy-bin metadata for 'E_min'/'E_max'." + ) from exc + if not np.isfinite(e_min) or not np.isfinite(e_max) or e_min >= e_max: + raise ValueError( + "Classification model file " + f"'{model_file}' has invalid energy-bin metadata: require finite E_min < E_max." + ) + return {"E_min": e_min, "E_max": e_max} def _update_parameters(full_params, zenith_bins, energy_bin, e_bin_number): @@ -507,6 +522,18 @@ def apply_classification_models(df, model_configs, threshold_keys): if e_bin_lo == -1 or e_bin_hi == -1: _logger.warning("Skipping events with invalid energy interpolation bins") continue + if "ze_bin" in group_df: + zenith_values = pd.to_numeric(group_df["ze_bin"], errors="coerce") + valid_zenith = zenith_values.notna() & (zenith_values >= 0) + if not valid_zenith.all(): + _logger.warning( + "Skipping %d events with invalid/out-of-range zenith bins during " + "classification apply", + int((~valid_zenith).sum()), + ) + group_df = group_df.loc[valid_zenith] + if group_df.empty: + continue _logger.info( "Processing %d events with interpolation bins (%d, %d)", @@ -526,8 +553,15 @@ def apply_classification_models(df, model_configs, threshold_keys): ) model_lo = models[e_bin_lo]["model"] model_hi = models[e_bin_hi]["model"] - flatten_lo = flatten_data.reindex(columns=models[e_bin_lo]["features"]) - flatten_hi = flatten_data.reindex(columns=models[e_bin_hi]["features"]) + missing_lo = sorted(set(models[e_bin_lo]["features"]) - set(flatten_data.columns)) + missing_hi = sorted(set(models[e_bin_hi]["features"]) - set(flatten_data.columns)) + if missing_lo or missing_hi: + raise ValueError( + "Classification model/input feature schema mismatch: " + f"low-bin missing={missing_lo}, high-bin missing={missing_hi}." + ) + flatten_lo = flatten_data.loc[:, models[e_bin_lo]["features"]] + flatten_hi = flatten_data.loc[:, models[e_bin_hi]["features"]] class_probs_lo = model_lo.predict_proba(flatten_lo)[:, 1] if e_bin_lo == e_bin_hi: @@ -1008,93 +1042,164 @@ def train_regression(df, model_configs): def train_classification(df, model_configs): - """ - Train a single XGBoost model for gamma/hadron classification. - - Parameters - ---------- - df : list of pd.DataFrame - Training data. - model_configs : dict - Dictionary of model configurations. - """ + """Train a single XGBoost model for gamma/hadron classification.""" if df[0].empty or df[1].empty: raise ValueError( "Classification training requires non-empty signal and background data. " f"signal_events={len(df[0])}, background_events={len(df[1])}." ) + if set(df[0].columns) != set(df[1].columns): + raise ValueError( + "Signal/background classification schemas differ. " + f"Only signal: {sorted(set(df[0].columns) - set(df[1].columns))}; " + f"only background: {sorted(set(df[1].columns) - set(df[0].columns))}" + ) - df[0]["label"] = 1 - df[1]["label"] = 0 - full_df = pd.concat([df[0], df[1]], ignore_index=True) - ze_data = full_df["ze_bin"] if "ze_bin" in full_df.columns else None - x_data = full_df.drop(columns=["label"]) - if model_configs.get("ignore_ze_bin", False): - if model_configs.get("balance_class_zenith_weights", False): - raise ValueError("Cannot use ignore_ze_bin with balance_class_zenith_weights.") - if "ze_bin" in x_data.columns: - _logger.info("Removing ze_bin from classification training features.") - x_data = x_data.drop(columns=["ze_bin"]) - _logger.info(f"Features ({len(x_data.columns)}): {', '.join(x_data.columns)}") - model_configs["features"] = list(x_data.columns) + signal = df[0].copy() + background = df[1].copy() + signal["label"] = 1 + background["label"] = 0 + full_df = pd.concat([signal, background], ignore_index=True) y_data = full_df["label"] - - split_inputs = [x_data, y_data] + ze_data = full_df.get("ze_bin") if ze_data is not None: - split_inputs.append(ze_data) + numeric_zenith = pd.to_numeric(ze_data, errors="coerce") + invalid_zenith = numeric_zenith.isna() | (numeric_zenith < 0) + if invalid_zenith.any(): + raise ValueError( + "Classification training contains out-of-range or invalid zenith bins: " + f"{int(invalid_zenith.sum())} events." + ) - split_result = train_test_split( - *split_inputs, - train_size=model_configs.get("train_test_fraction", 0.5), - random_state=model_configs.get("random_state", None), - stratify=y_data, + profile = ( + "extended" + if model_configs.get("tmva_style", False) + else model_configs.get("feature_profile", "extended") + ) + feature_columns = features.classification_feature_columns( + full_df.columns, + profile=profile, + ignore_ze_bin=model_configs.get("ignore_ze_bin", False), + ) + all_nan = [ + column + for column in feature_columns + if signal[column].isna().all() or background[column].isna().all() + ] + if all_nan: + raise ValueError( + f"Classification features must contain values in both classes: {', '.join(all_nan)}" + ) + + x_data = full_df.loc[:, feature_columns] + model_configs["features"] = feature_columns + _logger.info("Features (%d): %s", len(feature_columns), ", ".join(feature_columns)) + + train_idx, validation_idx, test_idx, split_method = _classification_split_indices( + y_data, + full_df.get("__source_file"), + model_configs.get("train_test_fraction", 0.5), + model_configs.get("random_state"), + ) + x_train, x_validation, x_test = ( + x_data.iloc[index] for index in (train_idx, validation_idx, test_idx) + ) + y_train, y_validation, y_test = ( + y_data.iloc[index] for index in (train_idx, validation_idx, test_idx) + ) + ze_test = ze_data.iloc[test_idx] if ze_data is not None else None + model_configs["classification_split"] = { + "method": split_method, + "n_train": len(train_idx), + "n_validation": len(validation_idx), + "n_test": len(test_idx), + } + _logger.info( + "Classification split: train=%d validation=%d test=%d (%s)", + len(x_train), + len(x_validation), + len(x_test), + split_method, ) - if ze_data is None: - x_train, x_test, y_train, y_test = split_result - ze_test = None - else: - x_train, x_test, y_train, y_test, _, ze_test = split_result - _logger.info(f"Training events: {len(x_train)}, Testing events: {len(x_test)}") weights_train = None if model_configs.get("balance_class_zenith_weights", False): - weights_train = _class_zenith_balance_weights(x_train, y_train) - _logger.info( - "Using class/zenith sample weights " - f"(mean={weights_train.mean():.3f}, std={weights_train.std():.3f}, " - f"min={weights_train.min():.3f}, max={weights_train.max():.3f})" - ) - eval_set = [(x_train, y_train), (x_test, y_test)] + weights_train = _class_zenith_balance_weights(full_df.iloc[train_idx], y_train) + eval_set = [(x_train, y_train), (x_validation, y_validation)] for name, cfg in model_configs.get("models", {}).items(): - _logger.info(f"Training {name}") + _logger.info("Training %s", name) model = xgb.XGBClassifier(**cfg.get("hyper_parameters", {})) fit_kwargs = {"eval_set": eval_set, "verbose": True} if weights_train is not None: fit_kwargs["sample_weight"] = weights_train model.fit(x_train, y_train, **fit_kwargs) - shap_importance = evaluate_classification_model( - model, - x_test, - y_test, - full_df, - x_data.columns.tolist(), - name, - ) cfg["model"] = model - cfg["features"] = x_data.columns.tolist() # Store feature names for diagnostics - efficiency_all, efficiencies_by_zenith = evaluation_efficiency( + cfg["features"] = feature_columns + cfg["shap_importance"] = evaluate_classification_model( + model, x_test, y_test, full_df, feature_columns, name + ) + efficiency, efficiencies_by_zenith = evaluation_efficiency( name, model, x_test, y_test, return_by_zenith=True, ze_bins=ze_test ) - cfg["efficiency"] = efficiency_all + cfg["efficiency"] = efficiency for ze_bin, ze_efficiency in efficiencies_by_zenith.items(): cfg[f"efficiency_ze{ze_bin}"] = ze_efficiency - cfg["shap_importance"] = shap_importance return model_configs +def _classification_split_indices(labels, groups, train_fraction, random_state): + """Return source-grouped train, validation, and test row indices.""" + if not 0 < train_fraction < 1: + raise ValueError("train_test_fraction must be between zero and one.") + + indices = np.arange(len(labels)) + if groups is not None: + group_labels = pd.DataFrame({"group": groups, "label": labels}).drop_duplicates() + groups_are_class_specific = not group_labels["group"].duplicated().any() + if groups_are_class_specific: + try: + train_groups, holdout_groups = train_test_split( + group_labels, + train_size=train_fraction, + random_state=random_state, + stratify=group_labels["label"], + ) + validation_groups, test_groups = train_test_split( + holdout_groups, + train_size=0.5, + random_state=random_state, + stratify=holdout_groups["label"], + ) + return ( + indices[groups.isin(train_groups["group"])], + indices[groups.isin(validation_groups["group"])], + indices[groups.isin(test_groups["group"])], + "grouped_source_file", + ) + except ValueError: + _logger.warning( + "Not enough source files for a grouped classification split; " + "falling back to a stratified event split." + ) + + train_idx, holdout_idx = train_test_split( + indices, + train_size=train_fraction, + random_state=random_state, + stratify=labels, + ) + validation_idx, test_idx = train_test_split( + holdout_idx, + train_size=0.5, + random_state=random_state, + stratify=labels.iloc[holdout_idx], + ) + return train_idx, validation_idx, test_idx, "stratified_event" + + def _class_zenith_balance_weights(x_train, y_train): """Compute sample weights that equalize class distributions over ze_bin.""" if "ze_bin" not in x_train.columns: diff --git a/tests/test_classification_apply_interpolation.py b/tests/test_classification_apply_interpolation.py index 198d71e..1b4196a 100644 --- a/tests/test_classification_apply_interpolation.py +++ b/tests/test_classification_apply_interpolation.py @@ -72,6 +72,35 @@ def test_apply_classification_models_interpolates_probabilities_and_thresholds(m np.testing.assert_array_equal(is_gamma[50], np.array([0, 1], dtype=np.uint8)) +def test_apply_leaves_out_of_range_zenith_events_invalid(monkeypatch): + """Events outside the trained zenith range must not be scored by an edge model.""" + df = pd.DataFrame( + { + "Erec": [1.0], + "e_bin_lo": [0], + "e_bin_hi": [0], + "e_alpha": [0.0], + "ze_bin": [-1], + "dummy": [1.0], + } + ) + model_configs = { + "models": { + 0: { + "model": DummyXGBClassifier(1.0), + "features": ["dummy"], + "thresholds": {50: 0.5}, + } + } + } + monkeypatch.setattr(models, "flatten_feature_data", lambda *args, **kwargs: df[["dummy"]]) + + class_probability, is_gamma = models.apply_classification_models(df, model_configs, [50]) + + assert np.isnan(class_probability[0]) + assert is_gamma[50][0] == 0 + + def test_extra_columns_skip_tmva_only_size_second_max_when_branch_missing(): """Standard classification should not synthesize the TMVA-only SizeSecondMax column.""" df = pd.DataFrame( diff --git a/tests/test_classification_robustness.py b/tests/test_classification_robustness.py new file mode 100644 index 0000000..8a4d4ad --- /dev/null +++ b/tests/test_classification_robustness.py @@ -0,0 +1,79 @@ +"""Focused tests for classification hardening.""" + +import numpy as np +import pandas as pd + +from eventdisplay_ml import data_processing, features, models + + +def test_telescope_config_comparison_tolerates_float_noise_and_nan(): + first = { + "tel_ids": np.array([1, 2]), + "mirror_area": np.array([100.0, np.nan]), + "tel_x": np.array([0.0, 10.0]), + "tel_y": np.array([1.0, 2.0]), + } + second = { + "tel_ids": np.array([1, 2]), + "mirror_area": np.array([100.0 + 1e-8, np.nan]), + "tel_x": np.array([0.0, 10.0 + 1e-8]), + "tel_y": np.array([1.0, 2.0]), + } + + assert data_processing._telescope_configs_match(first, second) + + +def test_robust_profile_excludes_routing_and_activity_columns(): + columns = [ + "MSCW", + "MSCL", + "width_0", + "length_0", + "tel_active_0", + "mirror_area_0", + "ze_bin", + "Erec", + "__source_file", + ] + + assert features.classification_feature_columns(columns, profile="robust") == [ + "MSCW", + "MSCL", + "width_0", + "length_0", + "ze_bin", + ] + + +def test_extended_profile_retains_historical_features_but_excludes_routing(): + columns = [ + "MSCW", + "DispNImages", + "ze_bin", + "Erec", + "tel_active_0", + "__source_file", + ] + + assert features.classification_feature_columns(columns) == [ + "MSCW", + "DispNImages", + "ze_bin", + "tel_active_0", + ] + + +def test_grouped_split_keeps_source_files_disjoint(): + labels = pd.Series(np.repeat([0, 1], 80)) + groups = pd.Series( + np.concatenate([np.repeat(np.arange(8), 10), np.repeat(np.arange(8, 16), 10)]) + ) + + train, validation, test, method = models._classification_split_indices( + labels, groups, train_fraction=0.5, random_state=7 + ) + + assert method == "grouped_source_file" + assert set(groups.iloc[train]).isdisjoint(groups.iloc[validation]) + assert set(groups.iloc[train]).isdisjoint(groups.iloc[test]) + assert set(groups.iloc[validation]).isdisjoint(groups.iloc[test]) diff --git a/tests/test_data_processing.py b/tests/test_data_processing.py index e565113..a0c6fe3 100644 --- a/tests/test_data_processing.py +++ b/tests/test_data_processing.py @@ -2,18 +2,19 @@ import numpy as np import pandas as pd +import pytest from eventdisplay_ml.data_processing import energy_interpolation_bins, zenith_in_bins -def test_zenith_in_bins_numeric_edges_clips_and_handles_boundaries(): - """Numeric bin edges should clip out-of-range values and place edge values consistently.""" +def test_zenith_in_bins_numeric_edges_marks_invalid_and_handles_boundaries(): + """Numeric bin edges should mark out-of-range values instead of clipping them.""" zenith_angles = np.array([-5.0, 0.0, 9.9, 10.0, 19.9, 20.0, 42.0], dtype=float) bins = [0.0, 10.0, 20.0, 30.0] result = zenith_in_bins(zenith_angles, bins) - np.testing.assert_array_equal(result, np.array([0, 0, 0, 1, 1, 2, 2], dtype=np.int32)) + np.testing.assert_array_equal(result, np.array([-1, 0, 0, 1, 1, 2, -1], dtype=np.int32)) assert result.dtype == np.int32 @@ -28,10 +29,34 @@ def test_zenith_in_bins_dict_bins_matches_numeric_definition(): result = zenith_in_bins(zenith_angles, dict_bins) - np.testing.assert_array_equal(result, np.array([0, 0, 0, 1, 1, 2, 2], dtype=np.int32)) + np.testing.assert_array_equal(result, np.array([-1, 0, 0, 1, 1, 2, -1], dtype=np.int32)) assert result.dtype == np.int32 +def test_zenith_in_bins_accepts_one_dictionary_bin(): + result = zenith_in_bins([0.0, 10.0, 20.0], [{"Ze_min": 0.0, "Ze_max": 20.0}]) + + np.testing.assert_array_equal(result, np.array([0, 0, 0], dtype=np.int32)) + + +def test_zenith_in_bins_rejects_noncontiguous_dict_bins(): + bins = [ + {"Ze_min": 0.0, "Ze_max": 10.0}, + {"Ze_min": 12.0, "Ze_max": 20.0}, + ] + with pytest.raises(ValueError, match="ordered and contiguous"): + zenith_in_bins([5.0], bins) + + +def test_zenith_in_bins_rejects_invalid_dict_bin_bounds(): + bins = [ + {"Ze_min": 10.0, "Ze_max": 0.0}, + {"Ze_min": 0.0, "Ze_max": 20.0}, + ] + with pytest.raises(ValueError, match="strictly increasing"): + zenith_in_bins([5.0], bins) + + def test_energy_interpolation_bins_interpolates_and_clamps_with_invalid_events(): """Interpolation bins should handle interior, edge, and invalid energies robustly.""" df_chunk = pd.DataFrame({"Erec": [0.0, 0.1, 1.0, 10.0, 100.0]}) diff --git a/tests/test_train_classification_shap.py b/tests/test_train_classification_shap.py index 0bfe0d0..c0d6277 100644 --- a/tests/test_train_classification_shap.py +++ b/tests/test_train_classification_shap.py @@ -122,6 +122,15 @@ def test_class_zenith_balance_weights_equalize_class_zenith_distributions(): assert ze1 / (ze0 + ze1) == pytest.approx(0.5) +def test_train_classification_rejects_invalid_zenith_bins(): + """Invalid zenith routing states must fail before model fitting.""" + signal = pd.DataFrame({"f1": [1.0, 2.0, 3.0], "ze_bin": [-1, 0, 0]}) + background = pd.DataFrame({"f1": [-1.0, -2.0, -3.0], "ze_bin": [0, 0, 0]}) + + with pytest.raises(ValueError, match="out-of-range or invalid zenith bins"): + models.train_classification([signal, background], {"models": {}}) + + def test_train_classification_applies_class_zenith_weights(): """The optional class/zenith balance weights should be passed to XGBoost.""" signal_df = pd.DataFrame(