"""W11: neural lighting cache for the voxel world (ray-traced AO -> .pt). The ray-tracing tie-in for the Neural World demo. Sky visibility (ambient occlusion) is TRACED: from each ground point we shoot a hemisphere of rays and count how many escape past the tall blocks (house, trees) to the sky. That expensive per-point integral is distilled into a tiny MLP that the browser evaluates per tile in real time — the same "trace once, learn the field, look it up" idea as the W9 radiance cache, specialized to this scene's geometry. Layout MUST match the demo's hand-placed world (index.html): house 3x3 at rows/cols 2..4, ten trees at fixed spots — everything else is flat. """ import sys, os, time sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import numpy as np import torch import torch.nn as nn torch.manual_seed(0) HERE = os.path.dirname(os.path.abspath(__file__)) ROOT = os.path.dirname(HERE) DEV = "cuda" if torch.cuda.is_available() else "cpu" N = 32 NFREQ = 6 # tall occluders (cx, cz, half, top) — must match the demo world layout. # house: 6x6 perimeter centred at (14.5, 14.5); trees: the demo's fixed set. # x=col, z=row (three.js places tile (row r, col c) at x=c, z=r). TREES = [[3, 4], [5, 20], [7, 27], [2, 10], [4, 28], [9, 3], [6, 24], [10, 29], [20, 4], [22, 27], [25, 10], [27, 22], [29, 6], [24, 29], [19, 28], [28, 15], [3, 16], [8, 8], [23, 3], [26, 26], [21, 20], [18, 3], [29, 29], [2, 24]] HR0, HC0, HS = 12, 12, 6 BOXES = [(HC0 + HS / 2, HR0 + HS / 2, HS / 2, 2.2)] # house block (x,z,half,top) for (r, c) in TREES: if not (HR0 - 2 <= r < HR0 + HS + 2 and HC0 - 2 <= c < HC0 + HS + 2): BOXES.append((c + 0.5, r + 0.5, 0.6, 2.4)) # tree (x=col, z=row) BOX = torch.tensor(BOXES, device=DEV) # (B,4): cx,cz,half,top def sky_visibility(pts, n_rays=64): """pts (P,2) ground xz -> AO in [0,1] via hemisphere ray marching.""" P = len(pts) g = torch.Generator(device=DEV); g.manual_seed(1) # cosine-ish hemisphere directions (upper) u1 = torch.rand(n_rays, device=DEV, generator=g) u2 = torch.rand(n_rays, device=DEV, generator=g) r = torch.sqrt(u1); phi = 2 * np.pi * u2 dirs = torch.stack([r * torch.cos(phi), torch.sqrt((1 - u1).clamp_min(0)), # up = +y r * torch.sin(phi)], 1) # (n_rays,3) o = torch.stack([pts[:, 0], torch.full((P,), 0.55, device=DEV), pts[:, 1]], 1) # (P,3) vis = torch.ones(P, n_rays, device=DEV) cx, cz, half, top = BOX[:, 0], BOX[:, 1], BOX[:, 2], BOX[:, 3] # march: does ray (o,d) hit any box slab before escaping upward? for b in range(len(BOX)): # slab intersection in xz, check the entry y is below the box top dx = dirs[:, 0][None, :]; dz = dirs[:, 2][None, :]; dy = dirs[:, 1][None, :] ox = o[:, 0][:, None]; oz = o[:, 2][:, None]; oy = o[:, 1][:, None] inv_x = 1.0 / torch.where(dx.abs() < 1e-6, torch.full_like(dx, 1e-6), dx) inv_z = 1.0 / torch.where(dz.abs() < 1e-6, torch.full_like(dz, 1e-6), dz) tx1 = (cx[b] - half[b] - ox) * inv_x; tx2 = (cx[b] + half[b] - ox) * inv_x tz1 = (cz[b] - half[b] - oz) * inv_z; tz2 = (cz[b] + half[b] - oz) * inv_z tmin = torch.maximum(torch.minimum(tx1, tx2), torch.minimum(tz1, tz2)) tmax = torch.minimum(torch.maximum(tx1, tx2), torch.maximum(tz1, tz2)) hit_xz = (tmax > tmin) & (tmax > 0) t_enter = tmin.clamp_min(0) y_at = oy + dy * t_enter blocked = hit_xz & (y_at < top[b]) & (t_enter > 1e-3) vis = torch.where(blocked, torch.zeros_like(vis), vis) return vis.mean(1) # (P,) AO def freq_encode(xz): bands = (2.0 ** torch.arange(NFREQ, device=xz.device)) * np.pi proj = xz[..., None] * bands return torch.cat([xz, torch.sin(proj).flatten(-2), torch.cos(proj).flatten(-2)], -1) class LightNet(nn.Module): def __init__(self, hidden=64): super().__init__() din = 2 + 2 * 2 * NFREQ self.net = nn.Sequential(nn.Linear(din, hidden), nn.SiLU(), nn.Linear(hidden, hidden), nn.SiLU(), nn.Linear(hidden, 3)) def forward(self, xz): return torch.nn.functional.softplus(self.net(freq_encode(xz))) # ---- ground-truth AO -> ambient RGB on a dense grid ---- print(f"=== tracing sky visibility on {DEV} ===") t0 = time.time() G = 112 gx, gz = torch.meshgrid(torch.linspace(0, N, G, device=DEV), torch.linspace(0, N, G, device=DEV), indexing="ij") pts = torch.stack([gx.ravel(), gz.ravel()], 1) ao = sky_visibility(pts, n_rays=96) SKY = torch.tensor([0.55, 0.70, 1.0], device=DEV) # cool sky GROUND = torch.tensor([0.30, 0.32, 0.28], device=DEV) # warm bounce floor amb = (0.30 + 0.70 * ao)[:, None] * SKY + (1 - ao)[:, None] * GROUND * 0.4 xz_norm = pts / N * 2 - 1 print(f" {G*G} points, mean AO {ao.mean():.3f} " f"(min {ao.min():.3f}) ({time.time()-t0:.1f}s)") # ---- fit the cache ---- net = LightNet().to(DEV) print(f"=== training light cache ({sum(p.numel() for p in net.parameters())} params) ===") opt = torch.optim.Adam(net.parameters(), lr=3e-3) sched = torch.optim.lr_scheduler.CosineAnnealingLR(opt, T_max=2500) for it in range(2500): idx = torch.randint(0, len(xz_norm), (8192,), device=DEV) pred = net(xz_norm[idx]) loss = ((pred - amb[idx]) ** 2).mean() opt.zero_grad(); loss.backward(); opt.step(); sched.step() if it % 500 == 0 or it == 2499: print(f" it {it:4d} MSE {loss.item():.6f}") with torch.no_grad(): err = (net(xz_norm) - amb).abs().mean().item() print(f" final mean abs error {err:.4f}") # ---- figure: traced vs learned ambient ---- try: import matplotlib; matplotlib.use("Agg"); import matplotlib.pyplot as plt with torch.no_grad(): pr = net(xz_norm).reshape(G, G, 3).clamp(0, 1).cpu().numpy() gt = amb.reshape(G, G, 3).clamp(0, 1).cpu().numpy() fig, ax = plt.subplots(1, 3, figsize=(14, 4.6)) ax[0].imshow(ao.reshape(G, G).cpu().numpy(), cmap="bone"); ax[0].set_title("traced sky visibility (AO)") ax[1].imshow(gt); ax[1].set_title("target ambient (AO -> RGB)") ax[2].imshow(pr); ax[2].set_title(f"neural cache (world_light.pt), err {err:.3f}") for a in ax: a.set_xticks([]); a.set_yticks([]) fig.suptitle("Neural world lighting — traced ambient occlusion distilled to a .pt") fig.tight_layout() out = os.path.join(ROOT, "world_lighting_validation.png") fig.savefig(out, dpi=110); print(f"wrote {out}") except Exception as e: print("plot skipped:", e) torch.save({"state_dict": net.state_dict(), "arch": {"nfreq": NFREQ, "hidden": 64}, "scene": "neural_world_32"}, os.path.join(HERE, "world_light.pt")) print("saved experiments/world_light.pt")