"""Small MLX experiment: train adapters or evaluate generated SQL answers."""
import argparse
import json
import os
from pathlib import Path
import time
import types


def demonstrations(records, count):
    """Three reviewed training examples: lookup, sum, then two filters.

    Selected on development data before any final-test model evaluations.
    IDs refer to the pinned WikiSQL training source, never the validation/test set.
    """
    if not count:
        return []
    by_source_row = {r["source_row"]:r for r in records}
    return [by_source_row[i] for i in (41941, 17406, 42163)]


def main():
    p = argparse.ArgumentParser()
    p.add_argument("mode", choices=["train", "eval"])
    p.add_argument("--model", required=True)
    p.add_argument("--data", type=Path, required=True)
    p.add_argument("--source", type=Path, required=True)
    p.add_argument("--out", type=Path, required=True)
    p.add_argument("--steps", type=int, default=2000)
    p.add_argument("--batch-size", type=int, default=1)
    p.add_argument("--split", choices=["valid", "test"], default="valid")
    p.add_argument("--limit", type=int, default=0)
    p.add_argument("--adapter")
    p.add_argument("--few-shot", type=int, choices=[0, 3], default=0)
    args = p.parse_args()
    if args.out.exists() and any(args.out.iterdir()):
        p.error("Choose an empty output directory; previous evidence must not be overwritten")
    args.out.mkdir(parents=True, exist_ok=True)
    import mlx.core as mx
    # Set allocator limits before importing/loading the model. GPU and CPU share RAM.
    budget = float(os.environ.get("WIKISQL_MLX_LIMIT_GIB", "3.0"))
    mx.set_memory_limit(int(budget * 1024**3))
    mx.set_cache_limit(128 * 1024**2)
    mx.set_wired_limit(int(budget * 1024**3))
    from mlx_lm import load, stream_generate
    from mlx_lm.sample_utils import make_sampler
    from mlx.utils import tree_flatten
    import numpy as np
    from sql_task import score

    np.random.seed(20260929)
    mx.random.seed(20260929)
    started = time.monotonic()
    model, tokenizer = load(args.model, adapter_path=args.adapter, tokenizer_config={"trust_remote_code": False})
    mx.eval(model.parameters())
    print("MODEL_LOADED", json.dumps({"seconds": time.monotonic()-started, "mlx_active_bytes": mx.get_active_memory(), "mlx_peak_bytes": mx.get_peak_memory()}), flush=True)
    summary = {"mode": args.mode, "model": args.model, "adapter": args.adapter, "few_shot":args.few_shot, "seed": 20260929, "mlx_limit_gib": budget, "parameters": sum(v.size for _,v in tree_flatten(model.parameters()))}
    if args.mode == "train":
        from functools import partial
        import mlx_lm.lora as lora_module
        from mlx_lm.tuner.trainer import train
        from loss import response_loss
        # Narrow local override; installed packages are left untouched. The same
        # corrected loss is used for training and all validation measurements.
        lora_module.train = partial(train, loss=response_loss)
        from mlx_lm.lora import CONFIG_DEFAULTS, train_model
        from mlx_lm.tuner.datasets import load_dataset
        from mlx_lm.tuner.callbacks import TrainingCallback
        config = dict(CONFIG_DEFAULTS)
        config.update(model=args.model, data=str(args.data), train=True, seed=20260929,
                      num_layers=16, batch_size=args.batch_size, iters=args.steps,
                      learning_rate=1e-4, mask_prompt=True, max_seq_length=512,
                      val_batches=-1, steps_per_eval=200, steps_per_report=10,
                      save_every=200, grad_checkpoint=True, adapter_path=str(args.out),
                      lora_parameters={"rank":8,"dropout":0.0,"scale":16.0})
        namespace = types.SimpleNamespace(**config)
        training, valid, _ = load_dataset(namespace, tokenizer)
        class Recorder(TrainingCallback):
            def write(self, kind, info):
                with (args.out / "metrics.jsonl").open("a") as file:
                    file.write(json.dumps({"kind":kind, **info, "elapsed_seconds":time.monotonic()-started}) + "\n")
            def on_train_loss_report(self, info):
                self.write("train", info)
            def on_val_loss_report(self, info):
                self.write("validation", info)
        train_model(namespace, model, training, valid, Recorder())
        summary.update(steps=args.steps, batch_size=args.batch_size,
                       trainable_parameters=sum(v.size for _,v in tree_flatten(model.trainable_parameters())),
                       adapter_bytes=(args.out/"adapters.safetensors").stat().st_size)
    else:
        records = [json.loads(l) for l in (args.data/f"{args.split}.records.jsonl").read_text().splitlines()]
        if args.limit:
            records = records[:args.limit]
        results = []
        examples = demonstrations([json.loads(l) for l in (args.data/"train.records.jsonl").read_text().splitlines()], args.few_shot)
        summary["few_shot_source_rows"] = [r["source_row"] for r in examples]
        model.eval()
        with (args.out/"predictions.jsonl").open("w") as output:
            for i, record in enumerate(records, 1):
                chat = [record["messages"][0]]
                for example in examples:
                    chat.extend([example["messages"][1], {"role":"assistant", "content":example["gold_sql"]}])
                chat.append(record["messages"][1])
                prompt = tokenizer.apply_chat_template(chat, add_generation_prompt=True, tokenize=False)
                text, final = "", None
                tic = time.monotonic()
                for chunk in stream_generate(model, tokenizer, prompt, max_tokens=128, sampler=make_sampler(temp=0)):
                    text += chunk.text
                    final = chunk
                result = {**record, "evaluation_messages":chat, "raw_output":text, "seconds":time.monotonic()-tic,
                          "generation_tokens":final.generation_tokens if final else 0,
                          "finish_reason":final.finish_reason if final else None,
                          **score(text, record, args.source/f"{record['source_split']}.db")}
                output.write(json.dumps(result, ensure_ascii=False)+"\n")
                output.flush()
                results.append(result)
                mx.clear_cache()
                if i % 10 == 0 or i == len(records):
                    print(f"EVAL {i}/{len(records)} correct={sum(r['answer_correct'] for r in results)} elapsed={time.monotonic()-started:.1f}s", flush=True)
        summary.update(split=args.split, count=len(results), valid_sql=sum(r["valid_sql"] for r in results),
                       answer_correct=sum(r["answer_correct"] for r in results), exact_sql=sum(r["exact_sql"] for r in results),
                       answer_correct_without_fence=sum(r["answer_correct_without_fence"] for r in results),
                       empty_gold_answers=sum(not r["gold_answer"] for r in results),
                       nonempty_correct=sum(r["answer_correct"] and bool(r["gold_answer"]) for r in results),
                       generation_tokens=sum(r["generation_tokens"] for r in results))
        summary["generation_limit_hits"] = sum(r["finish_reason"] == "length" for r in results)
    summary.update(elapsed_seconds=time.monotonic()-started, mlx_peak_bytes=mx.get_peak_memory())
    (args.out/"summary.json").write_text(json.dumps(summary, indent=2))
    print("SUMMARY", json.dumps(summary), flush=True)


if __name__ == "__main__":
    main()
