File size: 6,998 Bytes
ac9d706 | 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 | """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")
|