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
|
Reconstruct the local model folder the user picked from a cached model. |
|
Run the region classifier and average its per-fold probabilities. |
|
Resolve the folds directory for a local model directory. |
|
Download a named S3 model and return its folds directory. |
Return the cached inference model, loading or prompting for it on first use. |
|
Drop cached inference predictions on a shank so the next click recomputes. |
|
Return whether an inference model is cached on the plugin. |
|
Load the region-classifier model from a local directory or download it from S3. |
|
Run the full inference-model load GUI and cache the result. |
|
Predict the per-channel region (argmax over fold-averaged probabilities). |
|
Return cumulative region probabilities by depth for the stacked-band view. |
|
Ensure the features DataFrame has every column the model expects. |
|
Validate a fold-based model directory and extract its feature/class contract. |
|
Check the model directory has the expected fold structure. |
Classes
A loaded fold-based region classifier and its feature/class contract. |
|
|
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:
objectA 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; theshankandconfigkeywords injected by the decorator are absorbed via**kwargs.- Parameters:
controller (AlignmentGUIController) – The main application controller.
items (ShankController) – The shank whose cached predictions are cleared.
- 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
InferenceModelunder[MODEL_NAME]['model'], plus the chosen source (local_inference_dirandmodel_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 nofoldssubdirectory).model_name (str or None) – S3 model name to download when
model_diris None. May be a nested name such asxgboost_channels/2026_W12_Cosmos_careless-clover-dingo. Defaults toMODEL_VINTAGEwhen 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:
controller (AlignmentGUIController) – The main application controller.
items (ShankController) – The shank being predicted on.
- 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:
controller (AlignmentGUIController) – The main application controller.
items (ShankController) – The shank being predicted on.
- 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
outsidecolumn with0.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
outsideadded if it was missing.- Return type:
pandas.DataFrame
- Raises:
ValueError – If a required column other than
outsideis 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.yamlunderfolds_pathand asserts every fold agrees on itsFEATURESandCLASSESlists. 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
FOLD0Xsubdirectories, each with ameta.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.yamlis found underfolds_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_dircontainsfolds/FOLD00orFOLD00directly, else False.- Return type:
bool