Files
ViewDesignEngine/python/vde/torch_sdf.py
T
茂之钳 08b5854c56
CI / Build & Test (push) Failing after 16m59s
CI / Release Build (push) Failing after 42s
feat(sdf): S12-C PyTorch integration + Python bindings
Add batch evaluation/gradient functions for SDF trees (sdf_torch.h/cpp).
Add sdf_gradient.h with finite-difference gradient utilities.
Add Python modules vde.sdf (high-level SDF API) and vde.torch_sdf (PyTorch autograd).
Add test_sdf_torch_bridge.cpp with batch eval/gradient tests.
Integrate into vde_sdf library and test suite.
2026-07-24 07:20:45 +00:00

186 lines
5.8 KiB
Python

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