"""pathtracer-diff: a differentiable Monte Carlo path tracer as one kernel. Forward renders a linear-radiance image with a megakernel path tracer: Lambertian diffuse, GGX conductor (VNDF-sampled, height-correlated Smith), smooth dielectric, rough plastic, and rough dielectric materials; area lights with per-texel emission and an importance-sampled equirectangular environment map; an optional homogeneous participating medium; multiple importance sampling by default; per-texel albedo with bilinear filtering; binned-SAH BVH. One thread owns one pixel, so the image is bitwise deterministic. Backward is exact path replay: sampling distributions depend only on geometry, frozen material parameters (roughness, ior), the environment CDF and medium rate (detached at Scene construction), and the counter-based Philox stream, never on the differentiable parameters. Every radiance term is a product of per-bounce factors that are affine in their vertex's albedo texels (Schlick Fresnel is affine in F0), times a linear emission or environment texel, with medium sigma-derivatives in closed form, so the replay differentiates exactly and scatters through the bilinear-footprint adjoint. render() gradients flow to albedo texels, emission texels, environment texels, and medium sigmas. geometry_grad() adds d(loss)/d(vertex positions) for direct lighting: a dual-number interior term plus edge-sampled shadow and camera silhouette boundary terms. from kernels import get_kernel ptd = get_kernel("phanerozoic/pathtracer-diff", version=1, trust_remote_code=True) scene = ptd.Scene(vertices, faces, material_ids, albedo, emission, uvs=uvs, material_types=types, roughness=rough, ior=ior, env=env_map) cam = ptd.Camera(position=(2.78, 2.73, -8.0), look_at=(2.78, 2.73, 2.8)) img = ptd.render(scene, cam, 512, 512, spp=64) img.sum().backward() # grads on albedo texels / emission / env texels gv = ptd.geometry_grad(scene, cam, dL_dimage) # grads on vertices """ import math import torch from ._ops import ops MAX_MATERIALS = 64 MAX_BOUNCES = 16 DIFFUSE, CONDUCTOR, DIELECTRIC, PLASTIC, ROUGH_DIELECTRIC = 0, 1, 2, 3, 4 _MODES = {"brdf": 0, "nee": 1, "mis": 2} __all__ = ["Scene", "Camera", "render", "geometry_grad", "ops", "MAX_MATERIALS", "MAX_BOUNCES", "DIFFUSE", "CONDUCTOR", "DIELECTRIC", "PLASTIC", "ROUGH_DIELECTRIC"] def _build_bvh(tris, leaf_size=4, sah_bins=16): """Binned-SAH BVH over [F, 9] triangles, built level-wise with fully vectorized torch ops (segmented reductions + stable-sort partitioning), so million-triangle scenes build in seconds. Returns (nodes_f [N, 6], nodes_i [N, 3], order [F]). nodes_i rows are (left, right, axis << 1) for internal nodes and (start, count, 1) for leaves; internal children are the (lower, upper) coordinate halves along the split axis, so traversal can visit the near child first.""" dev = "cuda" if torch.cuda.is_available() else "cpu" tris = tris.to(torch.float32).to(dev) F = tris.shape[0] v = tris.reshape(F, 3, 3) tlo = v.amin(dim=1) thi = v.amax(dim=1) cen = (tlo + thi) * 0.5 perm = torch.arange(F, device=dev) # active segments: contiguous [start, end) ranges of perm, one per node starts = torch.zeros(1, dtype=torch.long, device=dev) ends = torch.full((1,), F, dtype=torch.long, device=dev) ids = torch.zeros(1, dtype=torch.long, device=dev) INF = float("inf") nf_parts, ni_parts = [], [] # per-node rows appended in id order n_nodes = 1 def half_area(lo, hi): d = (hi - lo).clamp(min=0) return d[..., 0] * d[..., 1] + d[..., 1] * d[..., 2] + \ d[..., 2] * d[..., 0] while ids.numel() > 0: A = ids.numel() seg_len = ends - starts total = int(seg_len.sum()) # active positions and per-position segment index csum = torch.cumsum(seg_len, 0) - seg_len pos = starts.repeat_interleave(seg_len) + \ (torch.arange(total, device=dev) - csum.repeat_interleave(seg_len)) s = torch.arange(A, device=dev).repeat_interleave(seg_len) tri = perm[pos] # per-segment node bounds and centroid bounds nlo = torch.full((A, 3), INF, device=dev) nhi = torch.full((A, 3), -INF, device=dev) clo = torch.full((A, 3), INF, device=dev) chi = torch.full((A, 3), -INF, device=dev) s3 = s.unsqueeze(1).expand(-1, 3) nlo.scatter_reduce_(0, s3, tlo[tri], reduce="amin", include_self=True) nhi.scatter_reduce_(0, s3, thi[tri], reduce="amax", include_self=True) clo.scatter_reduce_(0, s3, cen[tri], reduce="amin", include_self=True) chi.scatter_reduce_(0, s3, cen[tri], reduce="amax", include_self=True) pad = 1e-5 * (1.0 + nlo.abs() + nhi.abs()) bounds = torch.cat([nlo - pad, nhi + pad], dim=1) # [A, 6] leaf_mask = seg_len <= leaf_size split_mask = ~leaf_mask # choose axis (largest centroid extent) and bin the centroids ext = (chi - clo).clamp(min=0) axis = ext.argmax(dim=1) # [A] ax_pos = axis[s] coord = cen[tri].gather(1, ax_pos.unsqueeze(1)).squeeze(1) seg_lo = clo.gather(1, axis.unsqueeze(1)).squeeze(1)[s] seg_ext = ext.gather(1, axis.unsqueeze(1)).squeeze(1)[s] t = (coord - seg_lo) / seg_ext.clamp(min=1e-30) bin_i = (t * sah_bins).long().clamp(0, sah_bins - 1) degen = seg_ext <= 1e-12 # all centroids equal on the axis # per-(segment, bin) counts and bounds key = s * sah_bins + bin_i counts = torch.bincount(key, minlength=A * sah_bins) \ .reshape(A, sah_bins) blo = torch.full((A * sah_bins, 3), INF, device=dev) bhi = torch.full((A * sah_bins, 3), -INF, device=dev) k3 = key.unsqueeze(1).expand(-1, 3) blo.scatter_reduce_(0, k3, tlo[tri], reduce="amin", include_self=True) bhi.scatter_reduce_(0, k3, thi[tri], reduce="amax", include_self=True) blo = blo.reshape(A, sah_bins, 3) bhi = bhi.reshape(A, sah_bins, 3) # prefix (left) and suffix (right) running bounds and counts plo = torch.cummin(blo, dim=1).values phi = torch.cummax(bhi, dim=1).values pn = torch.cumsum(counts, dim=1) slo = torch.flip(torch.cummin(torch.flip(blo, [1]), dim=1).values, [1]) shi = torch.flip(torch.cummax(torch.flip(bhi, [1]), dim=1).values, [1]) sn = torch.flip(torch.cumsum(torch.flip(counts, [1]), dim=1), [1]) # SAH cost of splitting after bin b (b = 0..bins-2) nl = pn[:, :-1].float() nr = sn[:, 1:].float() cost = half_area(plo[:, :-1], phi[:, :-1]) * nl + \ half_area(slo[:, 1:], shi[:, 1:]) * nr cost = torch.where((nl > 0) & (nr > 0), cost, torch.full_like(cost, INF)) best_cost, best_bin = cost.min(dim=1) sah_ok = torch.isfinite(best_cost) & ~degen[csum] # per segment # side: left if bin <= best_bin (SAH) else lower positional half # (median fallback for degenerate segments) side_left = bin_i <= best_bin[s] within = pos - starts[s] med_left = within < (seg_len[s] // 2) use_sah = sah_ok[s] side_left = torch.where(use_sah, side_left, med_left) side_left = side_left & split_mask[s] # per-segment left count; guard SAH splits that put everything on # one side (possible only via numerics) with the median fallback nl_seg = torch.zeros(A, dtype=torch.long, device=dev) nl_seg.scatter_add_(0, s, side_left.long()) bad = split_mask & ((nl_seg == 0) | (nl_seg == seg_len)) if bool(bad.any()): fix = bad[s] side_left = torch.where(fix, med_left & split_mask[s], side_left) nl_seg = torch.zeros(A, dtype=torch.long, device=dev) nl_seg.scatter_add_(0, s, side_left.long()) # partition each segment in place: stable sort by (segment, right?) k2 = s * 2 + (~side_left).long() order2 = torch.argsort(k2, stable=True) perm[pos] = tri[order2] # allocate children for splitting segments, emit node rows n_split = int(split_mask.sum()) child_rank = torch.cumsum(split_mask.long(), 0) - 1 left_ids = n_nodes + 2 * child_rank right_ids = left_ids + 1 ni = torch.empty(A, 3, dtype=torch.int64, device=dev) ni[:, 0] = torch.where(split_mask, left_ids, starts) ni[:, 1] = torch.where(split_mask, right_ids, seg_len) ni[:, 2] = torch.where(split_mask, axis * 2, torch.ones_like(axis)) nf_parts.append((ids, bounds)) ni_parts.append((ids, ni)) # next level mid = starts + nl_seg new_starts = torch.cat([starts[split_mask], mid[split_mask]]) new_ends = torch.cat([mid[split_mask], ends[split_mask]]) new_ids = torch.cat([left_ids[split_mask], right_ids[split_mask]]) n_nodes += 2 * n_split starts, ends, ids = new_starts, new_ends, new_ids nodes_f = torch.empty(n_nodes, 6, device=dev) nodes_i = torch.empty(n_nodes, 3, dtype=torch.int64, device=dev) for idv, rows in nf_parts: nodes_f[idv] = rows for idv, rows in ni_parts: nodes_i[idv] = rows return (nodes_f.cpu(), nodes_i.to(torch.int32).cpu(), perm.cpu()) def _build_bvh_reference(tris, leaf_size=4, force_split=8, sah_bins=16): """Recursive reference builder (kept for cross-checking the vectorized build in development).""" tris = tris.to(torch.float32) F = tris.shape[0] v = tris.reshape(F, 3, 3) lo = v.amin(dim=1) hi = v.amax(dim=1) cen = (lo + hi) * 0.5 nodes_f, nodes_i, order = [], [], [] def half_area(blo, bhi): d = (bhi - blo).clamp(min=0) return d[0] * d[1] + d[1] * d[2] + d[2] * d[0] stack = [] def alloc(idx): nid = len(nodes_f) nodes_f.append(None) nodes_i.append(None) stack.append((nid, idx)) return nid root = alloc(torch.arange(F)) while stack: nid, idx = stack.pop() n = idx.numel() blo = lo[idx].amin(dim=0) bhi = hi[idx].amax(dim=0) pad = 1e-5 * (1.0 + blo.abs() + bhi.abs()) nodes_f[nid] = torch.cat([blo - pad, bhi + pad]) if n <= leaf_size: start = len(order) order.extend(idx.tolist()) nodes_i[nid] = (start, n, 1) continue cb_lo = cen[idx].amin(dim=0) cb_hi = cen[idx].amax(dim=0) ext = cb_hi - cb_lo axis = int(ext.argmax()) split_pos = None if float(ext[axis]) > 1e-12: c = cen[idx, axis] edges = torch.linspace(float(cb_lo[axis]), float(cb_hi[axis]), sah_bins + 1) b = torch.bucketize(c, edges[1:-1]) counts = torch.bincount(b, minlength=sah_bins) # per-bin bounds over the node's triangles binlo = torch.full((sah_bins, 3), float("inf")) binhi = torch.full((sah_bins, 3), float("-inf")) binlo.scatter_reduce_(0, b.unsqueeze(1).expand(-1, 3), lo[idx], reduce="amin", include_self=True) binhi.scatter_reduce_(0, b.unsqueeze(1).expand(-1, 3), hi[idx], reduce="amax", include_self=True) best_cost, best_bin = None, None nl = 0 llo = torch.full((3,), float("inf")) lhi = torch.full((3,), float("-inf")) pre = [] for i in range(sah_bins - 1): if counts[i] > 0: llo = torch.minimum(llo, binlo[i]) lhi = torch.maximum(lhi, binhi[i]) nl += int(counts[i]) pre.append((nl, llo.clone(), lhi.clone())) nr = 0 rlo = torch.full((3,), float("inf")) rhi = torch.full((3,), float("-inf")) for i in range(sah_bins - 1, 0, -1): if counts[i] > 0: rlo = torch.minimum(rlo, binlo[i]) rhi = torch.maximum(rhi, binhi[i]) nr += int(counts[i]) nl_i, llo_i, lhi_i = pre[i - 1] if nl_i == 0 or nr == 0: continue cost = (half_area(llo_i, lhi_i) * nl_i + half_area(rlo, rhi) * nr) if best_cost is None or float(cost) < best_cost: best_cost = float(cost) best_bin = i - 1 if best_bin is not None: parent_area = half_area(blo, bhi) leaf_cost = float(n) * float(parent_area) split_cost = 0.125 * float(parent_area) + best_cost if split_cost < leaf_cost or n > force_split: mask = b <= best_bin left_idx = idx[mask] right_idx = idx[~mask] if left_idx.numel() > 0 and right_idx.numel() > 0: split_pos = (left_idx, right_idx) if split_pos is None: if n > force_split and float(ext[axis]) > 1e-12: srt = idx[torch.argsort(cen[idx, axis], stable=True)] mid = n // 2 split_pos = (srt[:mid], srt[mid:]) else: start = len(order) order.extend(idx.tolist()) nodes_i[nid] = (start, n, 1) continue left = alloc(split_pos[0]) right = alloc(split_pos[1]) nodes_i[nid] = (left, right, axis << 1) assert root == 0 return (torch.stack(nodes_f), torch.tensor(nodes_i, dtype=torch.int32), torch.tensor(order, dtype=torch.int64)) class Camera: """Pinhole camera. `tensor(H, W)` packs (pos, forward, right*tan(v/2)*aspect, up*tan(v/2)) as a float32 [12] tensor for the kernel.""" def __init__(self, position, look_at, up=(0.0, 1.0, 0.0), vfov_deg=40.0): p = torch.tensor(position, dtype=torch.float64) t = torch.tensor(look_at, dtype=torch.float64) u = torch.tensor(up, dtype=torch.float64) f = t - p f = f / f.norm() r = torch.linalg.cross(f, u) r = r / r.norm() uu = torch.linalg.cross(r, f) self.position, self.forward, self.right, self.up = p, f, r, uu self.vfov_deg = float(vfov_deg) def tensor(self, H, W, device="cuda"): th = math.tan(math.radians(self.vfov_deg) * 0.5) rs = self.right * (th * W / H) us = self.up * th return torch.cat([self.position, self.forward, rs, us]).to( device=device, dtype=torch.float32).contiguous() def _as_texture(t, device): t = torch.as_tensor(t, dtype=torch.float32) if t.dim() == 1 and t.numel() == 3: t = t.reshape(1, 1, 3) if t.dim() != 3 or t.shape[2] != 3 or t.shape[0] < 1 or t.shape[1] < 1: raise ValueError("each albedo texture must be [H, W, 3]") return t.to(device) class Scene: """Triangle-mesh scene. vertices [V, 3], faces [F, 3], material_ids [F]. albedo: [M, 3] constants or a list of M textures [Hm, Wm, 3] (constants are 1x1 textures; both forms receive gradients). emission [M, 3] (may require grad). uvs: None, [V, 2], or [F, 3, 2]; textures wrap. material_types [M] of DIFFUSE|CONDUCTOR|DIELECTRIC (default all DIFFUSE); roughness [M] GGX alpha for conductors (default 0.3, clamped >= 0.01); ior [M] for dielectrics (default 1.5). env: optional [Eh, Ew, 3] equirect radiance map (may require grad); its importance-sampling CDF is built detached at construction. Texture and env shapes are fixed at construction.""" def __init__(self, vertices, faces, material_ids, albedo, emission, uvs=None, material_types=None, roughness=None, ior=None, env=None, medium=None, device="cuda"): vertices = torch.as_tensor(vertices, dtype=torch.float32) faces = torch.as_tensor(faces, dtype=torch.int64) material_ids = torch.as_tensor(material_ids, dtype=torch.int64) is_texture_list = (isinstance(albedo, (list, tuple)) and any(torch.is_tensor(t) for t in albedo)) if is_texture_list: self.albedo = None self.albedo_textures = [_as_texture(t, device) for t in albedo] else: if not torch.is_tensor(albedo): albedo = torch.tensor(albedo, dtype=torch.float32, device=device) if albedo.dim() != 2 or albedo.shape[1] != 3: raise ValueError("albedo must be [M, 3] or a list of textures") self.albedo = albedo.to(device) self.albedo_textures = None M = (len(self.albedo_textures) if self.albedo is None else self.albedo.shape[0]) if M > MAX_MATERIALS: raise ValueError(f"at most {MAX_MATERIALS} materials, got {M}") emi_texture_list = (isinstance(emission, (list, tuple)) and any(torch.is_tensor(t) for t in emission)) if emi_texture_list: self.emission = None self.emission_textures = [_as_texture(t, device) for t in emission] if len(self.emission_textures) != M: raise ValueError("emission list must have M entries") else: if not torch.is_tensor(emission): emission = torch.tensor(emission, dtype=torch.float32, device=device) if emission.shape != (M, 3): raise ValueError("emission must be [M, 3]") self.emission = emission.to(device) self.emission_textures = None if int(material_ids.max()) >= M or int(material_ids.min()) < 0: raise ValueError("material_ids out of range") def mvec(x, default, dtype): if x is None: return torch.full((M,), default, dtype=dtype, device=device) t = torch.as_tensor(x, dtype=dtype).reshape(-1).to(device) if t.numel() != M: raise ValueError("per-material array must have M entries") return t self.mat_type = mvec(material_types, DIFFUSE, torch.int32).contiguous() if int(self.mat_type.max()) > 4 or int(self.mat_type.min()) < 0: raise ValueError("material_types must be 0..4") self.mat_rough = mvec(roughness, 0.3, torch.float32).contiguous() self.mat_ior = mvec(ior, 1.5, torch.float32).contiguous() def build_hdr(constants, textures): shapes = ([(1, 1)] * M if constants is not None else [(int(t.shape[0]), int(t.shape[1])) for t in textures]) hdr, off = [], 0 for (h, w) in shapes: hdr.append((off, w, h)) off += h * w return (torch.tensor(hdr, dtype=torch.int32, device=device).contiguous(), off) self.tex_hdr, self.n_texels = build_hdr(self.albedo, self.albedo_textures) self.emi_hdr, self.n_emi_texels = build_hdr(self.emission, self.emission_textures) # homogeneous medium: live sigmas (may require grad), frozen # (detached) sampling rate if medium is not None: if env is not None: raise ValueError("medium and env are mutually exclusive") sa, ss = medium sa = sa if torch.is_tensor(sa) else torch.tensor( sa, dtype=torch.float32, device=device) ss = ss if torch.is_tensor(ss) else torch.tensor( ss, dtype=torch.float32, device=device) if sa.numel() != 3 or ss.numel() != 3: raise ValueError("medium sigmas must be [3]") self.med_sa = sa.to(device) self.med_ss = ss.to(device) st = (sa.detach() + ss.detach()).reshape(-1) self.med_sbar = float(st.mean().clamp(min=1e-6)) else: self.med_sa = None self.med_ss = None self.med_sbar = 0.0 tris = vertices[faces].reshape(-1, 9) if uvs is None: uv_c = torch.zeros(tris.shape[0], 3, 2, dtype=torch.float32) else: uvs = torch.as_tensor(uvs, dtype=torch.float32) if uvs.dim() == 2 and uvs.shape[1] == 2: uv_c = uvs[faces] elif uvs.dim() == 3 and uvs.shape[1:] == (3, 2): uv_c = uvs else: raise ValueError("uvs must be [V, 2] or [F, 3, 2]") nodes_f, nodes_i, order = _build_bvh(tris) tris = tris[order] uv_c = uv_c[order] mat_ids = material_ids[order].to(torch.int32) self.face_verts = faces[order].to(torch.int32).to(device).contiguous() self.n_verts = int(vertices.shape[0]) self._vertices_cpu = vertices.cpu() self._faces_cpu = faces[order].cpu() self._edges = None if self.emission is not None: em_max = self.emission.detach().cpu().amax(dim=1) else: em_max = torch.stack([t.detach().max().cpu() for t in self.emission_textures]) emissive = (em_max > 0)[mat_ids.long()] lf = torch.nonzero(emissive, as_tuple=False).flatten().to(torch.int32) if lf.numel() > 0: t = tris[lf.long()].reshape(-1, 3, 3).double() e1 = t[:, 1] - t[:, 0] e2 = t[:, 2] - t[:, 0] areas = 0.5 * torch.linalg.cross(e1, e2).norm(dim=1) total = float(areas.sum()) cdf = (areas.cumsum(0) / areas.sum()).float() else: total = 0.0 cdf = torch.zeros(0, dtype=torch.float32) # environment map + detached sampling tables if env is not None: env = torch.as_tensor(env, dtype=torch.float32) if env.dim() != 3 or env.shape[2] != 3 or env.shape[0] < 2 or \ env.shape[1] < 2: raise ValueError("env must be [Eh, Ew, 3]") self.env = env.to(device) eh, ew = int(env.shape[0]), int(env.shape[1]) lum = self.env.detach().mean(dim=2).cpu().double() + 1e-8 sint = torch.sin((torch.arange(eh, dtype=torch.float64) + 0.5) * math.pi / eh).clamp(min=1e-4) roww = (lum.sum(dim=1) * sint) row_p = roww / roww.sum() cdf_m = row_p.cumsum(0).float() col_p = lum / lum.sum(dim=1, keepdim=True) cdf_c = col_p.cumsum(1).float() pdf_img = (row_p.unsqueeze(1) * col_p).float() # sums to 1 self.env_w, self.env_h = ew, eh self.env_cdf_m = cdf_m.to(device).contiguous() self.env_cdf_c = cdf_c.reshape(-1).to(device).contiguous() self.env_pdf = pdf_img.reshape(-1).to(device).contiguous() else: self.env = None self.env_w = self.env_h = 0 z = torch.zeros(0, dtype=torch.float32, device=device) self.env_cdf_m = z self.env_cdf_c = z self.env_pdf = z self.device = device self.tris = tris.to(device).contiguous() self.mat_ids = mat_ids.to(device).contiguous() self.uvs = uv_c.to(device).contiguous() self.nodes_f = nodes_f.to(device).contiguous() self.nodes_i = nodes_i.to(device).contiguous() self.light_faces = lf.to(device).contiguous() self.light_cdf = cdf.to(device).contiguous() self.total_light_area = total @property def n_faces(self): return self.tris.shape[0] @staticmethod def _flatten(constants, textures, hdr, what): if constants is not None: if constants.shape[0] != hdr.shape[0]: raise ValueError(f"{what}/material count changed") return constants.reshape(-1, 3) parts = [] for m, t in enumerate(textures): off, w, h = (int(x) for x in hdr[m]) if (int(t.shape[0]), int(t.shape[1])) != (h, w): raise ValueError(f"{what} shapes are fixed at construction") parts.append(t.reshape(-1, 3)) return torch.cat(parts, dim=0) def _flat_albedo(self): return self._flatten(self.albedo, self.albedo_textures, self.tex_hdr, "albedo") def _flat_emission(self): return self._flatten(self.emission, self.emission_textures, self.emi_hdr, "emission") def _flat_env(self): if self.env is None: return torch.zeros(0, 3, dtype=torch.float32, device=self.device) if (int(self.env.shape[0]), int(self.env.shape[1])) != \ (self.env_h, self.env_w): raise ValueError("env shape is fixed at construction") return self.env.reshape(-1, 3) class _RenderFn(torch.autograd.Function): @staticmethod def _args(scene, tex_flat, emi_flat, env_flat, med_sa, med_ss, cam_t, spp, max_bounces, mode, seed): return (scene.tris, scene.mat_ids, scene.uvs, scene.nodes_f, scene.nodes_i, scene.light_faces, scene.light_cdf, float(scene.total_light_area), tex_flat.detach().contiguous(), scene.tex_hdr, emi_flat.detach().contiguous(), scene.emi_hdr, scene.mat_type, scene.mat_rough, scene.mat_ior, med_sa.detach().contiguous(), med_ss.detach().contiguous(), float(scene.med_sbar), env_flat.detach().contiguous(), scene.env_cdf_m, scene.env_cdf_c, scene.env_pdf, scene.env_w, scene.env_h, cam_t, spp, max_bounces, mode, seed) @staticmethod def forward(ctx, tex_flat, emi_flat, env_flat, med_sa, med_ss, scene, cam_t, H, W, spp, max_bounces, mode, seed): image = torch.empty(H, W, 3, device=tex_flat.device, dtype=torch.float32) ops.pt_forward(*_RenderFn._args(scene, tex_flat, emi_flat, env_flat, med_sa, med_ss, cam_t, spp, max_bounces, mode, seed), image) ctx.scene = scene ctx.cam_t = cam_t ctx.params = (H, W, spp, max_bounces, mode, seed) ctx.save_for_backward(tex_flat, emi_flat, env_flat, med_sa, med_ss) return image @staticmethod def backward(ctx, grad_image): tex_flat, emi_flat, env_flat, med_sa, med_ss = ctx.saved_tensors scene, cam_t = ctx.scene, ctx.cam_t H, W, spp, max_bounces, mode, seed = ctx.params ga = torch.zeros_like(tex_flat) ge = torch.zeros_like(emi_flat) genv = torch.zeros_like(env_flat) gmed = torch.zeros(6 if med_sa.numel() else 0, device=tex_flat.device) ops.pt_backward(*_RenderFn._args(scene, tex_flat, emi_flat, env_flat, med_sa, med_ss, cam_t, spp, max_bounces, mode, seed), grad_image.contiguous(), ga, ge, genv, gmed) gsa = gmed[:3] if med_sa.numel() else None gss = gmed[3:] if med_sa.numel() else None return (ga, ge, genv, gsa, gss, None, None, None, None, None, None, None, None) def _scene_edges(scene): if scene._edges is None: fv = scene._faces_cpu vp = scene._vertices_cpu first = {} rows = [] for f in range(fv.shape[0]): ids = fv[f] for (ca, cb) in ((0, 1), (1, 2), (2, 0)): a, b = int(ids[ca]), int(ids[cb]) key = (min(a, b), max(a, b)) if key in first: rows[first[key]][3] = f else: first[key] = len(rows) rows.append([ca, cb, f, -1]) edges = torch.tensor(rows, dtype=torch.int32) p0 = vp[fv[edges[:, 2].long(), edges[:, 0].long()].long()] p1 = vp[fv[edges[:, 2].long(), edges[:, 1].long()].long()] lens = (p1 - p0).norm(dim=1).double() total = float(lens.sum()) cdf = (lens.cumsum(0) / lens.sum()).float() cdf = torch.cat([cdf, torch.tensor([total])]) scene._edges = (edges.to(scene.device).contiguous(), cdf.to(scene.device).contiguous()) return scene._edges def geometry_grad(scene, camera, grad_image, spp=16, edge_samples=1 << 16, seed=0): """d(loss)/d(vertex positions) for the DIRECT-lighting transport term: the detached-sampling interior derivative (dual-number shading + light area measure) plus edge-sampled shadow-silhouette and primary (camera-silhouette) boundary terms. Diffuse receivers/emitters carry the radiance jumps; indirect bounces are not differentiated with respect to geometry. grad_image is dLoss/dImage [H, W, 3]; returns grad_verts [V, 3].""" edges, ecdf = _scene_edges(scene) grad_verts = torch.zeros(scene.n_verts, 3, device=scene.device) cam_t = camera.tensor(grad_image.shape[0], grad_image.shape[1], device=scene.device) ops.pt_geometry_grad(scene.tris, scene.mat_ids, scene.uvs, scene.nodes_f, scene.nodes_i, scene.light_faces, scene.light_cdf, float(scene.total_light_area), scene._flat_albedo().detach().contiguous(), scene.tex_hdr, scene._flat_emission().detach().contiguous(), scene.emi_hdr, scene.mat_type, scene.mat_rough, scene.mat_ior, scene.face_verts, edges, ecdf, cam_t, int(spp), int(edge_samples), int(seed), grad_image.detach().contiguous(), grad_verts) return grad_verts def render(scene, camera, height, width, spp=64, max_bounces=4, estimator=None, nee=None, seed=0): """Render a linear-radiance image [H, W, 3] float32, differentiable with respect to the scene's albedo texels, emission, and environment texels. estimator: "mis" (default), "nee", or "brdf"; the legacy `nee` bool maps True->"nee", False->"brdf". A fixed seed renders the same paths every call, so the Monte Carlo objective is a deterministic function of the parameters.""" if not (1 <= max_bounces <= MAX_BOUNCES): raise ValueError(f"max_bounces must be in [1, {MAX_BOUNCES}]") if estimator is not None and nee is not None: raise ValueError("pass estimator or nee, not both") if estimator is None: estimator = "mis" if nee is None else ("nee" if nee else "brdf") if estimator not in _MODES: raise ValueError('estimator must be "mis", "nee", or "brdf"') if estimator == "nee" and scene.env is not None: raise ValueError('environment maps require the "mis" or "brdf" estimator') cam_t = camera.tensor(height, width, device=scene.device) if scene.med_sa is not None: med_sa, med_ss = scene.med_sa, scene.med_ss else: med_sa = torch.zeros(0, device=scene.device) med_ss = torch.zeros(0, device=scene.device) return _RenderFn.apply(scene._flat_albedo(), scene._flat_emission(), scene._flat_env(), med_sa, med_ss, scene, cam_t, height, width, int(spp), int(max_bounces), _MODES[estimator], int(seed))