loglens-learnability / simreal /train_composite.py
resoajoe's picture
Upload folder using huggingface_hub
021d920 verified
Raw
History Blame Contribute Delete
11.8 kB
import os
"""Idea B — TriClock-Real: synthetic change events composited onto REAL cam0
footage. Infinite exact labels + real nuisance statistics (the two things the
transfer autopsy said were missing, each patched by the other's data source).
Episode: query time t in the pre-gold window; memory slots at log wall-clock
ages use the REAL frames at those times (so slots contain real person motion,
lighting drift, real scene changes = nuisance). Three synthetic objects are
drawn on top; each may change once, age log-uniform in [1s, 24h]:
lamp: color toggle agent: position swap
furniture: appears/disappears (25% never present)
Label per object: none / fast(<=60s) / med(<=1h) / slow(<=24h).
The model must date the synthetic changes while IGNORING everything real.
python3 -m simreal.train_composite --steps 6000 # pretrain + heldout eval
python3 -m simreal.train_composite --steps 6000 --finetune-real # + realdata finetune/gold
"""
import argparse, bisect, json, os, time
import numpy as np
import torch
import torch.nn.functional as F
from triclock.model import LogLensNet
class ViTEncoder(torch.nn.Module):
"""~8px patches -> linear embed -> 2-layer transformer -> CLS. Nano ViT."""
def __init__(self, d=64):
super().__init__()
self.proj = torch.nn.Linear(8 * 8 * 3, d)
self.cls = torch.nn.Parameter(torch.zeros(1, 1, d))
self.pos = torch.nn.Parameter(torch.zeros(1, 65, d))
layer = torch.nn.TransformerEncoderLayer(d, 4, d * 4, batch_first=True,
dropout=0.0, norm_first=True)
self.tr = torch.nn.TransformerEncoder(layer, 2)
def forward(self, x): # x [N,3,64,64]
N = x.shape[0]
p = x.unfold(2, 8, 8).unfold(3, 8, 8) # [N,3,8,8,8,8]
p = p.permute(0, 2, 3, 1, 4, 5).reshape(N, 64, -1)
z = self.proj(p) + self.pos[:, 1:]
z = torch.cat([self.cls.expand(N, -1, -1) + self.pos[:, :1], z], 1)
return self.tr(z)[:, 0]
HORIZON_S = 86_400
FAST_S, MED_S = 60, 3_600
BUCKETS = ["none", "fast", "med", "slow"]
def log_ages(K=12, T=HORIZON_S):
b = T ** (1.0 / (K - 1))
return sorted({max(1, int(round(b ** i))) for i in range(K)})
def bucket(age_s):
if age_s is None or age_s > HORIZON_S:
return 0
return 1 if age_s <= FAST_S else (2 if age_s <= MED_S else 3)
class RealBank:
def __init__(self, d="simreal", bank="simreal/frames64.npy"):
self.f = np.load(bank, mmap_mode="r")
self.ts = np.load(bank.replace("frames64", "ts64"))
def at(self, t_ms):
i = bisect.bisect_left(self.ts, t_ms)
i = min(max(i, 0), len(self.ts) - 1)
return np.asarray(self.f[i], dtype=np.float32) / 255.0
def draw(frame, kind, pos, state, color, s=6, contrast=1.0):
x, y = pos
if contrast < 1.0:
import numpy as _np
def blend(region, col):
return col * contrast + region.mean(axis=(0, 1)) * (1 - contrast)
else:
blend = None
if kind == "lamp":
c_ = (color if state == 0 else 1.0 - color)
frame[y:y + s, x:x + s] = (blend(frame[y:y + s, x:x + s], c_) if blend else c_)
elif kind == "agent":
x2, y2 = (x + 22) % 52, (y + 17) % 52
if state == 0:
frame[y:y + s - 1, x:x + s - 1] = (blend(frame[y:y + s - 1, x:x + s - 1], color) if blend else color)
else:
frame[y2:y2 + s - 1, x2:x2 + s - 1] = (blend(frame[y2:y2 + s - 1, x2:x2 + s - 1], color) if blend else color)
elif kind == "furniture":
if state == 0:
frame[y:y + s - 1, x:x + s + 2] = (blend(frame[y:y + s - 1, x:x + s + 2], color) if blend else color)
return frame
def make_episode(bank, t_ms, ages_s, rng, static_bg=False, black_bg=False, objscale=6, contrast=1.0, noise_bg="off", noise_family=0):
objs = []
labels = []
for kind in ("lamp", "agent", "furniture"):
pos = (int(rng.integers(2, 52)), int(rng.integers(2, 52)))
color = rng.uniform(0.15, 0.95, size=3).astype(np.float32)
if kind == "furniture" and rng.random() < 0.25:
objs.append((kind, pos, None, color)); labels.append(0); continue
if rng.random() < 0.2: # static, never changes
objs.append((kind, pos, np.inf, color)); labels.append(0); continue
age = float(np.exp(rng.uniform(0.0, np.log(HORIZON_S))))
objs.append((kind, pos, age, color)); labels.append(bucket(age))
frames = []
import numpy as _np
bg0 = bank.at(int(t_ms)) if static_bg else None
for a in [0.0] + list(ages_s):
if noise_bg == "static":
sd = int(t_ms) % noise_family if noise_family else int(t_ms) % (2**31)
fr = _np.clip(_np.random.default_rng(sd).uniform(0, 1, (64, 64, 3)), 0, 1).astype(_np.float32)
if os.environ.get("NOISE_BLUR"):
import cv2 as _cv2
k = int(os.environ["NOISE_BLUR"])
fr = _cv2.GaussianBlur(fr, (k, k), 0)
elif noise_bg == "moving":
fr = _np.clip(_np.random.default_rng((int(t_ms) + int(a)) % (2**31)).uniform(0, 1, (64, 64, 3)), 0, 1).astype(_np.float32)
elif black_bg:
fr = _np.zeros((64, 64, 3), _np.float32)
else:
fr = (bg0.copy() if static_bg else bank.at(int(t_ms - a * 1000)).copy())
for (kind, pos, chg, color) in objs:
if chg is None:
if kind != "furniture":
fr = draw(fr, kind, pos, 1, color, objscale, contrast) # absent-forever furniture: skip
continue
state = 0 if a < chg else 1 # post-change if frame younger than change
# state semantics: 0 = new state, 1 = old state
fr = draw(fr, kind, pos, state, color, objscale, contrast)
frames.append(fr)
return np.stack(frames), np.array(labels, dtype=np.int64)
def make_batch(bank, lo, hi, ages, rng, n, static_bg=False, black_bg=False, objscale=6, contrast=1.0, noise_bg="off", noise_family=0):
ts = rng.integers(lo, hi, size=n)
fs, ys, ag = [], [], []
for t in ts:
aj = [max(1.0, a * float(rng.uniform(0.85, 1.15))) for a in ages]
f, y = make_episode(bank, int(t), aj, rng, static_bg, black_bg, objscale, contrast, noise_bg, noise_family)
fs.append(f); ys.append(y); ag.append(np.array([0.0] + aj, dtype=np.float32))
return np.stack(fs), np.stack(ag), np.stack(ys)
def to_torch(f, a, y, dev):
f = torch.from_numpy(f).permute(0, 1, 4, 2, 3).contiguous().to(dev)
return f, torch.from_numpy(a).to(dev), torch.from_numpy(y).to(dev)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--steps", type=int, default=6000)
ap.add_argument("--batch", type=int, default=32)
ap.add_argument("--budget", type=int, default=12)
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--eval-n", type=int, default=400)
ap.add_argument("--finetune-real", action="store_true")
ap.add_argument("--static-bg", action="store_true")
ap.add_argument("--black-bg", action="store_true")
ap.add_argument("--objscale", type=int, default=6)
ap.add_argument("--contrast", type=float, default=1.0)
ap.add_argument("--bank", default="simreal/frames64.npy")
ap.add_argument("--noise-bg", choices=["off","static","moving"], default="off")
ap.add_argument("--lr", type=float, default=1e-3)
ap.add_argument("--encoder", choices=["cnn","vit"], default="cnn")
ap.add_argument("--anneal-banks", default="", help="comma list of banks; switch by step fraction over first 75%")
ap.add_argument("--anneal-contrast-from", type=float, default=0, help="start contrast, anneal to --contrast by 75% of steps")
ap.add_argument("--anneal-from", type=int, default=0, help="start objscale, anneal to --objscale by 75% of steps")
ap.add_argument("--noise-family", type=int, default=0, help="draw static noise from K fixed textures (0=fresh per episode)")
ap.add_argument("--out", default="results_simreal")
args = ap.parse_args()
dev = "cuda" if torch.cuda.is_available() else "cpu"
torch.manual_seed(args.seed)
bank = RealBank(bank=args.bank)
banks_sched = [RealBank(bank=b) for b in args.anneal_banks.split(",")] if args.anneal_banks else None
ages = log_ages(args.budget)
span = bank.ts[-1] - bank.ts[0]
gold_after = bank.ts[-1] - int(span * 0.25)
lo = bank.ts[0] + HORIZON_S * 1000
net = LogLensNet().to(dev)
if args.encoder == "vit":
net.enc = ViTEncoder(64).to(dev)
opt = torch.optim.AdamW(net.parameters(), lr=args.lr)
rng = np.random.default_rng(args.seed)
t0 = time.time()
for step in range(args.steps):
if banks_sched:
frac = min(1.0, step / (0.75 * args.steps))
bank = banks_sched[min(len(banks_sched) - 1, int(frac * len(banks_sched)))]
if args.anneal_from:
frac = min(1.0, step / (0.75 * args.steps))
cur = int(round(args.anneal_from + (args.objscale - args.anneal_from) * frac))
else:
cur = args.objscale
if args.anneal_contrast_from:
frac = min(1.0, step / (0.75 * args.steps))
cur_ct = args.anneal_contrast_from + (args.contrast - args.anneal_contrast_from) * frac
else:
cur_ct = args.contrast
f, a, y = to_torch(*make_batch(bank, lo, gold_after, ages, rng, args.batch, args.static_bg, args.black_bg, cur, cur_ct, args.noise_bg, args.noise_family), dev)
loss = F.cross_entropy(net(f, a).flatten(0, 1), y.flatten())
opt.zero_grad(); loss.backward(); opt.step()
if (step + 1) % 250 == 0:
print(f"[simreal] step {step+1}/{args.steps} loss {loss.item():.4f} "
f"({time.time()-t0:.0f}s)", flush=True)
# held-out composited eval (fresh rng, query times INSIDE gold window's
# frames but with synthetic labels -> tests nuisance-robust dating)
net.eval()
erng = np.random.default_rng(777)
hit = np.zeros(4); tot = np.zeros(4)
with torch.no_grad():
for _ in range(args.eval_n // args.batch):
f, a, y = to_torch(*make_batch(bank, gold_after, bank.ts[-1], ages, erng,
args.batch, args.static_bg, args.black_bg, args.objscale, args.contrast), dev)
pred = net(f, a).argmax(-1).cpu().numpy()
yy = y.cpu().numpy()
for b in range(pred.shape[0]):
for o in range(3):
tot[yy[b, o]] += 1; hit[yy[b, o]] += pred[b, o] == yy[b, o]
acc = {b: round(float(h / t), 4) for b, h, t in zip(BUCKETS, hit, tot) if t}
mean = round(float(np.mean([h / t for h, t in zip(hit, tot) if t])), 4)
res = {"heldout_composited_acc": acc, "heldout_mean": mean,
"steps": args.steps, "seed": args.seed, "device": dev}
os.makedirs(args.out, exist_ok=True)
torch.save(net.state_dict(), os.path.join(args.out, f"simreal{'_black' if args.black_bg else ('_static' if args.static_bg else '')}_o{args.objscale}_ct{args.contrast:g}_{args.noise_bg}{args.noise_family}_{os.path.basename(args.bank)[:9]}_{args.encoder}_an{args.anneal_from}c{args.anneal_contrast_from:g}{'B' if args.anneal_banks else ''}_s{args.seed}.pt"))
json.dump(res, open(os.path.join(args.out, f"simreal{'_black' if args.black_bg else ('_static' if args.static_bg else '')}_o{args.objscale}_ct{args.contrast:g}_{args.noise_bg}{args.noise_family}_{os.path.basename(args.bank)[:9]}_{args.encoder}_an{args.anneal_from}c{args.anneal_contrast_from:g}{'B' if args.anneal_banks else ''}_s{args.seed}.json"), "w"), indent=2)
print(json.dumps(res, indent=2), flush=True)
if __name__ == "__main__":
main()