File size: 10,994 Bytes
c5da886 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 | """dexp010_controller.py — exp010: THE STEP-GATED DIFFUSION CONTROLLER at
inference + the generation-side band-lesion battery (the instrument upgrade).
The controller (stage plan phase 3): during sampling, each denoising step
sets the band windows from the CURRENT timestep — band experts activate
coarse-to-fine across the trajectory exactly as trained. This file ships the
controller as a usable module (StepGatedSampler) and judges what the bands DO
IN IMAGE SPACE, where per-step eps-MSE may be gauge-blind (exp009 verdict).
Arms (trained mb3 stack, exp008 s0): all_on | lesion_band0 (LOW-noise/detail)
| lesion_band1 (MID) | lesion_band2 (HIGH-noise/structure) | frozen.
Image-space gauges (deterministic, paired seeds):
G1 round-trip: CLIP-L cond-vs-shuffled per arm (does lesioning band b
change grounding?).
G2 structure-vs-detail decomposition of the lesion effect: for the SAME
(prompt, seed), LP-MSE(lesioned img, all_on img) = coarse-structure
change; HP-MSE = fine-detail change. PREREG: lesion_band2 (HIGH) moves
LP most (coarse/blob role); lesion_band0 (LOW) moves HP most (detail
role); band1 intermediate. This is the image-space test of the
coarse-to-fine thesis.
G3 within-prompt diversity: 6 prompts x 4 seeds, mean pairwise CLIP
distance per arm — PREREG: lesion_band2 reduces diversity (the HIGH
band's diversity role).
Pod: bash pod2/run_exp010.sh
"""
from __future__ import annotations
import io
import json
import os
import sys
sys.path[:0] = ["pod2", "."]
import numpy as np
import torch
import torch.nn.functional as F
from pod_ledger import ledger_run, note, burn_down
from d1_substrate import MEM_FRACTION
from aleph_diffusion_core import derangement
from dexp008_multiband import (MultibandDelta, band_weights, attach,
load_unet, N_BANDS)
SD_BASE = "stable-diffusion-v1-5/stable-diffusion-v1-5"
DEXP8_CKPT = ("/workspace/ckpts2/dexp008" if os.path.isdir("/workspace")
else "./data/dexp008")
N_JUDGE, STEPS_GEN, GUIDANCE = 24, 30, 7.5
CKPT_FILE = os.environ.get("DEXP10_CKPT", "mb3_s0.pt")
CKPT_DIR_ENV = os.environ.get("DEXP10_CKPT_DIR", "")
TAG = os.environ.get("DEXP10_TAG", "mb3_s0")
N_DIV_PROMPTS, N_DIV_SEEDS = 6, 4
SEED = 1234
DATA_DIR = ("/workspace/data/dexp010" if os.path.isdir("/workspace")
else os.path.join(os.environ.get("GEOLIP_DATA", "./data"),
"dexp010"))
class StepGatedSampler:
"""THE CONTROLLER: DDIM sampling with per-step band gating — each step
sets the crossfade windows from the current timestep, activating band
experts coarse-to-fine across the trajectory."""
def __init__(self, unet, wraps, device):
from diffusers import DDIMScheduler
self.unet, self.wraps, self.device = unet, wraps, device
self.sched = DDIMScheduler.from_pretrained(SD_BASE,
subfolder="scheduler")
@torch.no_grad()
def sample(self, ehs_cond, seed, steps=STEPS_GEN, guidance=GUIDANCE):
B = ehs_cond.shape[0]
self.sched.set_timesteps(steps, device=self.device)
g = torch.Generator(device=self.device).manual_seed(seed)
x = torch.randn(B, 4, 64, 64, generator=g, device=self.device)
x = x * self.sched.init_noise_sigma
ehs_in = torch.cat([ehs_cond, torch.zeros_like(ehs_cond)], dim=0)
for t in self.sched.timesteps:
s01 = torch.full((2 * B,), float(t) / 1000.0, device=self.device)
w = band_weights(s01)
for wr in self.wraps:
wr.w_bands = w # the step gate
xin = self.sched.scale_model_input(torch.cat([x, x]), t)
eps = self.unet(xin, t, ehs_in, return_dict=False)[0]
e_c, e_u = eps.chunk(2)
eps = e_u + guidance * (e_c - e_u)
x = self.sched.step(eps, t, x).prev_sample
return x
def hp_img(x):
return x - F.avg_pool2d(x, 3, stride=1, padding=1)
def lp_img(x):
return F.avg_pool2d(x, 7, stride=1, padding=3)
def run(device="cuda"):
torch.cuda.set_per_process_memory_fraction(MEM_FRACTION, 0)
os.makedirs(DATA_DIR, exist_ok=True)
from PIL import Image
from diffusers import AutoencoderKL
from transformers import (CLIPTextModel, CLIPTokenizer, CLIPModel,
CLIPProcessor)
from d1_exp000b_natural import load_rows, encode_nl77
from d1_exp000_baselines import judge_selftest
rows = load_rows(N_JUDGE)
prompts = [r["prompt"] for r in rows]
perm = derangement(N_JUDGE, seed=SEED)
tok = CLIPTokenizer.from_pretrained(SD_BASE, subfolder="tokenizer")
te = CLIPTextModel.from_pretrained(
SD_BASE, subfolder="text_encoder",
torch_dtype=torch.float32).to(device).eval()
vae = AutoencoderKL.from_pretrained(
SD_BASE, subfolder="vae", torch_dtype=torch.float32).to(device).eval()
clip = CLIPModel.from_pretrained(
"openai/clip-vit-large-patch14",
torch_dtype=torch.float32).to(device).eval()
cproc = CLIPProcessor.from_pretrained("openai/clip-vit-large-patch14")
feat = judge_selftest(clip, cproc, device)
with torch.no_grad():
ehs_c = encode_nl77(prompts, tok, te, device)
ehs_s = encode_nl77([prompts[i] for i in perm.tolist()], tok, te,
device)
orig = torch.cat([
feat(Image.open(io.BytesIO(r["image_bytes"])).convert("RGB"))
for r in rows])
unet = load_unet(device)
mods, wraps = attach(unet, lambda d: MultibandDelta(d))
ck_dir = CKPT_DIR_ENV or DEXP8_CKPT
ck = torch.load(os.path.join(ck_dir, CKPT_FILE),
map_location="cpu", weights_only=True)
for m, sd in zip(mods, ck["mods"]):
m.load_state_dict(sd, strict=True)
sampler = StepGatedSampler(unet, wraps, device)
def set_arm(arm):
for m in mods:
m.enabled = arm != "frozen"
m.band_enabled = [True] * N_BANDS
if arm.startswith("lesion_band"):
m.band_enabled[int(arm[-1])] = False
@torch.no_grad()
def decode_imgs(lat):
img = vae.decode(lat / 0.18215).sample
return ((img / 2 + 0.5).clamp(0, 1))
arms = ["all_on", "lesion_band0", "lesion_band1", "lesion_band2",
"frozen"]
results = {"config": {"steps": STEPS_GEN, "guidance": GUIDANCE,
"seed": SEED, "ckpt": TAG,
"sampler": "DDIM step-gated"}}
gen_store = {}
with ledger_run("dexp010 controller battery", budget_h=1.5) as h:
for arm in arms:
set_arm(arm)
imgs_all, cos = [], []
for i in range(0, N_JUDGE, 6):
lat = sampler.sample(ehs_c[i:i + 6], seed=SEED + i)
px = decode_imgs(lat)
imgs_all.append(px.cpu())
pil = [Image.fromarray(
(p.permute(1, 2, 0).numpy() * 255).astype(np.uint8))
for p in px.cpu()]
f = torch.cat([feat(p) for p in pil])
cos += (f * orig[i:i + 6]).sum(-1).tolist()
cond_mean = sum(cos) / len(cos)
cos_s = []
for i in range(0, N_JUDGE, 6):
lat = sampler.sample(ehs_s[i:i + 6], seed=SEED + i)
px = decode_imgs(lat)
pil = [Image.fromarray(
(p.permute(1, 2, 0).numpy() * 255).astype(np.uint8))
for p in px.cpu()]
f = torch.cat([feat(p) for p in pil])
cos_s += (f * orig[i:i + 6]).sum(-1).tolist()
gen_store[arm] = torch.cat(imgs_all)
results[arm] = {
"round_trip": {"cond": round(cond_mean, 4),
"shuffled": round(sum(cos_s) / len(cos_s), 4),
"gap": round(cond_mean
- sum(cos_s) / len(cos_s), 4)}}
print(f"[exp010] {arm}: rt gap "
f"{results[arm]['round_trip']['gap']:+.4f}", flush=True)
# G2: structure-vs-detail decomposition of each lesion (vs all_on)
base_imgs = gen_store["all_on"]
for arm in ("lesion_band0", "lesion_band1", "lesion_band2",
"frozen"):
d = gen_store[arm] - base_imgs
results[arm]["vs_all_on"] = {
"lp_change": round(float((lp_img(gen_store[arm])
- lp_img(base_imgs)).pow(2).mean()),
6),
"hp_change": round(float((hp_img(gen_store[arm])
- hp_img(base_imgs)).pow(2).mean()),
6),
"total_change": round(float(d.pow(2).mean()), 6)}
# G3: within-prompt diversity (6 prompts x 4 seeds)
for arm in arms:
set_arm(arm)
div = []
for pi in range(N_DIV_PROMPTS):
fs = []
for si in range(N_DIV_SEEDS):
lat = sampler.sample(ehs_c[pi:pi + 1],
seed=SEED + 1000 + 97 * si)
px = decode_imgs(lat).cpu()[0]
pil = Image.fromarray(
(px.permute(1, 2, 0).numpy() * 255).astype(np.uint8))
fs.append(feat(pil))
fs = torch.cat(fs)
sim = fs @ fs.T
n = fs.shape[0]
div.append(1 - ((sim.sum() - n) / (n * (n - 1))).item())
results[arm]["diversity"] = round(sum(div) / len(div), 5)
h["verdict"] = json.dumps(
{a: results[a]["round_trip"]["gap"] for a in arms})
# prereg checks
lb0 = results["lesion_band0"]["vs_all_on"]
lb2 = results["lesion_band2"]["vs_all_on"]
results["prereg"] = {
"G2_high_band_moves_LP_most":
lb2["lp_change"] > lb0["lp_change"],
"G2_low_band_moves_HP_most":
lb0["hp_change"] > lb2["hp_change"],
"G3_high_lesion_cuts_diversity":
results["lesion_band2"]["diversity"]
< results["all_on"]["diversity"],
"note": "image-space coarse-to-fine test; 1-seed battery",
}
with open(os.path.join(DATA_DIR, "results.json" if TAG == "mb3_s0" else f"results_{TAG}.json"), "w") as f:
json.dump(results, f, indent=2)
note(f"dexp010: {json.dumps(results['prereg'])}")
print(json.dumps(results["prereg"], indent=2))
burn_down()
return results
def smoke():
x = torch.randn(2, 3, 32, 32)
assert hp_img(x).shape == x.shape and lp_img(x).shape == x.shape
w = band_weights(torch.tensor([0.02, 0.5, 0.98]))
assert w[0].argmax() == 0 and w[2].argmax() == 2
print("dexp010 smoke PASSED (filters + windows; GPU run is pod work)")
if __name__ == "__main__":
if "--run" in sys.argv:
run()
else:
smoke()
|