"""avatar-E2E-inference 추론 가중치를 GitHub Release용 번들로 묶는다.

가중치는 .gitignore로 git에서 제외되므로(*.pth *.pt *.npy *.npz *.pkl *.vrm)
fresh clone에는 코드만 온다. 이 스크립트가 만드는 3개 tar를 릴리스 에셋으로 올리고,
새 머신에서는 fetch-weights.sh가 그것을 받아 저장소 상대경로 그대로 푼다.

flow 체크포인트는 optimizer_state_dict를 벗겨서 담는다(620MB -> 210MB). 추론 경로는
`ck.get("model_state_dict", ck)`로 읽으므로 코드 수정 없이 동작한다. 대신 번들의
best.pth로는 **학습 재개가 불가능**하다 — 원본은 학습 머신에 보존해야 한다.
VQ-VAE는 encoder(추론 미사용, 4개 합 139MB)까지 벗길 수 있지만 loader를
strict=False로 완화해야 해서 원본 그대로 담는다.

    python pack_weights.py                      # -> dist/ 에 3개 tar + SHA256SUMS
    python pack_weights.py --out /tmp/bundles
"""
import argparse
import hashlib
import os
import shutil
import subprocess
import sys
import tempfile

HERE = os.path.dirname(os.path.abspath(__file__))
ROOT = os.path.abspath(os.path.join(HERE, "..", ".."))            # 저장소 루트
WH = "package/motion-blender/experiments/wave-hands"
F3 = "package/face/animasync-face-v3"

# (번들 이름, [저장소 상대경로]) — 경로 구조를 그대로 tar에 담아 루트에서 풀면 제자리로 간다.
BUNDLES = {
    "kemix-body-v4-vel": [
        f"{WH}/outputs/gen_train/vqvae_kemix_v4_vel/best_spine.pth",
        f"{WH}/outputs/gen_train/vqvae_kemix_v4_vel/best_arms.pth",
        f"{WH}/outputs/gen_train/vqvae_kemix_v4_vel/best_legs.pth",
        f"{WH}/outputs/gen_train/vqvae_kemix_v4_vel/best_fingers.pth",
        f"{WH}/outputs/kemix_npz_v4/mean_std",
        f"{WH}/outputs/gen_train/seed_bank/seed_avg.npz",
        f"{WH}/outputs/kemix_rest_table.npy",
        f"{WH}/outputs/dummy_lang/weights/vocab.pkl",
    ],
    "kemix-face-v3": [
        f"{F3}/models/v3_face/checkpoints/best_expression_v14.pt",
    ],
    "kemix-avatar": [
        f"{F3}/avatar/GG_11.vrm",
    ],
}
# 슬림화해서 body 번들에 넣는 항목: (원본, 번들 내 경로)
SLIM_FLOW = (f"{WH}/outputs/gen_train/weights_v4_vel/best.pth",
             f"{WH}/outputs/gen_train/weights_v4_vel/best.pth")


def sha256(path, chunk=1 << 20):
    h = hashlib.sha256()
    with open(path, "rb") as f:
        for block in iter(lambda: f.read(chunk), b""):
            h.update(block)
    return h.hexdigest()


def stage_slim_flow(stage):
    """optimizer_state_dict를 뺀 flow 체크포인트를 staging에 쓴다."""
    import torch                                   # 이 스크립트에서만 필요
    src, rel = SLIM_FLOW
    src_abs = os.path.join(ROOT, src)
    dst = os.path.join(stage, rel)
    os.makedirs(os.path.dirname(dst), exist_ok=True)
    ck = torch.load(src_abs, map_location="cpu", weights_only=False)
    if "model_state_dict" not in ck:
        sys.exit(f"{src}: model_state_dict 키가 없다 — 포맷이 바뀌었는지 확인할 것")
    kept = {"epoch": ck.get("epoch"), "model_state_dict": ck["model_state_dict"]}
    torch.save(kept, dst)
    before, after = os.path.getsize(src_abs), os.path.getsize(dst)
    print(f"  slim  {rel}\n        {before/1e6:.1f}MB -> {after/1e6:.1f}MB "
          f"(optimizer {(before-after)/1e6:.0f}MB 제거)")
    return rel


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--out", default=os.path.join(HERE, "dist"))
    a = ap.parse_args()
    os.makedirs(a.out, exist_ok=True)

    missing = [p for paths in BUNDLES.values() for p in paths
               if not os.path.exists(os.path.join(ROOT, p))]
    if not os.path.exists(os.path.join(ROOT, SLIM_FLOW[0])):
        missing.append(SLIM_FLOW[0])
    if missing:
        sys.exit("원본 파일 없음:\n  " + "\n  ".join(missing))

    with tempfile.TemporaryDirectory(prefix="kemix-pack-") as stage:
        print("staging…")
        slim_rel = stage_slim_flow(stage)
        for paths in BUNDLES.values():
            for rel in paths:
                src, dst = os.path.join(ROOT, rel), os.path.join(stage, rel)
                os.makedirs(os.path.dirname(dst), exist_ok=True)
                (shutil.copytree if os.path.isdir(src) else shutil.copy2)(src, dst)

        print("\n번들 생성…")
        artifacts = []
        for name, paths in BUNDLES.items():
            members = list(paths)
            if name == "kemix-body-v4-vel":
                members.insert(0, slim_rel)
            tar = os.path.join(a.out, f"{name}.tgz")
            # .pth는 이미 zip 컨테이너라 gzip 이득이 8% 뿐 — 최저 압축으로 시간만 아낀다.
            subprocess.run(["tar", "-I", "gzip -1", "-cf", tar, "-C", stage, *members],
                           check=True)
            artifacts.append(tar)
            print(f"  {os.path.basename(tar):28s} {os.path.getsize(tar)/1e6:8.1f} MB")

    sums = os.path.join(a.out, "SHA256SUMS")
    with open(sums, "w") as f:
        for tar in artifacts:
            f.write(f"{sha256(tar)}  {os.path.basename(tar)}\n")
    print(f"\n-> {a.out}/  ({len(artifacts)} bundles + SHA256SUMS)")
    print(open(sums).read(), end="")


if __name__ == "__main__":
    main()
