"""Load our checkpoints from either the original .pt files or the .safetensors releases, returning the same dict layout the scripts expect.
  head    : {"model": state_dict, "cfg": {...}}
  NAR LoRA: {"lora": [A0,B0,...] (layer-major: nar_self_attn q,k,v,o then nar_mlp gate,up,down), "io": {"vae2llm": sd, "llm2vae": sd}, "rank": int}
  AR LoRA : {"lora": [A0,B0,...] (layer-major: self_attn q,k,v,o then mlp gate,up,down), "rank": int}"""
import torch
def _mods(prefix): return [(f"{prefix}self_attn",n) for n in ("q_proj","k_proj","v_proj","o_proj")]+[(f"{prefix}mlp",n) for n in ("gate_proj","up_proj","down_proj")]
def load_ckpt(path, map_location="cpu"):
    if not str(path).endswith(".safetensors"): return torch.load(path, map_location=map_location, weights_only=False)
    from safetensors.torch import load_file; from safetensors import safe_open
    t=load_file(path, device=str(map_location))
    with safe_open(path, "pt") as f: meta=f.metadata() or {}
    if any(k.endswith(".lora_A") for k in t):
        prefix="nar_" if any(".nar_self_attn." in k for k in t) else ""
        layers=sorted({int(k.split(".")[1]) for k in t if k.startswith("layers.")}); lora=[]
        for L in layers:
            for blk,proj in _mods(prefix): lora+=[t[f"layers.{L}.{blk}.{proj}.lora_A"], t[f"layers.{L}.{blk}.{proj}.lora_B"]]
        out={"lora":lora, "rank":int(meta.get("rank", lora[0].shape[0]))}
        if prefix: out["io"]={m:{k.split(".",1)[1]:v for k,v in t.items() if k.startswith(m+".")} for m in ("vae2llm","llm2vae")}
        return out
    return {"model":t, "cfg":{"instnorm": meta.get("input","").find("instnorm=true")>=0}}
