import os, time, torch
os.environ["OMP_NUM_THREADS"]="2"; os.environ["MKL_NUM_THREADS"]="2"
torch.set_num_threads(2)
from diffusers import StableDiffusionPipeline
print("loading fp16 pipe...", flush=True)
pipe = StableDiffusionPipeline.from_pretrained("rv6_diffusers", torch_dtype=torch.float16,
        safety_checker=None, requires_safety_checker=False, disable_mmap=True)
pipe.to("cpu")
latent = torch.randn(1,4,32,32)   # 256px latent
t = torch.tensor([981], dtype=torch.long)
emb = torch.randn(1,77,768)
def bench(mod, name):
    for _ in range(1):
        with torch.no_grad():
            t0=time.time(); mod(latent.to(mod.dtype), t, encoder_hidden_states=emb.to(mod.dtype)); dt=time.time()-t0
    print(f"{name}: {dt:.2f}s/step", flush=True)
bench(pipe.unet, "UNet fp16")
pipe.unet = pipe.unet.to(torch.float32)
latent = latent.to(torch.float32); emb = emb.to(torch.float32)
bench(pipe.unet, "UNet fp32")
print("peak swap now:", flush=True)
import subprocess; print(subprocess.run(["free","-h"],capture_output=True,text=True).stdout)
