"""One instrumented, unchanged Diffusers SD1.5 call for Lecture 04 LAB06.

No model training, paid API, manual replacement of pipeline outputs, or full grid.
Dependencies match the existing image lab. Captures real model calls with hooks.
Run: python run_component_trace.py --source-results PATH --output results
"""
from __future__ import annotations
import argparse
import hashlib
import importlib.metadata
import json
import os
from pathlib import Path
import platform
import time
import traceback
from datetime import datetime, timezone

MODEL = "stable-diffusion-v1-5/stable-diffusion-v1-5"
REVISION = "451f4fe16113bff5a5d2269ed5ad43b0592e9a14"
PROMPT = "A red ceramic mug on a wooden table beside a window, soft daylight, still life photograph"
LATENT_HASH = "5dcb0b9f4d7b65d2458eae198fedc392f08d61d32defd303f465b779fd0945e1"

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

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

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

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--source-results", type=Path, required=True)
    parser.add_argument("--output", type=Path, default=Path("results"))
    parser.add_argument("--model-path", default=MODEL)
    parser.add_argument("--threads", type=int, default=8)
    args = parser.parse_args()
    out = args.output
    out.mkdir(parents=True, exist_ok=True)
    if (out / "trace.json").exists():
        raise RuntimeError("Use a new output directory; do not overwrite a recorded run.")
    record = {"status": "started", "started_at": utc(), "model": MODEL,
              "revision": REVISION, "prompt": PROMPT, "negative_prompt": "",
              "seed": 11, "guidance_scale": 4.0, "guidance_rescale": 0.0,
              "num_inference_steps": 20, "eta": 0.0, "height": 512, "width": 512,
              "method": "Hooks around the original pipeline; all returned values remain unchanged.",
              "text_encoder_calls": [], "steps": [], "vae_encode_call_count": 0}
    write_json(out / "trace.json", record)
    try:
        import torch
        import diffusers
        from diffusers import StableDiffusionPipeline, DDIMScheduler
        torch.set_num_threads(args.threads)
        torch.set_num_interop_threads(1)
        record["environment"] = {"python": platform.python_version(),
            "device": "cpu", "dtype": "bfloat16", "threads": torch.get_num_threads(),
            "packages": {x: importlib.metadata.version(x) for x in
                         ["torch", "diffusers", "transformers", "accelerate", "safetensors", "numpy", "Pillow"]}}
        for name in ["cpu.max", "memory.max"]:
            path = Path("/sys/fs/cgroup") / name
            if path.exists(): record["environment"][name] = path.read_text().strip()
        source = args.source_results / "initial-latent-seed-11.pt"
        if sha(source) != LATENT_HASH:
            raise RuntimeError("The saved seed-11 starting latent does not match the existing lab.")
        initial = torch.load(source, weights_only=True, map_location="cpu")
        if initial.dtype != torch.bfloat16: raise RuntimeError("Expected the saved bfloat16 latent.")
        tensors = {}

        def summary(value):
            cpu = value.detach().cpu().contiguous()
            f = cpu.float()
            return {"shape": list(cpu.shape), "dtype": str(cpu.dtype),
                    "min": float(f.min()), "max": float(f.max()), "mean": float(f.mean()),
                    "std_population": float(f.std(unbiased=False)), "l2_norm": float(f.norm()),
                    "tensor_bytes_sha256": hashlib.sha256(cpu.view(torch.uint8).numpy().tobytes()).hexdigest()}

        def keep(name, value):
            tensors[name] = value.detach().cpu().clone()
            return summary(value)

        record["initial_latent_source_sha256"] = sha(source)
        record["initial_latent"] = keep("initial_latent", initial)
        load_start = time.perf_counter()
        print("Loading fixed pretrained components", flush=True)
        pipe = StableDiffusionPipeline.from_pretrained(
            args.model_path, revision=REVISION if args.model_path == MODEL else None,
            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()
        pipe.set_progress_bar_config(disable=True)
        record["load_seconds"] = time.perf_counter() - load_start
        record["scheduler"] = {"class": type(pipe.scheduler).__name__, "config": dict(pipe.scheduler.config)}
        record["vae_scaling_factor"] = float(pipe.vae.config.scaling_factor)
        pending = {}

        def text_hook(module, args_, kwargs, result):
            ids = args_[0] if args_ else kwargs["input_ids"]
            index = len(record["text_encoder_calls"])
            branch = "conditional" if index == 0 else "empty_text"
            entry = {"branch": branch, "token_ids": ids.cpu().tolist(),
                     "tokens": [pipe.tokenizer.convert_ids_to_tokens(row) for row in ids.cpu().tolist()],
                     "input": keep(f"text_{branch}_ids", ids),
                     "output": keep(f"text_{branch}_embedding", result[0])}
            record["text_encoder_calls"].append(entry)

        def unet_hook(module, args_, kwargs, result):
            sample, timestep = args_[0], args_[1]
            prediction = result[0]
            uncond, cond = prediction.chunk(2)
            index = len(record["steps"])
            pending.clear()
            pending.update(index=index, timestep=int(timestep),
                           uncond=uncond.detach().clone(), cond=cond.detach().clone())
            if index == 0:
                record["first_unet_input"] = keep("step0_unet_input", sample)
                record["first_unet_text_input"] = keep("step0_unet_text_input", kwargs["encoder_hidden_states"])
                record["first_unet_branches_use_identical_latent"] = bool(torch.equal(sample[0], sample[1]))
                keep("step0_empty_prediction", uncond)
                keep("step0_conditional_prediction", cond)

        original_step = pipe.scheduler.step
        def traced_step(model_output, timestep, sample, *rest, **kwargs):
            before = sample.detach().clone()
            result = original_step(model_output, timestep, sample, *rest, **kwargs)
            after = result[0]
            expected = pending["uncond"] + record["guidance_scale"] * (pending["cond"] - pending["uncond"])
            cfg_exact = torch.equal(expected, model_output)
            step = {"index": pending["index"], "timestep": int(timestep),
                    "sample_before": summary(before), "empty_prediction": summary(pending["uncond"]),
                    "conditional_prediction": summary(pending["cond"]),
                    "guided_prediction": summary(model_output), "sample_after": summary(after),
                    "cfg_matches_original_arithmetic_exactly": bool(cfg_exact),
                    "cfg_max_abs_difference": float((expected.float()-model_output.float()).abs().max()),
                    "update_l2_norm": float((after.float()-before.float()).norm())}
            if not cfg_exact: raise AssertionError("CFG capture does not match the original pipeline output")
            if pending["index"] == 0:
                keep("step0_latent_before", before)
                keep("step0_guided_prediction", model_output)
                keep("step0_latent_after", after)
                # Stable, disclosed coordinates; not selected to exaggerate a difference.
                coordinates = [[0, 0, 0, 0], [0, 0, 16, 16], [0, 1, 32, 32], [0, 2, 48, 48], [0, 3, 63, 63]]
                step["coordinates"] = [{"index": c, "latent_before": float(before[tuple(c)]),
                    "empty_prediction": float(pending["uncond"][tuple(c)]),
                    "conditional_prediction": float(pending["cond"][tuple(c)]),
                    "guided_prediction": float(model_output[tuple(c)]),
                    "latent_after": float(after[tuple(c)])} for c in coordinates]
            record["steps"].append(step)
            if len(record["steps"]) % 5 == 0:
                print("Recorded", len(record["steps"]), "scheduler updates", flush=True)
            return result

        original_decode = pipe.vae.decode
        def traced_decode(z, *rest, **kwargs):
            record["vae_decode_input"] = keep("vae_decode_input", z)
            result = original_decode(z, *rest, **kwargs)
            output = result[0] if isinstance(result, tuple) else result.sample
            record["vae_decoded_tensor"] = keep("vae_decoded_tensor", output)
            return result

        original_encode = pipe.vae.encode
        def traced_encode(*args_, **kwargs):
            record["vae_encode_call_count"] += 1
            return original_encode(*args_, **kwargs)

        text_handle = pipe.text_encoder.register_forward_hook(text_hook, with_kwargs=True)
        unet_handle = pipe.unet.register_forward_hook(unet_hook, with_kwargs=True)
        pipe.scheduler.step = traced_step
        pipe.vae.decode = traced_decode
        pipe.vae.encode = traced_encode

        def callback(pipeline, step, timestep, callback_kwargs):
            if step == 19: record["final_latent"] = keep("final_latent", callback_kwargs["latents"])
            return callback_kwargs

        print("Starting one unchanged 20-step pipeline inference", flush=True)
        start = time.perf_counter()
        with torch.inference_mode():
            result = pipe(prompt=PROMPT, negative_prompt="", height=512, width=512,
                          num_inference_steps=20, guidance_scale=4.0, guidance_rescale=0.0,
                          eta=0.0, latents=initial.clone(), generator=torch.Generator("cpu").manual_seed(11),
                          callback_on_step_end=callback)
        record["inference_seconds"] = time.perf_counter() - start
        text_handle.remove()
        unet_handle.remove()
        result.images[0].save(out / "component-trace-final.png")
        torch.save(tensors, out / "component-tensors.pt")
        record["scheduler_timesteps"] = pipe.scheduler.timesteps.cpu().tolist()
        record["safety_checker_flags"] = result.nsfw_content_detected
        record["final_image"] = {"file": "component-trace-final.png", "sha256": sha(out / "component-trace-final.png")}
        record["tensor_file"] = {"file": "component-tensors.pt", "sha256": sha(out / "component-tensors.pt")}
        old = args.source_results / "txt-baseline-seed-11.png"
        if old.exists():
            from PIL import Image
            import numpy as np
            old_array = np.array(Image.open(old)).astype("float32")
            new_array = np.array(result.images[0]).astype("float32")
            record["historical_image_comparison"] = {"historical_sha256": sha(old),
                "identical_png_sha256": sha(old) == record["final_image"]["sha256"],
                "mean_absolute_pixel_difference_0_to_255": float(abs(old_array-new_array).mean()),
                "max_absolute_pixel_difference_0_to_255": float(abs(old_array-new_array).max())}
        record.update(status="success", completed_at=utc())
        write_json(out / "trace.json", record)
        print("SUCCESS", record["inference_seconds"], "seconds", flush=True)
    except Exception as exc:
        record.update(status="failed", completed_at=utc(), error_type=type(exc).__name__, error=str(exc))
        write_json(out / "trace.json", record)
        traceback.print_exc()
        raise

if __name__ == "__main__":
    main()
