"""Public helper module.

Conceptually, this module assembles the higher-level decoding workflows used by
the notebook pipelines.

It exists as a separate unit so trial selection, feature preparation, model
training, and output packaging can be shared across multiple decoding analyses.

It connects entry-level neural payloads to the low-level SVM kernels and
returns notebook-friendly result structures.
"""

from __future__ import annotations

from typing import Any, Dict, Iterable, List, Mapping, Optional, Set, Tuple

import numpy as np

from .signal_utils import nearest_bin
from .stats_utils import downsample_balance, zscore_cols
from .subspace import ensure_trials_by_cells
from .selection import get_ccf_mask, get_celltype_mask, get_completion_mask, get_quiet_mask
from .svm import decode_one_bin_svm, decode_one_bin_svm_matlab_compat

__all__ = [
    "run_afterwhisker_decoding_pipeline",
    "run_delay_decoding_pipeline",
    "run_nb_neurons_decoding_pipeline",
    "run_prewhisk_training_decoding_pipeline",
]


def run_nb_neurons_decoding_pipeline(
    entries_by_area: Mapping[str, List[Dict[str, Any]]],
    params: Mapping[str, Any],
    nb_neurons_list: Iterable[int],
    num_repetitions: int,
    area_list: Optional[Mapping[str, Set[str]]] = None,
    enable_ccf_filter: bool = True,
) -> Dict[str, Any]:
    """Run the repeated-neuron-count decoding pipeline.

    Parameters
    ----------
    entries_by_area:
        Probe entries grouped by area.
    params:
        Decoding configuration mapping containing region list, window centers,
        trial filters, and SVM settings.
    nb_neurons_list:
        Neuron-count settings evaluated independently.
    num_repetitions:
        Number of repeated random neuron subsets per neuron-count setting.
    area_list:
        Optional area-to-CCF mapping used by the unit filter.
    enable_ccf_filter:
        Whether CCF-based unit filtering is active.

    Returns
    -------
    dict[str, Any]
        Notebook-friendly accuracy bundle containing real and shuffled
        accuracies plus the shared window centers.
    """

    nb_neurons_list = [int(x) for x in nb_neurons_list]
    num_repetitions = int(num_repetitions)
    window_centers = np.asarray(params["windowCenters"], dtype=np.float32)
    n_bins = int(window_centers.size)

    accuracy: Dict[str, Any] = {}
    accuracy_shuffled: Dict[str, Any] = {}

    for i_nb in nb_neurons_list:
        name = f"Nb_neurons{i_nb}"
        accuracy[name] = {"sessionaddress": {}, "selected_neurons": {}}
        accuracy_shuffled[name] = {"sessionaddress": {}}

    # Iterate area by area so the output keeps the same nested structure as the
    # downstream notebooks and saved result files.
    for current_area in params["regionlist"]:
        probe_entries = list(entries_by_area.get(current_area, []))
        n_probes = len(probe_entries)

        for i_nb in nb_neurons_list:
            name = f"Nb_neurons{i_nb}"
            accuracy[name][current_area] = np.full((n_probes, n_bins, num_repetitions), np.nan, dtype=np.float32)
            accuracy_shuffled[name][current_area] = np.full(
                (n_probes, n_bins, num_repetitions), np.nan, dtype=np.float32
            )
            accuracy[name]["sessionaddress"][current_area] = np.full((n_probes,), None, dtype=object)
            accuracy_shuffled[name]["sessionaddress"][current_area] = np.full((n_probes,), None, dtype=object)
            accuracy[name]["selected_neurons"][current_area] = np.full((n_probes,), None, dtype=object)

        for iprobe, e in enumerate(probe_entries):
            trial = np.asarray(e["trial"]).reshape(-1)
            lick = np.asarray(e["lick"]).reshape(-1)

            completed_trials_ind = get_completion_mask(e, str(params["completion_state"]))
            qind = get_quiet_mask(e, str(params["quietstate"]))

            class1 = (trial == 1) & (lick == 1)
            class2 = (trial == 3) & (lick == 0)
            class1 = class1 & qind & completed_trials_ind
            class2 = class2 & qind & completed_trials_ind

            celltype_ind = get_celltype_mask(e, str(params["celltype"]))
            ccf_ind = get_ccf_mask(
                e,
                str(current_area),
                area_list=area_list,
                enable_ccf_filter=enable_ccf_filter,
            )
            curr_cell_ind = celltype_ind & ccf_ind

            currsig_all = np.asarray(e["spike_counts"], dtype=np.float32)[:, :, curr_cell_ind]
            n_units_avail = int(currsig_all.shape[2])

            n1 = int(np.sum(class1))
            n2 = int(np.sum(class2))

            for i_nb in nb_neurons_list:
                name = f"Nb_neurons{i_nb}"

                if n_units_avail < i_nb:
                    continue

                selected_neurons_mat = np.full((num_repetitions, i_nb), -1, dtype=np.int32)

                # Sample one neuron subset per repetition and evaluate every
                # time bin on that same subset.
                for i_repetition in range(1, num_repetitions + 1):
                    rng_rep = np.random.default_rng(i_repetition)
                    selected_one_based = rng_rep.permutation(n_units_avail)[:i_nb] + 1
                    selected_zero_based = selected_one_based - 1
                    selected_neurons_mat[i_repetition - 1, :] = selected_one_based.astype(np.int32, copy=False)

                    if min(n1, n2) < 5:
                        continue

                    currsig = currsig_all[:, :, selected_zero_based]
                    val_class1 = currsig[:, class1, :]
                    val_class2 = currsig[:, class2, :]

                    for i_bin in range(n_bins):
                        X1 = np.squeeze(val_class1[i_bin, :, :])
                        X2 = np.squeeze(val_class2[i_bin, :, :])

                        X1 = ensure_trials_by_cells(X1, n1)
                        X2 = ensure_trials_by_cells(X2, n2)

                        acc_val, sh_val = decode_one_bin_svm_matlab_compat(
                            X1,
                            X2,
                            mintrial=int(params["mintrial"]),
                            balance_method=str(params["balance_method"]),
                            zscoring=bool(params["zscoring"]),
                            run_seed=0,
                            rng_obj=rng_rep,
                        )
                        accuracy[name][current_area][iprobe, i_bin, i_repetition - 1] = acc_val
                        accuracy_shuffled[name][current_area][iprobe, i_bin, i_repetition - 1] = sh_val

                accuracy[name]["selected_neurons"][current_area][iprobe] = selected_neurons_mat
                accuracy[name]["sessionaddress"][current_area][iprobe] = e["session_id"]
                accuracy_shuffled[name]["sessionaddress"][current_area][iprobe] = e["session_id"]

    return {
        "Accuracy": accuracy,
        "Accuracy_shuffeled": accuracy_shuffled,
        "windowCenters": window_centers,
    }

