"""Write a synthetic group of drift-diffusion participants to fit.

Draws each participant's parameters from a population, simulates their trials, and writes the
stacked table `hierarchical_fitting.py` expects::

    python make_example_data.py --n-participants 6 --n-trials 40
    python hierarchical_fitting.py --data group_data.csv

Nothing here is needed to fit your own data; it exists so the example has something to run on.
The model is imported from `hierarchical_fitting` rather than repeated, so the data are
generated by the same composition that is later fitted.
"""

import argparse

import numpy as np
import pandas as pd

from hierarchical_fitting import FIT_PARAMS, FIT_RANGES, build_model, trial_inputs
from psyneulink.core.compositions.hierarchical.transforms import BoundedTransform

# Population the participants are drawn from, in the unconstrained space the group model uses:
# centred on the middle of each range, with this variance between participants.
GROUP_MEAN_Z = np.zeros(len(FIT_PARAMS))
GROUP_VARIANCE_Z = np.full(len(FIT_PARAMS), 0.36)


def simulate_participant(theta, n_trials, seed):
    """Generate one participant's trials at known parameters."""
    comp, _ = build_model(rate=float(theta[0]), threshold=float(theta[1]), seed=seed)
    comp.run(inputs={comp.nodes[0]: trial_inputs(n_trials)}, context=f"simulate-{seed}")
    data = pd.DataFrame(np.squeeze(np.array(comp.results)),
                        columns=["decision", "response_time"])
    data["decision"] = data["decision"].astype("category")
    return data


def make_group_data(n_participants, n_trials, seed=0):
    """Draw participants from the population and simulate each one.

    Returns the stacked table and the parameters it was generated from.
    """
    transform = BoundedTransform(
        lower=[FIT_RANGES[p][0] for p in FIT_PARAMS],
        upper=[FIT_RANGES[p][1] for p in FIT_PARAMS],
    )
    rng = np.random.default_rng(seed)
    z_true = rng.normal(GROUP_MEAN_Z, np.sqrt(GROUP_VARIANCE_Z),
                        size=(n_participants, len(FIT_PARAMS)))
    theta_true = np.vstack([transform.to_natural(z) for z in z_true])

    frames = []
    for s in range(n_participants):
        frame = simulate_participant(theta_true[s], n_trials, seed=1000 + s)
        frame.insert(0, "subject", f"S{s:02d}")
        frames.append(frame)
    return pd.concat(frames, ignore_index=True), theta_true


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--n-participants", type=int, default=6)
    parser.add_argument("--n-trials", type=int, default=40)
    parser.add_argument("--seed", type=int, default=0)
    parser.add_argument("--data-out", default="group_data.csv")
    args = parser.parse_args()

    data, theta_true = make_group_data(args.n_participants, args.n_trials, seed=args.seed)
    data.to_csv(args.data_out, index=False)
    print(f"wrote {len(data)} trials from {args.n_participants} participants "
          f"to {args.data_out}")
    print(f"drawn around mean parameters: {theta_true.mean(axis=0).round(4)}")


if __name__ == "__main__":
    main()
