This project implements an interpretable manifold dynamics model where the learnable object is a torsion-free connection (Christoffel symbols) parameterized as low-degree polynomials of the state. Trajectories are generated by integrating the geodesic equation, and additional losses encourage the learned connection to be consistent with a (pseudo-)Riemannian metric via metric-compatibility reconstruction and path-independence regularization.
The implementation is PyTorch + CUDA, batched/parallel, with an optional jointly-trained autoencoder whose latent can be constrained to lie on a chosen manifold (sphere / hyperboloid / none).
pip install -r requirements.txtInstall PyTorch with CUDA using the official command for your platform, then install the remaining dependencies via requirements.txt.
Default demo uses a synthetic dataset: inertial motion in Cartesian coordinates expressed in polar coordinates (r, theta).
python train.py --use_ae 0 --manifold none --p 2 --q 0 --epochs 30python inspect_christoffel.py --ckpt checkpoints/latest.pt --use_ae 0 --manifold none --p 2 --q 0Save plots if desired:
python inspect_christoffel.py --ckpt checkpoints/latest.pt --save_fig figs/gamma.png-
train.pyMain training loop. Loads data, encodes to latent (if AE enabled), computes initial(x0, v0), rolls out geodesic dynamics under learned Christoffels, computes losses (trajectory + metric/geometry regularizers), backpropagates, and saves checkpoints. -
config.pyCentral hyperparameters and switches: latent dimensiond, polynomial degree, training schedule, loss weights, and pseudo-Riemannian signature(p, q)for metric validity penalties. -
utils.pySmall utilities: seeding, device selection, and angle unwrapping (needed for polar-coordinate trajectories to avoid discontinuities attheta).
data.pySynthetic dataset for sanity-check experiments. Default: inertial Cartesian motion represented in polar coordinates(r, theta); the induced Christoffels are known in closed form.
-
poly_features.pyDeterministic polynomial feature mapphi(x)(monomial basis up to degree 2 by default). This is the interpretability anchor: Christoffels are explicit polynomials in state. -
christoffel.pyThe Christoffel network. ComputesGamma^i_{jk}(x)as a linear combination of featuresphi(x)with learned coefficients. Enforces torsion-free symmetry (Gamma^i_{jk} = Gamma^i_{kj}) by construction. -
integrators.pyDifferentiable geodesic rollout. Integrates the second-order geodesic equation (converted to a first-order system) via a ResNet-style discrete update. Optionally applies manifold projection/retraction each step. -
metric_recon.pyMetric reconstruction by integrating the metric-compatibility equation along a path, yielding a reconstructedg(x)given a base metricg(x0)and learnedGamma(x). Supports multi-path reconstruction to quantify path dependence (loop loss). -
losses.pyLoss functions: trajectory loss + geometric regularizers: loop loss, metric symmetry, nondegeneracy (log|det g| barrier), fixed signature penalty (pseudo-Riemannian), coefficient weight decay.
-
autoencoder.pyOptional MLP autoencoder trained jointly with the manifold dynamics. Latents are projected onto a chosen manifold viamanifold.py. -
manifold.pyManifold constraint operators:project(z)retraction/projection to the manifoldtangent_project(z, v)projects velocities onto the tangent space (optional stabilization)
-
inspect_christoffel.pyEvaluates learned Christoffel components on a grid and compares to ground truth for the polar demo:Gamma^r_{theta,theta} = -rGamma^theta_{r,theta} = Gamma^theta_{theta,r} = 1/r
-
inspect_metric.pyReconstructs the metric on a grid and compares to the polar target metricdiag(1, r^2)with optional loop error maps. -
inspect_metric_ae.pyPulls back the reconstructed latent metric via the encoder Jacobian for AE experiments.
Let x(t) in R^d be the state in a coordinate chart. Two modes:
- Direct-state mode:
x(t)are observed coordinates. - Autoencoder mode: you observe
y(t)and setx(t) = E(y(t))(latent coordinates).
Goal: model time evolution as geodesic flow governed by a learned connection.
Define monomial features phi(x) up to degree p:
phi(x) = [1, x_1, ..., x_d, x_1^2, x_1 x_2, ..., x_d^2] in R^P
Learn Christoffels as explicit polynomials:
Gamma^i_{jk}(x) = sum_{m=1}^P C^i_{jk,m} phi_m(x)
with trainable coefficients C^i_{jk,m}.
Enforce symmetry in the lower indices:
Gamma^i_{jk}(x) = Gamma^i_{kj}(x)
This makes the connection a plausible Levi-Civita candidate and reduces degrees of freedom.
Geodesics satisfy:
d2 x^i / dt^2 + Gamma^i_{jk}(x) dx^j/dt dx^k/dt = 0
Let v = dx/dt. Then the first-order system is:
dx/dt = v
dv^i/dt = -Gamma^i_{jk}(x) v^j v^k
Define acceleration:
a^i(x, v) = -Gamma^i_{jk}(x) v^j v^k
Interpretability: acceleration is quadratic in velocity and structured entirely by Gamma(x).
We discretize the geodesic ODE using a simple differentiable update:
v_{t+1} = v_t + dt * a(x_t, v_t)
x_{t+1} = x_t + dt * v_{t+1}
This is equivalent to stacking residual blocks that implement the geodesic flow.
If latents must lie on a manifold M subset R^d, apply a projection/retraction:
x <- Pi_M(x)
Supported:
- none:
Pi(x) = x - sphere:
Pi(x) = x / ||x||enforcing||x|| = 1 - hyperboloid (Lorentz model): enforce
<x, x>_L = -1withx_0 > 0
Optionally project velocity to the tangent space:
- sphere:
v <- v - <v, x> x - hyperboloid:
v <- v + <v, x>_L x(since<x, x>_L = -1)
This keeps discrete rollouts on the constraint manifold.
A torsion-free connection Gamma is not necessarily Levi-Civita for any metric. To encourage (pseudo-)Riemannian structure, reconstruct a metric g(x) compatible with Gamma(x).
partial_i g_{jk} = Gamma^l_{ij} g_{lk} + Gamma^l_{ik} g_{jl}
This is a first-order linear PDE. We convert it to a path ODE by integrating along a curve x(s) from basepoint x_0 to x:
d/ds g_{jk}(s) =
(Gamma^l_{ij}(x(s)) g_{lk}(s) + Gamma^l_{ik}(x(s)) g_{jl}(s)) * dx^i/ds
Given:
- basepoint
(x_0) - base metric
g(x_0) = g_0(parameterized to have fixed signature) - a chosen path
we integrate to get a reconstructed g(x).
If Gamma is not metric-realizable, reconstructed g(x) depends on path. We penalize this by using two paths:
- A:
(x_0 -> x) - B:
(x_0 -> x_m -> x)
L_loop = || g^(A)(x) - g^(B)(x) ||_F^2
This provides extra training signal pushing Gamma toward metrizable behavior.
- Direct-state mode:
L_traj = (1 / (B T)) sum_{b,t} || x_hat_{b,t} - x_{b,t} ||^2
- AE mode (reconstruction in observation space):
L_recon = (1 / (B T)) sum_{b,t} || D(z_hat_{b,t}) - y_{b,t} ||^2
Given reconstructed g(x):
- symmetry:
L_sym = || g - g^T ||^2 - nondegeneracy barrier:
L_det = softplus(alpha - log|det g|) - fixed signature: penalize eigenvalues that violate desired
(p, q)
Coefficient weight decay:
L_wd = || C ||_2^2
L =
L_traj/recon
+ lambda_loop L_loop
+ lambda_sym L_sym
+ lambda_det L_det
+ lambda_sig L_sig
+ lambda_wd L_wd
For each batch:
- Encode (optional):
x_{0:T} = E(y_{0:T}) - Initialize:
x_0 = x(0),v_0 approx (x_1 - x_0) / dt - Rollout: integrate geodesic dynamics under
Gamma(x)to getx_hat_{0:T} - Decode (optional):
y_hat_{0:T} = D(x_hat_{0:T}) - Reconstruct metric
g(x)along paths and compute loop/validity losses - Backprop to update
Gammacoefficients (and AE / base metric parameters if enabled)
- Start with the polar-coordinate synthetic benchmark (no AE, SPD signature).
- Confirm
inspect_christoffel.pyshowsGamma^r_{theta theta} approx -randGamma^theta_{r theta} approx 1/r. - Enable AE and then move to more complex observation models (e.g., CNN encoders).
- Swap integrator for RK4 or differentiable ODE solvers (
torchdiffeq) after pipeline validation. - Co-train an explicit metric network
g(x)and enforceGamma_poly approx Gamma(g)(strong Levi-Civita guarantee). - Apply to domains with strong existing autoencoders (e.g., CNN AEs) to learn curved latent dynamics in concept/feature latent spaces.