"""
Run the compute-heavy sliding-window image decoder used by FigureS14_2.

The defaults reproduce the notebook-era sustained-cluster
``*_throughomissionRS_sustained.npy`` configuration, with explicit paths for
public reruns.
"""

import argparse
from pathlib import Path
from typing import Optional, Sequence

import h5py
import numpy as np
import pandas as pd

import _bootstrap  # noqa: F401
import decoding_utils as du


DEFAULT_REGIONS = ("VISall",)
DEFAULT_UNIT_SAMPLE_SIZES = (40,)
DEFAULT_CLUSTER = (2, 4)


def parse_cluster(values: Sequence[str]) -> Sequence:
    if len(values) == 1 and values[0] == "all":
        return ("all",)
    return tuple(int(value) for value in values)


def image_decoding_sliding_window(
    session_id: int,
    active_tensor_file: Path,
    stim_table_file: Path,
    unit_table_file: Path,
    output_dir: Path,
    regions: Sequence[str] = DEFAULT_REGIONS,
    unit_sample_sizes: Sequence[int] = DEFAULT_UNIT_SAMPLE_SIZES,
    decode_window_end: int = 1500,
    decode_window_size: int = 50,
    decode_window_bin_size: int = 10,
    decode_window_sliding_step: int = 20,
    use_nonchange: bool = False,
    through_omission: bool = True,
    class_weight: Optional[str] = None,
    rs: bool = True,
    cluster: Sequence = DEFAULT_CLUSTER,
    cluster_name: str = "sustained",
    random_seed: Optional[int] = None,
) -> dict:
    if random_seed is not None:
        np.random.seed(random_seed)

    stim_table = pd.read_csv(stim_table_file)
    unit_table = pd.read_csv(unit_table_file)
    output_dir.mkdir(parents=True, exist_ok=True)

    with h5py.File(active_tensor_file, "r") as unit_data:
        decoded = du.decodeImageSlidingWindow(
            session_id,
            unit_table,
            unit_data,
            stim_table,
            tuple(regions),
            list(unit_sample_sizes),
            decode_window_end,
            decode_window_size,
            decode_window_bin_size,
            decode_window_sliding_step,
            use_nonchange=use_nonchange,
            through_omission=through_omission,
            class_weight=class_weight,
            rs=rs,
            cluster=list(cluster),
        )

    np.save(output_dir / f"{session_id}_throughomissionRS_{cluster_name}.npy", decoded)
    return decoded


def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--session-id", type=int, required=True)
    parser.add_argument("--active-tensor-file", type=Path, required=True)
    parser.add_argument("--stim-table-file", type=Path, required=True)
    parser.add_argument("--unit-table-file", type=Path, required=True)
    parser.add_argument("--output-dir", type=Path, required=True)
    parser.add_argument("--regions", nargs="+", default=list(DEFAULT_REGIONS))
    parser.add_argument("--unit-sample-sizes", nargs="+", type=int, default=list(DEFAULT_UNIT_SAMPLE_SIZES))
    parser.add_argument("--decode-window-end", type=int, default=1500)
    parser.add_argument("--decode-window-size", type=int, default=50)
    parser.add_argument("--decode-window-bin-size", type=int, default=10)
    parser.add_argument("--decode-window-sliding-step", type=int, default=20)
    parser.add_argument("--use-nonchange", action="store_true")
    parser.add_argument("--no-through-omission", action="store_true")
    parser.add_argument("--class-weight", default=None)
    parser.add_argument("--include-fs", action="store_true", help="Use RS+FS units instead of RS-only units.")
    parser.add_argument("--cluster", nargs="+", default=[str(value) for value in DEFAULT_CLUSTER])
    parser.add_argument("--cluster-name", default="sustained")
    parser.add_argument(
        "--random-seed",
        type=int,
        default=None,
        help="Optional seed for unit subsampling. Defaults to notebook-like unseeded sampling.",
    )
    return parser


def main(argv: Optional[Sequence[str]] = None) -> None:
    args = build_parser().parse_args(argv)
    image_decoding_sliding_window(
        session_id=args.session_id,
        active_tensor_file=args.active_tensor_file,
        stim_table_file=args.stim_table_file,
        unit_table_file=args.unit_table_file,
        output_dir=args.output_dir,
        regions=args.regions,
        unit_sample_sizes=args.unit_sample_sizes,
        decode_window_end=args.decode_window_end,
        decode_window_size=args.decode_window_size,
        decode_window_bin_size=args.decode_window_bin_size,
        decode_window_sliding_step=args.decode_window_sliding_step,
        use_nonchange=args.use_nonchange,
        through_omission=not args.no_through_omission,
        class_weight=args.class_weight,
        rs=not args.include_fs,
        cluster=parse_cluster(args.cluster),
        cluster_name=args.cluster_name,
        random_seed=args.random_seed,
    )


if __name__ == "__main__":
    main()
