#!/usr/bin/env python3
"""Convert RV6 fp16 single-file ckpt + ft-mse vae -> diffusers fp16 directory."""
import os
import shutil
import time

os.environ.setdefault("OMP_NUM_THREADS", "2")
os.environ.setdefault("MKL_NUM_THREADS", "2")

import torch
from safetensors.torch import load_file
from transformers import CLIPTextConfig, CLIPTextModel, CLIPTokenizer
from diffusers import (
    AutoencoderKL,
    EulerDiscreteScheduler,
    StableDiffusionPipeline,
    UNet2DConditionModel,
)
from diffusers.loaders.single_file_utils import (
    convert_ldm_clip_checkpoint,
    convert_ldm_unet_checkpoint,
)

BASE = "/data/onlyfake/v3"
CKPT = os.path.join(BASE, "weights", "Realistic_Vision_V6.0_NV_B1_fp16.safetensors")
VAE = os.path.join(BASE, "weights", "sd-vae-ft-mse.safetensors")
CFG = os.path.join(BASE, "weights", "cfg")
OUT = os.path.join(BASE, "rv6_diffusers")


def load_cfg(name):
    import json
    with open(os.path.join(CFG, name)) as f:
        return json.load(f)


def main():
    torch.set_num_threads(2)
    t0 = time.time()

    if os.path.isdir(OUT):
        shutil.rmtree(OUT)
    os.makedirs(OUT)

    print("[conv] loading ckpt", flush=True)
    sd = load_file(CKPT)
    print("[conv] ckpt keys=%d dtype=%s (%.0fs)" % (len(sd), sd[list(sd)[0]].dtype, time.time() - t0), flush=True)

    print("[conv] unet", flush=True)
    unet_cfg = load_cfg("unet/config.json")
    unet = UNet2DConditionModel(**unet_cfg)
    unet_sd = convert_ldm_unet_checkpoint(sd, unet_cfg)
    miss, unexpected = unet.load_state_dict(unet_sd, strict=False, assign=True)
    print("  unet missing=%d unexpected=%d" % (len(miss), len(unexpected)), flush=True)
    assert len(miss) == 0, miss[:8]
    unet.requires_grad_(False)
    del unet_sd

    print("[conv] text_encoder", flush=True)
    te_cfg = load_cfg("text_encoder/config.json")
    te_sd = convert_ldm_clip_checkpoint(sd)
    if any(k.startswith("text_model.") for k in te_sd):
        te_sd = {k.replace("text_model.", "", 1) if k.startswith("text_model.") else k: v
                 for k, v in te_sd.items()}
    te = CLIPTextModel(CLIPTextConfig(**te_cfg))
    miss, unexpected = te.load_state_dict(te_sd, strict=False, assign=True)
    print("  text_encoder missing=%d unexpected=%d" % (len(miss), len(unexpected)), flush=True)
    assert len(miss) == 0, miss[:8]
    te.requires_grad_(False)
    del te_sd

    print("[conv] vae", flush=True)
    vae_cfg = load_cfg("vae/config.json")
    vae = AutoencoderKL(**vae_cfg)
    vae_sd = load_file(VAE)
    RENAME = {
        ".query.weight": ".to_q.weight",
        ".query.bias": ".to_q.bias",
        ".key.weight": ".to_k.weight",
        ".key.bias": ".to_k.bias",
        ".value.weight": ".to_v.weight",
        ".value.bias": ".to_v.bias",
        ".proj_attn.weight": ".to_out.0.weight",
        ".proj_attn.bias": ".to_out.0.bias",
    }
    renamed = {}
    for k, v in vae_sd.items():
        nk = k
        if ".attentions." in nk:
            for a, b in RENAME.items():
                nk = nk.replace(a, b)
        renamed[nk] = v
    vae_sd = renamed
    miss, unexpected = vae.load_state_dict(vae_sd, strict=False, assign=True)
    print("  vae missing=%d unexpected=%d" % (len(miss), len(unexpected)), flush=True)
    assert len(miss) == 0, miss[:8]
    vae.requires_grad_(False)
    del vae_sd

    print("[conv] tokenizer + scheduler", flush=True)
    tokenizer = CLIPTokenizer.from_pretrained(os.path.join(CFG, "tokenizer"))
    scheduler = EulerDiscreteScheduler.from_config(load_cfg("scheduler/scheduler_config.json"))

    print("[conv] assemble + save", flush=True)
    pipe = StableDiffusionPipeline(
        unet=unet,
        vae=vae,
        text_encoder=te,
        tokenizer=tokenizer,
        scheduler=scheduler,
        safety_checker=None,
        feature_extractor=None,
        requires_safety_checker=False,
    )
    pipe.save_pretrained(OUT, safe_serialization=True)
    del sd, unet, vae, te
    print("[conv] DONE -> %s (%.0fs)" % (OUT, time.time() - t0), flush=True)


if __name__ == "__main__":
    main()