def run_prewhisk_training_decoding_pipeline(
    entries_by_area: Mapping[str, List[Dict[str, Any]]],
    params: Mapping[str, Any],
    pre_bins_zero_based: np.ndarray,
    rng_seed: int = 0,
    area_list: Optional[Mapping[str, Set[str]]] = None,
    enable_ccf_filter: bool = True,
) -> Dict[str, Any]:
    """Run the decoding-from-prewhisk training pipeline.

    Parameters
    ----------
    entries_by_area:
        Probe entries grouped by area.
    params:
        Decoding configuration mapping.
    pre_bins_zero_based:
        Bin indices used to build the prewhisk training representation.
    rng_seed:
        Seed for the shared RNG driving balancing and CV splits.
    area_list:
        Optional area-to-CCF mapping used by the unit filter.
    enable_ccf_filter:
        Whether CCF-based unit filtering is active.

    Returns
    -------
    dict[str, Any]
        Accuracy bundle containing real and shuffled results plus window
        centers.
    """

    def zscore_safe(X: np.ndarray) -> np.ndarray:
        Z = zscore_cols(X)
        return np.nan_to_num(Z, nan=0.0, posinf=0.0, neginf=0.0)

    def build_train_matrix_prewhisk(val_class: np.ndarray, pre_bins: np.ndarray) -> np.ndarray:
        block = val_class[pre_bins, :, :]
        return np.asarray(np.sum(block, axis=0), dtype=np.float32)

    def fit_prewhisk_model(
        X_train: np.ndarray,
        y_train: np.ndarray,
        rng_obj: np.random.Generator,
    ) -> Tuple[Any, np.float32, np.float32]:
        from sklearn.model_selection import StratifiedKFold
        from sklearn.svm import NuSVC

        c1 = int(np.sum(y_train == 1))
        c2 = int(np.sum(y_train == -1))
        if min(c1, c2) < 5:
            return None, np.float32(np.nan), np.float32(np.nan)

        split_seed = int(rng_obj.integers(0, 2**31 - 1))
        skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=split_seed)

        acc = []
        acc_sh = []
        model_prewhisk = None

        try:
            for tr_idx, te_idx in skf.split(X_train, y_train):
                model = NuSVC(nu=0.5, kernel="linear")
                model.fit(X_train[tr_idx, :].astype(np.float64, copy=False), y_train[tr_idx])

                pred = model.predict(X_train[te_idx, :].astype(np.float64, copy=False))
                yy = y_train[te_idx]
                acc.append(100.0 * np.mean(pred == yy))

                yy_sh = yy.copy()
                yy_sh = yy_sh[rng_obj.permutation(yy_sh.size)]
                acc_sh.append(100.0 * np.mean(pred == yy_sh))

                model_prewhisk = model
        except ValueError:
            return None, np.float32(np.nan), np.float32(np.nan)

        return model_prewhisk, np.float32(np.mean(acc)), np.float32(np.mean(acc_sh))

    def eval_model_on_bin(
        model_prewhisk: Any,
        X_bin: np.ndarray,
        y_bin: np.ndarray,
        rng_obj: np.random.Generator,
    ) -> Tuple[np.float32, np.float32]:
        from sklearn.model_selection import StratifiedKFold

        if model_prewhisk is None:
            return np.float32(np.nan), np.float32(np.nan)

        c1 = int(np.sum(y_bin == 1))
        c2 = int(np.sum(y_bin == -1))
        if min(c1, c2) < 5:
            return np.float32(np.nan), np.float32(np.nan)

        split_seed = int(rng_obj.integers(0, 2**31 - 1))
        skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=split_seed)

        acc = []
        acc_sh = []

        try:
            for _, te_idx in skf.split(X_bin, y_bin):
                pred = model_prewhisk.predict(X_bin[te_idx, :].astype(np.float64, copy=False))
                yy = y_bin[te_idx]
                acc.append(100.0 * np.mean(pred == yy))

                yy_sh = yy.copy()
                yy_sh = yy_sh[rng_obj.permutation(yy_sh.size)]
                acc_sh.append(100.0 * np.mean(pred == yy_sh))
        except ValueError:
            return np.float32(np.nan), np.float32(np.nan)

        return np.float32(np.mean(acc)), np.float32(np.mean(acc_sh))

    rng_global = np.random.default_rng(rng_seed)
    pre_bins_zero_based = np.asarray(pre_bins_zero_based, dtype=np.int32).reshape(-1)
    window_centers = np.asarray(params["windowCenters"], dtype=np.float32)

    if pre_bins_zero_based.size == 0:
        raise ValueError("pre_bins_zero_based must not be empty")
    if int(np.max(pre_bins_zero_based)) >= window_centers.size:
        raise ValueError("prewhisk bins out of range for current windowCenters")

    accuracy: Dict[str, Any] = {"sessionaddress": {}}
    accuracy_shuffeled: Dict[str, Any] = {"sessionaddress": {}}

    # Train one prewhisk decoder per probe entry, then evaluate that trained
    # model across all bins for the same probe.
    for current_area in params["regionlist"]:
        probe_entries = list(entries_by_area.get(current_area, []))

        area_acc_list: List[np.ndarray] = []
        area_sh_list: List[np.ndarray] = []
        area_sessions: List[str] = []

        for e in probe_entries:
            trial = np.asarray(e["trial"]).reshape(-1)
            lick = np.asarray(e["lick"]).reshape(-1)

            completed_trials_ind = get_completion_mask(e, str(params["completion_state"]))
            qind = get_quiet_mask(e, str(params["quietstate"]))

            class1 = (trial == 1) & (lick == 1)
            class2 = (trial == 3) & (lick == 0)
            class1 = class1 & qind & completed_trials_ind
            class2 = class2 & qind & completed_trials_ind

            celltype_ind = get_celltype_mask(e, str(params["celltype"]))
            ccf_ind = get_ccf_mask(
                e,
                str(current_area),
                area_list=area_list,
                enable_ccf_filter=enable_ccf_filter,
            )
            curr_cell_ind = celltype_ind & ccf_ind

            if np.sum(curr_cell_ind) < 5:
                continue

            currsig = np.asarray(e["spike_counts"], dtype=np.float32)[:, :, curr_cell_ind]
            val_class1 = currsig[:, class1, :]
            val_class2 = currsig[:, class2, :]

            n1_train = int(np.sum(class1))
            n2_train = int(np.sum(class2))

            X1_train = build_train_matrix_prewhisk(val_class1, pre_bins_zero_based)
            X2_train = build_train_matrix_prewhisk(val_class2, pre_bins_zero_based)

            X1_train = ensure_trials_by_cells(X1_train, n1_train)
            X2_train = ensure_trials_by_cells(X2_train, n2_train)

            if (X1_train.shape[1] < int(params["mintrial"])) or (X2_train.shape[1] < int(params["mintrial"])):
                continue

            X_train, y_train = downsample_balance(X1_train, X2_train, rng_obj=rng_global)
            if params["zscoring"]:
                X_train = zscore_safe(X_train)

            model_prewhisk, _, _ = fit_prewhisk_model(X_train, y_train, rng_obj=rng_global)

            n_bins = currsig.shape[0]
            sess_acc = np.full((n_bins,), np.nan, dtype=np.float32)
            sess_sh = np.full((n_bins,), np.nan, dtype=np.float32)

            for i_bin in range(n_bins):
                X1_bin = np.squeeze(val_class1[i_bin, :, :])
                X2_bin = np.squeeze(val_class2[i_bin, :, :])

                n1 = int(np.sum(class1))
                n2 = int(np.sum(class2))
                X1_bin = ensure_trials_by_cells(X1_bin, n1)
                X2_bin = ensure_trials_by_cells(X2_bin, n2)

                if (X1_bin.shape[1] < int(params["mintrial"])) or (X2_bin.shape[1] < int(params["mintrial"])):
                    continue

                X_bin, y_bin = downsample_balance(X1_bin, X2_bin, rng_obj=rng_global)
                if params["zscoring"]:
                    X_bin = zscore_safe(X_bin)

                acc_val, sh_val = eval_model_on_bin(model_prewhisk, X_bin, y_bin, rng_obj=rng_global)
                sess_acc[i_bin] = acc_val
                sess_sh[i_bin] = sh_val

            area_acc_list.append(sess_acc)
            area_sh_list.append(sess_sh)
            area_sessions.append(str(e["session_id"]))

        if area_acc_list:
            accuracy[current_area] = np.stack(area_acc_list, axis=0)
            accuracy_shuffeled[current_area] = np.stack(area_sh_list, axis=0)
        else:
            n_bins = window_centers.size
            accuracy[current_area] = np.zeros((0, n_bins), dtype=np.float32)
            accuracy_shuffeled[current_area] = np.zeros((0, n_bins), dtype=np.float32)

        accuracy["sessionaddress"][current_area] = np.array(area_sessions, dtype=object)
        accuracy_shuffeled["sessionaddress"][current_area] = np.array(area_sessions, dtype=object)

    return {
        "Accuracy": accuracy,
        "Accuracy_shuffeled": accuracy_shuffeled,
        "windowCenters": window_centers,
    }

