ibl_alignment_gui.plugins.ephys_atlas.inference

Region-classifier inference for the channel-prediction plugin.

Loads a fold-based XGBoost region classifier — either from a local directory or downloaded from S3 — and runs ephysatlas.regionclassifier.infer_regions to predict per-channel Cosmos regions from an ephys-feature table. Heavy ephysatlas imports are deferred so this module can be imported in offline mode without it installed; it is only required when inference actually runs.

Functions

_current_local_dir

Reconstruct the local model folder the user picked from a cached model.

_fold_mean_probas

Run the region classifier and average its per-fold probabilities.

_get_model_path_from_local

Resolve the folds directory for a local model directory.

_get_model_path_from_s3

Download a named S3 model and return its folds directory.

get_model

Return the cached inference model, loading or prompting for it on first use.

invalidate_predictions

Drop cached inference predictions on a shank so the next click recomputes.

is_model_loaded

Return whether an inference model is cached on the plugin.

load_inference_model

Load the region-classifier model from a local directory or download it from S3.

load_model_dialog

Run the full inference-model load GUI and cache the result.

predict

Predict the per-channel region (argmax over fold-averaged probabilities).

predict_cumulative

Return cumulative region probabilities by depth for the stacked-band view.

validate_features

Ensure the features DataFrame has every column the model expects.

validate_model

Validate a fold-based model directory and extract its feature/class contract.

validate_model_folder

Check the model directory has the expected fold structure.

Classes

InferenceModel

A loaded fold-based region classifier and its feature/class contract.

_InferenceModelDialog

Inference-model selection dialog with optional dropdown and a local-folder picker.

class ibl_alignment_gui.plugins.ephys_atlas.inference.InferenceModel(features, classes, model_path, n_folds)[source]

Bases: object

A loaded fold-based region classifier and its feature/class contract.

Variables:
  • features (list of str) – Feature columns the model requires, in order.

  • classes (list of int) – Cosmos region ids the model predicts, in column order.

  • model_path (Path) – The folds directory passed to infer_regions.

  • n_folds (int) – Number of folds discovered during validation.

classes: list[int]
features: list[str]
model_path: Path
n_folds: int
ibl_alignment_gui.plugins.ephys_atlas.inference.get_model(controller)[source]

Return the cached inference model, loading or prompting for it on first use.

If no Channel Prediction model state exists yet, opens the load dialog and retries. If state exists but the model has not been built, rebuilds it from the cached local dir when present (otherwise prompts via the dialog) and retries. Returns None if the user cancels the dialog.

Parameters:

controller (AlignmentGUIController) – The main application controller.

Returns:

The loaded model, or None if the user cancelled loading.

Return type:

InferenceModel or None

ibl_alignment_gui.plugins.ephys_atlas.inference.invalidate_predictions(controller, items, **kwargs)[source]

Drop cached inference predictions on a shank so the next click recomputes.

Decorated with shank_loop(), so a single call iterates over every shank/config; the shank and config keywords injected by the decorator are absorbed via **kwargs.

Parameters:
Return type:

None

ibl_alignment_gui.plugins.ephys_atlas.inference.is_model_loaded(controller)[source]

Return whether an inference model is cached on the plugin.

Parameters:

controller (AlignmentGUIController) – The main application controller.

Returns:

True if a model has been loaded, else False.

Return type:

bool

ibl_alignment_gui.plugins.ephys_atlas.inference.load_inference_model(controller, model_dir=None, model_name=None, one=None)[source]

Load the region-classifier model from a local directory or download it from S3.

Caches the result on the Channel Prediction plugin: the loaded InferenceModel under [MODEL_NAME]['model'], plus the chosen source (local_inference_dir and model_name).

Parameters:
  • controller (AlignmentGUIController) – Provides ONE access for the S3 download path.

  • model_dir (str or Path or None) – If given, use this local model directory and skip S3. The fold models are expected under <model_dir>/folds/FOLD0X (or directly under <model_dir> if there is no folds subdirectory).

  • model_name (str or None) – S3 model name to download when model_dir is None. May be a nested name such as xgboost_channels/2026_W12_Cosmos_careless-clover-dingo. Defaults to MODEL_VINTAGE when not provided.

  • one (ONE or None) – ONE connection used for the S3 download.

Raises:

RuntimeError – If an S3 download is requested but no ONE/Alyx connection is available.

Return type:

None

ibl_alignment_gui.plugins.ephys_atlas.inference.load_model_dialog(controller)[source]

Run the full inference-model load GUI and cache the result.

Shows an intermediate dialog in both modes: offline offers only a local-folder picker, online adds a dropdown of named S3 models. A chosen local folder takes precedence over the dropdown. Calls load_inference_model() when the user selects a new or different model.

Parameters:

controller (AlignmentGUIController) – The main application controller.

Returns:

True if a model was loaded, False if the user cancelled.

Return type:

bool

ibl_alignment_gui.plugins.ephys_atlas.inference.predict(controller, items)[source]

Predict the per-channel region (argmax over fold-averaged probabilities).

Parameters:
Returns:

Predicted region ids and channel depths (µm), or None if no features are available.

Return type:

tuple of (np.ndarray, np.ndarray) or None

ibl_alignment_gui.plugins.ephys_atlas.inference.predict_cumulative(controller, items)[source]

Return cumulative region probabilities by depth for the stacked-band view.

Parameters:
Returns:

(cprobas, depths, colours, region_ids) or None if no features are available.

Return type:

tuple of (np.ndarray, np.ndarray, list of np.ndarray, np.ndarray) or None

ibl_alignment_gui.plugins.ephys_atlas.inference.validate_features(df, features)[source]

Ensure the features DataFrame has every column the model expects.

Fills an absent outside column with 0.0 (in-brain default, matching the upstream inference scripts); any other missing required column is a hard error.

Parameters:
  • df (pandas.DataFrame) – Per-channel features (one row per channel).

  • features (list of str) – Feature columns the model requires.

Returns:

The DataFrame, with outside added if it was missing.

Return type:

pandas.DataFrame

Raises:

ValueError – If a required column other than outside is absent.

ibl_alignment_gui.plugins.ephys_atlas.inference.validate_model(folds_path, max_folds=10)[source]

Validate a fold-based model directory and extract its feature/class contract.

Reads each FOLD0X/meta.yaml under folds_path and asserts every fold agrees on its FEATURES and CLASSES lists. This is what makes averaging across folds (and switching between models) safe — the GUI relies on a single, consistent FEATURES/CLASSES ordering.

Parameters:
  • folds_path (str or Path) – Directory containing FOLD0X subdirectories, each with a meta.yaml.

  • max_folds (int) – Upper bound on the number of folds to probe (folds are discovered, not assumed).

Returns:

The shared FEATURES list, the shared CLASSES list (Cosmos region ids), and the number of folds discovered.

Return type:

tuple of (list of str, list of int, int)

Raises:
  • FileNotFoundError – If no FOLD0X/meta.yaml is found under folds_path.

  • ValueError – If the folds disagree on FEATURES or CLASSES.

ibl_alignment_gui.plugins.ephys_atlas.inference.validate_model_folder(model_dir)[source]

Check the model directory has the expected fold structure.

Parameters:

model_dir (Path) – The local model directory to validate.

Returns:

True if model_dir contains folds/FOLD00 or FOLD00 directly, else False.

Return type:

bool