"""Lecture 04 LAB06: bounded real SD1.5 experiment, no training, no paid API.

Run: python run_image_lab.py --phase pilot|text|edit|inpaint|mask-followup|all
Existing completed configurations are loaded from disk; nothing is silently retried.
The raw initial latent is reused across prompts / CFG / step counts for each seed.
"""
from __future__ import annotations
import argparse, gc, hashlib, importlib, importlib.metadata, json, os, platform, time, traceback
from datetime import datetime, timezone
from pathlib import Path

ROOT = Path(__file__).resolve().parent
OUT = ROOT / "results"
OUT.mkdir(exist_ok=True)
MODEL = "stable-diffusion-v1-5/stable-diffusion-v1-5"
REVISION = "451f4fe16113bff5a5d2269ed5ad43b0592e9a14"
INPAINT_MODEL = "stable-diffusion-v1-5/stable-diffusion-inpainting"
INPAINT_REVISION = "8a4288a76071f7280aedbdb3253bdb9e9d5d84bb"
PROMPT = "A red ceramic mug on a wooden table beside a window, soft daylight, still life photograph"
PROMPT_BLUE = PROMPT.replace("red", "blue")
SEEDS = [11, 23, 37]
SIZE = 512
BASE_CFG = 4.0
BASE_STEPS = 20

def utc():
    return datetime.now(timezone.utc).isoformat()

def digest(path):
    return hashlib.sha256(Path(path).read_bytes()).hexdigest()

def save_json(path, value):
    Path(path).write_text(json.dumps(value, ensure_ascii=False, indent=2), encoding="utf-8")

def record_event(kind, **kwargs):
    with (OUT / "events.jsonl").open("a", encoding="utf-8") as f:
        f.write(json.dumps(dict(time=utc(), kind=kind, **kwargs), ensure_ascii=False) + "\n")

def environment():
    import torch
    result = dict(time=utc(), python=platform.python_version(), platform=platform.platform(),
                  cpu_count=os.cpu_count(), device="cpu", dtype="bfloat16", threads=torch.get_num_threads(),
                  cuda_available=torch.cuda.is_available(), packages={}, distribution_metadata={})
    for package,module in [(x,x) for x in ["torch", "diffusers", "transformers", "accelerate", "huggingface_hub", "safetensors", "numpy"]]+[("Pillow","PIL")]:
        try:
            loaded=importlib.import_module(module)
            result["packages"][package] = getattr(loaded,"__version__",importlib.metadata.version(package))
            result["distribution_metadata"][package] = importlib.metadata.version(package)
        except importlib.metadata.PackageNotFoundError: pass
    for name in ["cpu.max", "memory.max"]:
        p = Path("/sys/fs/cgroup") / name
        if p.exists(): result[name] = p.read_text().strip()
    if Path("/proc/cpuinfo").exists():
        lines = Path("/proc/cpuinfo").read_text().splitlines()
        result["cpu_model"] = next((x.split(":",1)[1].strip() for x in lines if x.startswith("model name")), "unknown")
        flags = next((x for x in lines if x.startswith("flags")), "")
        result["avx512_bf16"] = "avx512_bf16" in flags
    return result

def load_text_pipeline():
    import torch
    from diffusers import StableDiffusionPipeline, DDIMScheduler
    print("Loading pretrained SD1.5", flush=True)
    started = time.perf_counter()
    pipe = StableDiffusionPipeline.from_pretrained(
        MODEL, revision=REVISION, variant="fp16", torch_dtype=torch.bfloat16,
        use_safetensors=True, low_cpu_mem_usage=True,
    ).to("cpu")
    pipe.scheduler = DDIMScheduler.from_config(pipe.scheduler.config)
    pipe.enable_vae_slicing()
    # PyTorch SDPA is retained; attention slicing would remove its fast CPU path.
    save_json(OUT / "scheduler.json", dict(pipe.scheduler.config))
    record_event("model_loaded", model=MODEL, revision=REVISION, seconds=time.perf_counter()-started)
    return pipe

