"""第 16—19 课：可恢复、有独立验证题的训练循环。"""
import json
import math
import os
import random
import resource
import platform
import time
from dataclasses import asdict
from pathlib import Path

import torch

from .common import device_for, read_json, read_jsonl, write_json
from .data import Batcher, Tokenizer
from .model import ModelConfig, PoetryGPT


def seed_all(seed):
    random.seed(seed)
    torch.manual_seed(seed)
    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)


@torch.no_grad()
def evaluate_loss(model, records, tokenizer, batch_size, batches, device):
    was_training = model.training
    model.eval()
    batcher = Batcher(records, tokenizer, model.config.context, 413)
    total, count = 0.0, 0
    for _ in range(batches):
        x, y = batcher.batch(batch_size, device)
        _, loss = model(x, y)
        n = int((y != -100).sum().item())
        total += loss.item() * n
        count += n
    model.train(was_training)
    return total / count


def save_checkpoint(path, payload):
    path = Path(path)
    temp = path.with_suffix(".tmp")
    torch.save(payload, temp)
    os.replace(temp, path)


def load_model(path, device):
    path = Path(path)
    if not path.is_file():
        raise ValueError(f"找不到模型：{path}。先运行训练，或选择已交付的 artifacts/runs/poet/best.pt。")
    payload = torch.load(path, map_location="cpu", weights_only=True)
    model = PoetryGPT(ModelConfig(**payload["model_config"]))
    model.load_state_dict(payload["model"])
    model.to(device).eval()
    return model, Tokenizer(payload["vocabulary"]), payload


