""" PyTorch integration for differentiable SDF operations. Provides torch.autograd.Function wrappers that allow SDF evaluation within PyTorch computation graphs. Usage: import torch from vde.torch_sdf import SdfEvaluator evaluator = SdfEvaluator(node) x = torch.randn(100, 3, requires_grad=True) d = evaluator(x) # SDF values loss = d.abs().mean() loss.backward() # x.grad contains df/dx """ import torch class SdfEvaluator(torch.autograd.Function): """Differentiable SDF evaluation for PyTorch. Forward: evaluate SDF at input points. Backward: compute gradient of SDF w.r.t. input points. Usage: evaluator = SdfEvaluator.apply d = evaluator(node, points) # where node is an SdfNode """ @staticmethod def forward(ctx, node, points): """Forward pass: evaluate SDF. Args: node: SdfNode (from C++ binding) points: torch.Tensor of shape (N, 3) Returns: distances: torch.Tensor of shape (N,) """ import numpy as np from vde._vde import evaluate # Convert to numpy for C++ evaluation pts_np = points.detach().cpu().numpy().astype(np.float64) n = pts_np.shape[0] distances = np.zeros(n, dtype=np.float64) for i in range(n): p = pts_np[i] distances[i] = evaluate(node, p) ctx.save_for_backward(points) ctx.node = node # Store node reference for backward return torch.from_numpy(distances).to(points.device).to(points.dtype) @staticmethod def backward(ctx, grad_output): """Backward pass: gradient of SDF w.r.t. input points. df/dx = gradient of SDF at each input point Chain rule: dL/dx = dL/d(sdf) * d(sdf)/dx """ import numpy as np from vde._vde import evaluate points = ctx.saved_tensors[0] node = ctx.node pts_np = points.detach().cpu().numpy().astype(np.float64) n = pts_np.shape[0] gradients = np.zeros((n, 3), dtype=np.float64) h = 1e-6 for i in range(n): p = pts_np[i] # Central finite differences for j in range(3): pp = p.copy() pm = p.copy() pp[j] += h pm[j] -= h gradients[i, j] = (evaluate(node, pp) - evaluate(node, pm)) / (2.0 * h) grad_input = torch.from_numpy(gradients).to(points.device).to(points.dtype) grad_node = None # Don't backprop to SDF parameters (for now) return grad_node, grad_input * grad_output.unsqueeze(-1) class SdfShapeFitter: """Fit SDF shapes to target point clouds with gradient descent. Usage: fitter = SdfShapeFitter(sphere(1.0)) optimized = fitter.fit(target_points, lr=0.01, epochs=100) """ def __init__(self, node): self.node = node def fit(self, target_points, lr=0.01, epochs=100, verbose=False): """Fit SDF shape to target point cloud. Args: target_points: torch.Tensor of shape (N, 3) lr: learning rate epochs: number of optimization steps verbose: print progress Returns: optimized node (SdfNode) """ import numpy as np from vde._vde import evaluate pts = target_points.detach().cpu().numpy().astype(np.float64) n = pts.shape[0] h = 1e-6 # Get current parameters params = self.node.params param_list = self._params_to_list(params) for epoch in range(epochs): # Forward: evaluate SDF at all points loss = 0.0 grads = {k: 0.0 for k in self._param_names()} for i in range(n): p = pts[i] d = evaluate(self.node, p) loss += abs(d) # Compute numerical gradient w.r.t. parameters for name, idx in self._param_indices().items(): orig_val = param_list[idx] param_list[idx] = orig_val + h self._list_to_params(param_list, self.node.params) dp = evaluate(self.node, p) param_list[idx] = orig_val - h self._list_to_params(param_list, self.node.params) dm = evaluate(self.node, p) param_list[idx] = orig_val grads[name] += (abs(dp) - abs(dm)) / (2.0 * h) self._list_to_params(param_list, self.node.params) loss /= n for name in grads: grads[name] /= n # Gradient descent step for name, idx in self._param_indices().items(): param_list[idx] -= lr * grads[name] self._list_to_params(param_list, self.node.params) if verbose and epoch % 10 == 0: print(f"Epoch {epoch}: loss = {loss:.6f}") return self.node def _param_names(self): return ['radius', 'extent_x', 'extent_y', 'extent_z', 'height'] def _param_indices(self): return { 'radius': 0, 'extent_x': 1, 'extent_y': 2, 'extent_z': 3, 'height': 4, } def _params_to_list(self, params): return [params.radius, params.extents.x, params.extents.y, params.extents.z, params.height] def _list_to_params(self, lst, params): params.radius = lst[0] params.extents.x = lst[1] params.extents.y = lst[2] params.extents.z = lst[3] params.height = lst[4]