def latent(seed, pipe):
    import torch
    path = OUT / f"initial-latent-seed-{seed}.pt"
    if path.exists():
        return torch.load(path, weights_only=True)
    # Draw float32 on CPU first; store exactly the bfloat16 latent passed to the pipeline.
    value = torch.randn((1, pipe.unet.config.in_channels, SIZE//8, SIZE//8),
                        generator=torch.Generator("cpu").manual_seed(seed), dtype=torch.float32).to(torch.bfloat16)
    torch.save(value, path)
    return value

def completed(key):
    path = OUT / f"{key}.json"
    return path.exists() and json.loads(path.read_text()).get("status") == "success" and (OUT/f"{key}.png").exists()

def infer(pipe, key, spec, **inputs):
    import torch
    if completed(key):
        saved=json.loads((OUT/f"{key}.json").read_text())
        expected=dict(spec,width=SIZE,height=SIZE,device="cpu",dtype="bfloat16",
                      scheduler=type(pipe.scheduler).__name__,eta=0.0)
        changed=[field for field,value in expected.items() if saved.get(field)!=value]
        if changed or digest(OUT/f"{key}.png")!=saved.get("sha256"):
            raise RuntimeError(f"Cached output {key} does not match this configuration ({changed}). Use a new results directory; do not silently replace the original experiment.")
        print("Reuse", key, flush=True)
        return saved
    record = dict(spec, key=key, status="running", started_at=utc(), width=SIZE, height=SIZE,
                  device="cpu", dtype="bfloat16", scheduler=type(pipe.scheduler).__name__, eta=0.0)
    save_json(OUT/f"{key}.json", record)
    t0=time.perf_counter()
    def callback(pipeline, step, timestep, callback_kwargs):
        if step % 5 == 0: print(key, "step", step+1, "elapsed", round(time.perf_counter()-t0,1), flush=True)
        return callback_kwargs
    try:
        with torch.inference_mode():
            result = pipe(**inputs, callback_on_step_end=callback)
        image = result.images[0]
        image.save(OUT/f"{key}.png")
        record.update(status="success", completed_at=utc(), seconds=time.perf_counter()-t0,
                      output=f"{key}.png", sha256=digest(OUT/f"{key}.png"),
                      safety_checker_flag=result.nsfw_content_detected[0] if result.nsfw_content_detected else None)
        # Safety checker remains enabled. A flagged/black output is retained with its flag.
        save_json(OUT/f"{key}.json",record)
        print("SUCCESS", key, round(record["seconds"],1), "seconds", flush=True)
    except Exception as exc:
        record.update(status="failed", completed_at=utc(), seconds=time.perf_counter()-t0,
                      error_type=type(exc).__name__, error=str(exc))
        save_json(OUT/f"{key}.json",record)
        record_event("inference_failed", **record)
        raise
    return record

def run_text(pipe, pilot=False):
    for seed in SEEDS:
        configs = [("baseline",PROMPT,BASE_CFG,BASE_STEPS)]
        if not pilot:
            configs += [("prompt-blue",PROMPT_BLUE,BASE_CFG,BASE_STEPS)]
            configs += [(f"cfg-{cfg}",PROMPT,float(cfg),BASE_STEPS) for cfg in [1,8]]
            configs += [(f"steps-{steps}",PROMPT,BASE_CFG,steps) for steps in [10,40]]
        for condition,prompt,cfg,steps in configs:
            key=f"txt-{condition}-seed-{seed}"
            latent_path=OUT/f"initial-latent-seed-{seed}.pt"
            initial=latent(seed,pipe)
            infer(pipe,key,dict(task="txt2img", condition=condition, model=MODEL, revision=REVISION,
                               prompt=prompt, negative_prompt="", seed=seed, guidance_scale=cfg,
                               num_inference_steps=steps, initial_latent=latent_path.name,
                               initial_latent_sha256=digest(latent_path)),
                  prompt=prompt, negative_prompt="", height=SIZE,width=SIZE,
                  num_inference_steps=steps,guidance_scale=cfg,eta=0.0,
                  latents=initial.clone(),generator=__import__("torch").Generator("cpu").manual_seed(seed))
        if pilot: break

def run_edit(pipe):
    import torch
    from PIL import Image
    from diffusers import StableDiffusionImg2ImgPipeline, DDIMScheduler
    imgpipe = StableDiffusionImg2ImgPipeline(**pipe.components)
    imgpipe.scheduler=DDIMScheduler.from_config(pipe.scheduler.config)
    path=OUT/"txt-baseline-seed-11.png"
    source=Image.open(path).convert("RGB")
    for strength in [0.25,0.55,0.85]:
        key=f"img2img-strength-{strength:.2f}-seed-11"
        infer(imgpipe,key,dict(task="img2img",model=MODEL,revision=REVISION,prompt=PROMPT_BLUE,
                              negative_prompt="",seed=11,guidance_scale=4.0,num_inference_steps=20,
                              strength=strength,source=path.name,source_sha256=digest(path),
                              effective_steps=int(20*strength)),
              prompt=PROMPT_BLUE,negative_prompt="",image=source,strength=strength,
              num_inference_steps=20,guidance_scale=4.0,eta=0.0,
              generator=torch.Generator("cpu").manual_seed(11))

def run_inpaint(rectangle=(160,160,352,400), key="inpaint-fixed-mask-seed-11", follow_up_of=None):
    import torch
    from PIL import Image, ImageDraw
    from diffusers import StableDiffusionInpaintPipeline, DDIMScheduler
    source_path=OUT/"txt-baseline-seed-11.png"
    source=Image.open(source_path).convert("RGB")
    # Original case: predeclared central rectangle. Follow-up: hand-set object rectangle.
    mask=Image.new("L",(SIZE,SIZE),0)
    ImageDraw.Draw(mask).rectangle(rectangle, fill=255)
    mask_path=OUT/("inpainting-mask.png" if follow_up_of is None else "inpainting-object-mask.png"); mask.save(mask_path)
    pipe=StableDiffusionInpaintPipeline.from_pretrained(
        INPAINT_MODEL,revision=INPAINT_REVISION,variant="fp16",torch_dtype=torch.bfloat16,
        use_safetensors=True,low_cpu_mem_usage=True,
    ).to("cpu")
    pipe.scheduler=DDIMScheduler.from_config(pipe.scheduler.config)
    pipe.enable_vae_slicing()
    infer(pipe,key,dict(task="inpainting",model=INPAINT_MODEL,revision=INPAINT_REVISION,
                        prompt=PROMPT_BLUE,negative_prompt="",seed=11,guidance_scale=4.0,
                        num_inference_steps=20,strength=1.0,source=source_path.name,
                        source_sha256=digest(source_path),mask=mask_path.name,
                        mask_sha256=digest(mask_path),mask_rectangle=list(rectangle),
                        follow_up_of=follow_up_of,
                        post_composited=False,algorithm_note="StableDiffusionInpaintPipeline, not RePaint"),
          prompt=PROMPT_BLUE,negative_prompt="",image=source,mask_image=mask,
          height=SIZE,width=SIZE,num_inference_steps=20,guidance_scale=4.0,strength=1.0,eta=0.0,
          generator=torch.Generator("cpu").manual_seed(11))

def main():
    parser=argparse.ArgumentParser();parser.add_argument("--phase",choices=["pilot","text","edit","inpaint","mask-followup","all"],default="pilot")
    args=parser.parse_args()
    import torch
    torch.set_num_threads(min(8,os.cpu_count() or 1))
    torch.set_num_interop_threads(1)
    save_json(OUT/"environment.json",environment())
    record_event("start",phase=args.phase)
    if args.phase in ["pilot","text","edit","all"]:
        pipe=load_text_pipeline()
        if args.phase in ["pilot","text","all"]:run_text(pipe,pilot=args.phase=="pilot")
        if args.phase in ["edit","all"]:run_edit(pipe)
        del pipe;gc.collect()
    if args.phase in ["inpaint","all"]:run_inpaint()
    if args.phase in ["mask-followup","all"]:
        run_inpaint(rectangle=(24,272,208,472),key="inpaint-object-mask-seed-11",follow_up_of="inpaint-fixed-mask-seed-11")
    record_event("complete",phase=args.phase)

if __name__=="__main__":
    main()