def run_delay_decoding_pipeline(
    entries_by_area: Mapping[str, List[Dict[str, Any]]],
    params: Mapping[str, Any],
    area_list: Optional[Mapping[str, Set[str]]] = None,
    enable_ccf_filter: bool = True,
    rng_seed: int = 0,
) -> Dict[str, Any]:
    """Run the delay-window decoding pipeline.

    Parameters
    ----------
    entries_by_area:
        Probe entries grouped by area.
    params:
        Decoding configuration mapping.
    area_list:
        Optional area-to-CCF mapping used by the unit filter.
    enable_ccf_filter:
        Whether CCF-based unit filtering is active.
    rng_seed:
        Seed for the shared RNG driving balancing and CV splits.

    Returns
    -------
    dict[str, Any]
        Accuracy bundle with real and shuffled delay-decoding scores.
    """

    rng = np.random.default_rng(rng_seed)
    window_centers = np.asarray(params["windowCenters"], dtype=np.float32)
    n_bins = int(window_centers.size)
    min_cells = int(params.get("min_cells_per_session", 5))

    accuracy: Dict[str, Any] = {"sessionaddress": {}}
    accuracy_shuffeled: Dict[str, Any] = {"sessionaddress": {}}

    for current_area in params["regionlist"]:
        probe_entries = list(entries_by_area.get(current_area, []))
        n_probes = len(probe_entries)

        area_acc = np.full((n_probes, n_bins), np.nan, dtype=np.float32)
        area_sh = np.full((n_probes, n_bins), np.nan, dtype=np.float32)
        area_sessions = np.full((n_probes,), None, dtype=object)

        for iprobe, e in enumerate(probe_entries):
            trial = np.asarray(e["trial"]).reshape(-1)
            lick = np.asarray(e["lick"]).reshape(-1)

            completed_trials_ind = get_completion_mask(e, str(params["completion_state"]))
            qind = get_quiet_mask(e, str(params["quietstate"]))

            class1 = (trial == 1) & (lick == 1)
            class2 = (trial == 3) & (lick == 0)
            class1 = class1 & qind & completed_trials_ind
            class2 = class2 & qind & completed_trials_ind

            celltype_ind = get_celltype_mask(e, str(params["celltype"]))
            ccf_ind = get_ccf_mask(
                e,
                str(current_area),
                area_list=area_list,
                enable_ccf_filter=enable_ccf_filter,
            )
            curr_cell_ind = celltype_ind & ccf_ind

            if int(np.sum(curr_cell_ind)) < min_cells:
                continue

            currsig = np.asarray(e["spike_counts"], dtype=np.float32)[:, :, curr_cell_ind]
            val_class1 = currsig[:, class1, :]
            val_class2 = currsig[:, class2, :]

            n1 = int(np.sum(class1))
            n2 = int(np.sum(class2))

            area_sessions[iprobe] = e["session_id"]
            if min(n1, n2) < int(params["mintrial"]):
                continue

            for i_bin in range(n_bins):
                X1 = np.squeeze(val_class1[i_bin, :, :])
                X2 = np.squeeze(val_class2[i_bin, :, :])

                X1 = ensure_trials_by_cells(X1, n1)
                X2 = ensure_trials_by_cells(X2, n2)

                acc_val, sh_val = decode_one_bin_svm_matlab_compat(
                    X1,
                    X2,
                    mintrial=int(params["mintrial"]),
                    balance_method=str(params["balance_method"]),
                    zscoring=bool(params["zscoring"]),
                    run_seed=rng_seed,
                    rng_obj=rng,
                )
                area_acc[iprobe, i_bin] = acc_val
                area_sh[iprobe, i_bin] = sh_val

        accuracy[current_area] = area_acc
        accuracy_shuffeled[current_area] = area_sh
        accuracy["sessionaddress"][current_area] = area_sessions
        accuracy_shuffeled["sessionaddress"][current_area] = area_sessions

    return {
        "Accuracy": accuracy,
        "Accuracy_shuffeled": accuracy_shuffeled,
        "windowCenters": window_centers,
    }

