"""Recompute trace relationships from saved real tensors, without model inference."""
import argparse, hashlib, json
from pathlib import Path
import torch
from diffusers import DDIMScheduler

p = argparse.ArgumentParser()
p.add_argument("results", type=Path)
args = p.parse_args()
r = json.loads((args.results / "trace.json").read_text())
t = torch.load(args.results / "component-tensors.pt", weights_only=True, map_location="cpu")
s = DDIMScheduler.from_config(r["scheduler"]["config"])
s.set_timesteps(r["num_inference_steps"])
with torch.inference_mode():
    recombined = t["step0_empty_prediction"] + r["guidance_scale"] * (t["step0_conditional_prediction"] - t["step0_empty_prediction"])
    updated = s.step(t["step0_guided_prediction"], r["steps"][0]["timestep"], t["step0_latent_before"], eta=r["eta"]).prev_sample
checks = {
    "status_success": r["status"] == "success",
    "all_20_real_steps_recorded": len(r["steps"]) == r["num_inference_steps"] == 20,
    "two_actual_text_encoder_calls": len(r["text_encoder_calls"]) == 2,
    "text_branches_are_different": not torch.equal(t["text_empty_text_embedding"], t["text_conditional_embedding"]),
    "unet_empty_branch_uses_real_empty_embedding": torch.equal(t["step0_unet_text_input"][:1], t["text_empty_text_embedding"]),
    "unet_cond_branch_uses_real_cond_embedding": torch.equal(t["step0_unet_text_input"][1:], t["text_conditional_embedding"]),
    "same_initial_latent_on_both_unet_branches": torch.equal(t["step0_unet_input"][:1], t["step0_unet_input"][1:]),
    "initial_latent_reused_unchanged": torch.equal(t["initial_latent"], t["step0_latent_before"]),
    "cfg_recombination_exact": torch.equal(recombined, t["step0_guided_prediction"]),
    "scheduler_update_replayed_exactly": torch.equal(updated, t["step0_latent_after"]),
    "all_steps_cfg_match_original_pipeline": all(x["cfg_matches_original_arithmetic_exactly"] for x in r["steps"]),
    "step_chain_is_contiguous": all(a["sample_after"]["tensor_bytes_sha256"] == b["sample_before"]["tensor_bytes_sha256"] for a, b in zip(r["steps"], r["steps"][1:])),
    "final_latent_matches_last_step": r["final_latent"]["tensor_bytes_sha256"] == r["steps"][-1]["sample_after"]["tensor_bytes_sha256"],
    "vae_receives_scaled_final_latent": torch.equal(t["final_latent"] / r["vae_scaling_factor"], t["vae_decode_input"]),
    "no_input_photo_vae_encoding": r["vae_encode_call_count"] == 0,
    "image_hash_matches_record": hashlib.sha256((args.results / r["final_image"]["file"]).read_bytes()).hexdigest() == r["final_image"]["sha256"],
    "tensor_hash_matches_record": hashlib.sha256((args.results / r["tensor_file"]["file"]).read_bytes()).hexdigest() == r["tensor_file"]["sha256"],
}
report = {"all_passed": all(checks.values()), "checks": checks,
          "verification_method": "Load recorded tensors; replay CFG and the official DDIM step. No new model call."}
(args.results / "verification.json").write_text(json.dumps(report, indent=2))
print(json.dumps(report, indent=2))
if not report["all_passed"]: raise SystemExit(1)
