neural-raytracing / experiments /w11_world_lighting.py
Quazim0t0's picture
Import from Quazim0t0/neural-raytracing; repoint refs to NeuralVerified
ac9d706 verified
Raw History Blame Contribute Delete
7 kB
"""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")