Buckets:
| """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}} | |
Xet Storage Details
- Size:
- 1.64 kB
- Xet hash:
- bb2b6661e9ee5f5cfc67328d59b8a468d424001110472c19b43088f1c94adade
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.