def train(data_dir, run_dir, config_path, device_name="auto", steps=None, batch_size=None, resume=None):
    data_dir, run_dir = Path(data_dir), Path(run_dir)
    config = read_json(config_path)
    settings = config["training"].copy()
    manifest = read_json(data_dir / "manifest.json")
    vocabulary = read_json(data_dir / "vocab.json")
    checkpoint = None
    if resume:
        checkpoint = torch.load(resume, map_location="cpu", weights_only=True)
        if "optimizer" not in checkpoint: raise ValueError("继续训练请选 latest.pt；best.pt 用于写诗。")
        if checkpoint["data_fingerprint"] != manifest["split_sha256"] or checkpoint["vocabulary"] != vocabulary:
            raise ValueError("继续训练需要同一份材料和字表；检测到变化。")
        settings = checkpoint["training_config"].copy()
    elif (run_dir / "latest.pt").exists():
        raise ValueError("此训练目录已存在进度。请使用 --resume，或换一个 --run-dir。")
    if steps is not None: settings["steps"] = steps
    if batch_size is not None: settings["batch_size"] = batch_size
    if settings["steps"] < 1 or settings["batch_size"] < 1: raise ValueError("步数和每批数量必须大于 0。")
    device = device_for(device_name)
    if device.type == "cpu": torch.set_num_threads(min(int(os.environ.get("POETRY_CPU_THREADS", "8")), os.cpu_count() or 1))
    seed_all(settings["seed"])
    tokenizer = Tokenizer(vocabulary)
    model_config = ModelConfig(**checkpoint["model_config"]) if checkpoint else ModelConfig(vocab_size=len(vocabulary), **config["model"])
    model = PoetryGPT(model_config).to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=settings["learning_rate"], weight_decay=0.1)
    train_records, val_records = read_jsonl(data_dir / "train.jsonl"), read_jsonl(data_dir / "val.jsonl")
    batcher = Batcher(train_records, tokenizer, model_config.context, settings["seed"])
    completed, best = 0, float("inf")
    prior_seconds = 0.0
    if checkpoint:
        model.load_state_dict(checkpoint["model"])
        optimizer.load_state_dict(checkpoint["optimizer"])
        completed, best = checkpoint["step"], checkpoint["best_val"]
        if settings["steps"] <= completed: raise ValueError(f"已完成 {completed} 步，目标步数必须更大。")
        batcher.random.setstate(checkpoint["batch_random_state"])
        torch.set_rng_state(checkpoint["cpu_rng"])
        if device.type == "mps" and checkpoint.get("mps_rng") is not None:
            torch.mps.set_rng_state(checkpoint["mps_rng"])
        if device.type == "cuda" and checkpoint.get("cuda_rng") is not None:
            torch.cuda.set_rng_state_all(checkpoint["cuda_rng"])
        prior_seconds = checkpoint.get("elapsed_seconds", 0)
    run_dir.mkdir(parents=True, exist_ok=True)
    started = time.monotonic()
    started_step = completed
    write_json(run_dir / "run.json", {"status": "running", "step": completed,
                "target_steps": settings["steps"], "parameters": model.parameter_count(), "device": str(device)})
    print(json.dumps({"device": str(device), "parameters": model.parameter_count(), "start_step": completed,
                      "target_steps": settings["steps"], "train_poems": len(train_records)}, ensure_ascii=False), flush=True)

    def payload(full=False):
        result = {"format_version": 1, "model": model.state_dict(), "model_config": asdict(model_config),
                  "vocabulary": vocabulary, "step": completed, "best_val": best,
                  "training_config": settings, "data_fingerprint": manifest["split_sha256"],
                  "elapsed_seconds": prior_seconds + time.monotonic()-started,
                  "training_device": str(device), "torch_version": str(torch.__version__)}
        if full:
            result.update(optimizer=optimizer.state_dict(), batch_random_state=batcher.random.getstate(),
                          cpu_rng=torch.get_rng_state(),
                          mps_rng=torch.mps.get_rng_state() if device.type == "mps" else None,
                          cuda_rng=torch.cuda.get_rng_state_all() if device.type == "cuda" else None)
        return result

    def log(event):
        with (run_dir / "metrics.jsonl").open("a", encoding="utf-8") as stream:
            stream.write(json.dumps(event, ensure_ascii=False) + "\n")
        print(json.dumps(event, ensure_ascii=False), flush=True)

    def memory():
        peak = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
        report = {"peak_process_rss_mib": round(peak / (1024**2 if platform.system() == "Darwin" else 1024), 1)}
        if device.type == "mps":
            report["mps_driver_allocated_mib"] = round(torch.mps.driver_allocated_memory()/1024**2, 1)
        return report

    if not checkpoint:
        baseline = evaluate_loss(model, val_records, tokenizer, settings["batch_size"], settings["eval_batches"], device)
        log({"event": "baseline", "step": 0, "val_loss": baseline})
        save_checkpoint(run_dir / "initial.pt", payload())
    interrupted = False
    model.train()
    try:
        for step in range(completed + 1, settings["steps"] + 1):
            warmup = min(100, max(1, settings["steps"] // 10))
            progress = max(0, step-warmup) / max(1, settings["steps"]-warmup)
            factor = step/warmup if step < warmup else 0.15 + 0.85*0.5*(1+math.cos(math.pi*progress))
            lr = settings["learning_rate"] * factor
            for group in optimizer.param_groups: group["lr"] = lr
            x, y = batcher.batch(settings["batch_size"], device)
            optimizer.zero_grad(set_to_none=True)
            _, loss = model(x, y)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            optimizer.step()
            completed = step
            if step % 25 == 0 or step == 1:
                loss_value = loss.item()
                if not math.isfinite(loss_value): raise ValueError("训练出现非有限数值，请降低调整幅度并检查材料。")
                elapsed = time.monotonic() - started
                log({"event": "train", "step": step, "loss": loss_value, "lr": lr,
                     "seconds": round(prior_seconds+elapsed, 2),
                     "steps_per_second": round((step-started_step)/max(elapsed, 0.001), 3)})
            if step % settings["eval_every"] == 0 or step == settings["steps"]:
                val = evaluate_loss(model, val_records, tokenizer, settings["batch_size"], settings["eval_batches"], device)
                if val < best:
                    best = val
                    save_checkpoint(run_dir / "best.pt", payload())
                save_checkpoint(run_dir / "latest.pt", payload(True))
                log({"event": "validation", "step": step, "val_loss": val, "best_val": best, **memory()})
    except KeyboardInterrupt:
        interrupted = True
        print("收到中断，正在保存已完成的训练进度。", flush=True)
    finally:
        save_checkpoint(run_dir / "latest.pt", payload(True))
        write_json(run_dir / "run.json", {"status": "interrupted" if interrupted else "finished" if completed == settings["steps"] else "failed",
                    "step": completed, "parameters": model.parameter_count(), "device": str(device),
                    "seconds": round(prior_seconds+time.monotonic()-started, 2),
                    "model": asdict(model_config), "training": settings,
                    "best_val": best if math.isfinite(best) else None, **memory()})
    return {"run_dir": str(run_dir), "step": completed, "interrupted": interrupted}
