Download experiments/w11_world_lighting.py from NeuralVerified/neural-raytracing: direct link, hf CLI and curl.
- Browser
- Download file 7 kB
-
https://huggingface.co/NeuralVerified/neural-raytracing/resolve/main/experiments/w11_world_lighting.py
- Command line
-
hf download hf://NeuralVerified/neural-raytracing/experiments/w11_world_lighting.py
-
curl -L -o w11_world_lighting.py https://huggingface.co/NeuralVerified/neural-raytracing/resolve/main/experiments/w11_world_lighting.py
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") | |