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")