def run_afterwhisker_decoding_pipeline(
    entries_by_area: Mapping[str, List[Dict[str, Any]]],
    cfg: Mapping[str, Any],
    run_seed: int,
    area_list: Optional[Mapping[str, Set[str]]] = None,
    enable_ccf_filter: bool = True,
) -> Dict[str, Any]:
    """Run one after-whisker decoding pass for a single run seed.

    Parameters
    ----------
    entries_by_area:
        Probe entries grouped by area.
    cfg:
        Decoding configuration mapping for the after-whisker analysis.
    run_seed:
        Seed for the shared RNG driving balancing and shuffle controls.
    area_list:
        Optional area-to-CCF mapping used by the unit filter.
    enable_ccf_filter:
        Whether CCF-based unit filtering is active.

    Returns
    -------
    dict[str, Any]
        Accuracy bundle containing per-area real and shuffled after-whisker
        decoding scores plus the window centers used for the evaluated slice.
    """

    rng = np.random.default_rng(int(run_seed))
    window_centers = np.asarray(cfg["windowCenters"], dtype=np.float32)
    bin_slice = slice(int(cfg["bin_slice_start"]), int(cfg["bin_slice_stop"]))
    n_bins = int(bin_slice.stop - bin_slice.start)
    n_base = int(np.asarray(cfg["baselinelist"], dtype=np.float32).size)
    min_cells = int(cfg.get("min_cells_per_session", 5))

    accuracy: Dict[str, Any] = {"sessionaddress": {}}
    accuracy_shuffeled: Dict[str, Any] = {"sessionaddress": {}}

    for current_area in cfg["regionlist"]:
        probe_entries = list(entries_by_area.get(current_area, []))

        area_acc_list: List[np.ndarray] = []
        area_shuf_list: List[np.ndarray] = []
        area_sessions: List[str] = []

        for entry in probe_entries:
            trial = np.asarray(entry["trial"]).reshape(-1)
            lick = np.asarray(entry["lick"]).reshape(-1)
            n_trials = int(entry["n_trials"])

            completed_trials_ind = get_completion_mask(entry, str(cfg["completion_state"]))
            if completed_trials_ind.size != n_trials:
                completed_trials_ind = np.ones(n_trials, dtype=bool)

            qind = get_quiet_mask(entry, str(cfg["quietstate"]))
            if qind.size != n_trials:
                qind = np.ones(n_trials, dtype=bool)

            class1 = (trial == 1) & (lick == 1)
            class2 = (trial == 3) & (lick == 0)
            class1 = class1 & qind & completed_trials_ind
            class2 = class2 & qind & completed_trials_ind

            celltype_ind = get_celltype_mask(entry, str(cfg["celltype"]))
            ccf_ind = get_ccf_mask(
                entry,
                str(current_area),
                area_list=area_list,
                enable_ccf_filter=enable_ccf_filter,
            )
            curr_cell_ind = celltype_ind & ccf_ind

            if int(np.sum(curr_cell_ind)) < min_cells:
                continue

            currsig = np.asarray(entry["spike_counts"], dtype=np.float32)[:, :, curr_cell_ind]

            sess_acc = np.full((n_bins, n_base), np.nan, dtype=np.float32)
            sess_shuf = np.full((n_bins,), np.nan, dtype=np.float32)

            for ibase, base_t in enumerate(np.asarray(cfg["baselinelist"], dtype=np.float32)):
                currsig_win = currsig[bin_slice, :, :]

                if bool(cfg["baseline_subtraction"]):
                    bidx = nearest_bin(window_centers, float(base_t))
                    baseline_vals = currsig[bidx : bidx + 1, :, :]
                    currsig_win = currsig_win - baseline_vals

                val_class1 = currsig_win[:, class1, :]
                val_class2 = currsig_win[:, class2, :]

                n1 = int(np.sum(class1))
                n2 = int(np.sum(class2))

                for i_bin in range(n_bins):
                    X1 = np.squeeze(val_class1[i_bin, :, :])
                    X2 = np.squeeze(val_class2[i_bin, :, :])

                    X1 = ensure_trials_by_cells(X1, n1)
                    X2 = ensure_trials_by_cells(X2, n2)

                    # Decoding_afterwhisker_runX.m uses svmtrain('-s 1 -t 0') =
                    # nu-SVC (nu=0.5) linear with the size(X,2)>=mintrial guard, so
                    # use the MATLAB-compatible decoder here (matches the other
                    # decoding pipelines) rather than the generic C-SVC path.
                    acc_val, sh_val = decode_one_bin_svm_matlab_compat(
                        X1,
                        X2,
                        mintrial=int(cfg["mintrial"]),
                        balance_method=str(cfg["balance_method"]),
                        zscoring=bool(cfg["zscoring"]),
                        run_seed=int(run_seed),
                        rng_obj=rng,
                    )
                    sess_acc[i_bin, ibase] = acc_val
                    sess_shuf[i_bin] = sh_val

            area_acc_list.append(sess_acc)
            area_shuf_list.append(sess_shuf)
            area_sessions.append(str(entry.get("session_id", "unknown_session")))

        if area_acc_list:
            accuracy[current_area] = np.stack(area_acc_list, axis=0)
            accuracy_shuffeled[current_area] = np.stack(area_shuf_list, axis=0)
        else:
            accuracy[current_area] = np.zeros((0, n_bins, n_base), dtype=np.float32)
            accuracy_shuffeled[current_area] = np.zeros((0, n_bins), dtype=np.float32)

        accuracy["sessionaddress"][current_area] = np.array(area_sessions, dtype=object)
        accuracy_shuffeled["sessionaddress"][current_area] = np.array(area_sessions, dtype=object)

    return {
        "Accuracy": accuracy,
        "Accuracy_shuffeled": accuracy_shuffeled,
        "windowCenters": np.asarray(window_centers[bin_slice], dtype=np.float32),
    }
