Learning API#
The learning namespace provides validated data adaptation, differentiable geometry primitives, supervised and semi-supervised prediction, intrinsic statistics, robust and scalable summaries, clustering, inference, transport, dimension reduction, and metric learning. High-level algorithms operate on dense distance matrices unless their documentation says otherwise.
Data and capability contracts#
- class geojax.learning.ManifoldData(manifold, values, n_samples, batch_shape, event_shapes, report)#
Bases:
objectCanonical observations bound to the geometry that validated them.
Binding prevents a validated dataset from being silently reused under a different metric or representation with the same event shape. Reusing the object with its original geometry skips checks that are already at least as strong as the requested validation level.
- Parameters:
manifold (Any)
values (Any)
n_samples (int)
batch_shape (tuple[int, ...])
event_shapes (Any)
report (DataValidationReport)
- class geojax.learning.DataValidationReport(valid, check, n_samples, batch_shape, invalid_count=0, repaired_count=0, messages=())#
Bases:
objectEager validation summary for one adapted manifold dataset.
- Parameters:
valid (bool)
check (str)
n_samples (int)
batch_shape (tuple[int, ...])
invalid_count (int)
repaired_count (int)
messages (tuple[str, ...])
- class geojax.learning.ManifoldDataAdapterProtocol(*args, **kwargs)#
Bases:
ProtocolCallable protocol for user-registered representation adapters.
- class geojax.learning.EquivariantEmbeddingProtocol(*args, **kwargs)#
Bases:
ProtocolGeometry protocol for algorithms that require Euclidean embeddings.
- embed(x)#
Map a represented point to equivariant Euclidean coordinates.
- Parameters:
x (Any)
- Return type:
Any
- exception geojax.learning.LearningCapabilityError#
Bases:
ValueErrorRaised when an algorithm needs unavailable exact geometry operations.
- geojax.learning.as_manifold_data(manifold, values, *, sample_axis=None, representation='canonical', check='belongs', repair=False)#
Convert manifold observations to the canonical learning-data layout.
sample_axis=Nonedenotes the axis immediately before each geometry’s event dimensions. Product representations and axes may be pytrees matchingmanifold.factors. Python sequences of complete points requirerepresentation='point_sequence'so their interpretation is explicit.- Parameters:
manifold (Any)
values (Any)
sample_axis (Any)
representation (Any)
check (str)
repair (bool)
- Return type:
- geojax.learning.check_manifold_data(manifold, values, *, sample_axis=None, representation='canonical', check='belongs')#
Return a validation report without propagating data-validation errors.
- Parameters:
manifold (Any)
values (Any)
sample_axis (Any)
representation (Any)
check (str)
- Return type:
- geojax.learning.register_manifold_data_adapter(geometry_type, representation, adapter, *, overwrite=False)#
Register an explicit representation converter for a geometry class.
- Parameters:
geometry_type (type)
representation (str)
adapter (ManifoldDataAdapterProtocol)
overwrite (bool)
- Return type:
None
Geometric primitives#
- geojax.learning.pairwise_distances(manifold, x, y=None, *, squared=False, block_size=None)#
Return all pairwise exact distances between two point collections.
Collections use
batch_shape + (n_samples,) + event_shape. Product collections use the factor pytree and share their sample and batch axes.block_sizelimits the number of right-hand samples materialized in one geometry call while retaining a dense result.- Parameters:
manifold (Any)
x (Any)
y (Any | None)
squared (bool)
block_size (int | None)
- Return type:
Any
- geojax.learning.geodesic_interpolation(manifold, x, y, t)#
Evaluate
Exp_x(t Log_x(y))on the selected exact geodesic.- Parameters:
manifold (Any)
x (Any)
y (Any)
t (Any)
- Return type:
Any
- geojax.learning.tangent_space_map(source, target, x, *, source_base, target_base, transform)#
Apply a user transform between exact source and target tangent spaces.
- Parameters:
source (Any)
target (Any)
x (Any)
source_base (Any)
target_base (Any)
transform (Callable[[Any], Any])
- Return type:
Any
- geojax.learning.nearest_neighbors(manifold, data, queries=None, *, n_neighbors=5, exclude_self=True, block_size=None)#
Find exact-distance nearest neighbors in a dense manifold dataset.
- Parameters:
manifold (Any)
data (Any)
queries (Any | None)
n_neighbors (int)
exclude_self (bool)
block_size (int | None)
- Return type:
- class geojax.learning.NeighborsResult(distances, indices)#
Bases:
objectDistances and training-set indices for each nearest-neighbor query.
- Parameters:
distances (Any)
indices (Any)
Statistics and scalar-response regression#
- geojax.learning.frechet_mean(manifold, data, *, sample_weight=None, initial_point=None, solver=None, maxiter=200, tol=1e-07)#
Compute a local weighted Fréchet-mean minimizer.
The optimized objective is
sum_i w_i d(x, x_i)^2. On a general manifold it need not be geodesically convex, so convergence certifies a stationary local solution rather than a unique global mean.- Parameters:
manifold (Any)
data (Any)
sample_weight (Any | None)
initial_point (Any | None)
solver (Any | None)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.frechet_median(manifold, data, *, sample_weight=None, initial_point=None, smoothing=1e-08, maxiter=200, tol=1e-07)#
Compute a Huber-smoothed weighted geometric median.
The guarded Weiszfeld iteration minimizes
sum_i w_i rho_s(d(x, x_i)), whererho_s(r)=rforr >= sandrho_s(r)=r^2/(2s)+s/2otherwise. Thussmoothingcontrols the local approximation to the nonsmooth Fréchet-median objective.- Parameters:
manifold (Any)
data (Any)
sample_weight (Any | None)
initial_point (Any | None)
smoothing (float)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.minimum_enclosing_ball(manifold, data, *, initial_point=None, maxiter=500, tol=1e-07)#
Approximate the smallest enclosing geodesic ball by farthest-point updates.
- Parameters:
manifold (Any)
data (Any)
initial_point (Any | None)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.kernel_regression(manifold, data, targets, *, bandwidth, kernel=None)#
Fit Nadaraya-Watson regression with manifold-valued predictors.
- Parameters:
manifold (Any)
data (Any)
targets (Any)
bandwidth (float)
kernel (Callable[[Any, float], Any] | None)
- Return type:
- geojax.learning.select_kernel_bandwidth(manifold, data, targets, bandwidths, *, n_folds=5, key, kernel=None)#
Select a kernel bandwidth by deterministic-key K-fold mean squared error.
- Parameters:
manifold (Any)
data (Any)
targets (Any)
bandwidths (Any)
n_folds (int)
key (Any | int | None)
kernel (Callable[[Any, float], Any] | None)
- Return type:
- class geojax.learning.FrechetMeanResult(point, objective, gradient_norm, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectA local Fréchet-mean fit with objective and stationarity diagnostics.
- Parameters:
point (Any)
objective (Any)
gradient_norm (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- class geojax.learning.FrechetMedianResult(point, objective, gradient_norm, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectA Huber-smoothed intrinsic median fit and its terminal residual.
- Parameters:
point (Any)
objective (Any)
gradient_norm (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- class geojax.learning.EnclosingBallResult(center, radius, objective, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectA farthest-point enclosing-ball approximation, not a global certificate.
- Parameters:
center (Any)
radius (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- class geojax.learning.KernelRegressionModel(manifold, training_data, targets, bandwidth, kernel)#
Bases:
objectA fitted distance-kernel model for scalar Euclidean responses.
- Parameters:
manifold (Any)
training_data (Any)
targets (Any)
bandwidth (float)
kernel (Callable[[...], Any] | None)
- predict(data)#
Predict responses at canonical or adaptable manifold observations.
- Parameters:
data (Any)
- Return type:
Any
- class geojax.learning.KernelCVResult(model, bandwidth, scores, diagnostics=<factory>)#
Bases:
objectSelected kernel model, bandwidth, and validation scores.
- Parameters:
model (KernelRegressionModel)
bandwidth (float)
scores (Any)
diagnostics (Mapping[str, Any])
Supervised classification#
- geojax.learning.nearest_centroid_classifier(manifold, data, labels, *, sample_weight=None, maxiter=200, tol=1e-07)#
Fit one intrinsic Fréchet centroid per class.
- Parameters:
manifold (Any)
data (Any)
labels (Any)
sample_weight (Any | None)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.knn_classifier(manifold, data, labels, *, n_neighbors=5, weights='uniform')#
Fit a geodesic-distance k-nearest-neighbors classifier.
- Parameters:
manifold (Any)
data (Any)
labels (Any)
n_neighbors (int)
weights (str)
- Return type:
- geojax.learning.tangent_space_logistic_regression(manifold, data, labels, *, base_point=None, n_components=None, regularization=0.001, maxiter=500, tol=1e-07, learning_rate=1.0)#
Fit multinomial logistic regression in intrinsic tangent coordinates.
- Parameters:
manifold (Any)
data (Any)
labels (Any)
base_point (Any | None)
n_components (int | None)
regularization (float)
maxiter (int)
tol (float)
learning_rate (float)
- Return type:
- geojax.learning.tangent_space_discriminant_analysis(manifold, data, labels, *, method='lda', base_point=None, n_components=None, regularization=0.0001, priors=None)#
Fit LDA or QDA in intrinsic metric-orthonormal tangent coordinates.
- Parameters:
manifold (Any)
data (Any)
labels (Any)
method (str)
base_point (Any | None)
n_components (int | None)
regularization (float)
priors (Any | None)
- Return type:
- class geojax.learning.NearestCentroidModel(manifold, classes, centers, converged=True, reason='all class centroids converged', diagnostics=<factory>)#
Bases:
objectIntrinsic class centroids and their internal mean-fit status.
predict_probareturns normalized Gibbs distance scores. They are not calibrated posterior probabilities.- Parameters:
manifold (Any)
classes (Any)
centers (Any)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- predict(data)#
Predict the class of the closest intrinsic centroid.
- Parameters:
data (Any)
- Return type:
Any
- predict_proba(data)#
Return normalized, query-scaled Gibbs distance scores.
- Parameters:
data (Any)
- Return type:
Any
- class geojax.learning.KNearestNeighborsModel(manifold, training_data, classes, encoded_labels, n_neighbors, weights, diagnostics=<factory>)#
Bases:
objectA fitted geodesic-distance nearest-neighbors classifier.
- Parameters:
manifold (Any)
training_data (Any)
classes (Any)
encoded_labels (Any)
n_neighbors (int)
weights (str)
diagnostics (Mapping[str, Any])
- predict(data)#
Predict labels by uniform or inverse-distance voting.
- Parameters:
data (Any)
- Return type:
Any
- predict_proba(data)#
Return normalized class vote weights.
- Parameters:
data (Any)
- Return type:
Any
- class geojax.learning.TangentFeatureMap(manifold, base_point, basis, eigenvalues, converged=True, reason='reference point supplied', diagnostics=<factory>)#
Bases:
objectA metric-orthonormal coordinate chart at a fitted reference point.
- Parameters:
manifold (Any)
base_point (Any)
basis (tuple[Any, ...])
eigenvalues (Any)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- transform(data)#
Map manifold observations to the retained tangent coordinates.
- Parameters:
data (Any)
- Return type:
Any
- class geojax.learning.TangentSpaceClassifierModel(manifold, classes, feature_map, method, coefficients=None, intercept=None, location=None, scale=None, class_means=None, covariances=None, priors=None, objective=None, iterations=0, converged=True, reason='closed-form fit', diagnostics=<factory>)#
Bases:
objectA logistic, LDA, or QDA classifier in one intrinsic tangent chart.
- Parameters:
manifold (Any)
classes (Any)
feature_map (TangentFeatureMap)
method (str)
coefficients (Any)
intercept (Any)
location (Any)
scale (Any)
class_means (Any)
covariances (Any)
priors (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- predict(data)#
Predict labels after tangent-chart feature extraction.
- Parameters:
data (Any)
- Return type:
Any
- predict_proba(data)#
Return normalized class scores from the fitted tangent model.
- Parameters:
data (Any)
- Return type:
Any
Manifold-valued response regression#
- geojax.learning.geodesic_regression(manifold, predictors, responses, *, sample_weight=None, initial_point=None, solver=None, maxiter=200, tol=1e-07)#
Fit
Y(t) = Exp_p((t - t_bar) v)by joint intrinsic least squares.The intercept and an ambient parameterization of its tangent slope are optimized jointly. The slope is projected into
T_p Minside the objective, so this minimizes the stated nonlinear residual rather than a flat-space profile approximation.- Parameters:
manifold (Any)
predictors (Any)
responses (Any)
sample_weight (Any | None)
initial_point (Any | None)
solver (Any | None)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.local_polynomial_regression(manifold, predictors, responses, *, bandwidth, degree=1, kernel=None, maxiter=100, tol=1e-06)#
Fit local-constant or local-linear Fréchet regression.
- Parameters:
manifold (Any)
predictors (Any)
responses (Any)
bandwidth (float)
degree (int)
kernel (Callable[[Any, float], Any] | None)
maxiter (int)
tol (float)
- Return type:
- class geojax.learning.GeodesicRegressionModel(manifold, intercept, slope, predictor_mean, objective, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectA fitted one-predictor intrinsic geodesic regression curve.
- Parameters:
manifold (Any)
intercept (Any)
slope (Any)
predictor_mean (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- predict(predictors)#
Evaluate the fitted geodesic at scalar predictor values.
- Parameters:
predictors (Any)
- Return type:
Any
- class geojax.learning.LocalPolynomialRegressionModel(manifold, predictors, training_data, bandwidth, degree, kernel, maxiter, tol, diagnostics=<factory>)#
Bases:
objectA local-constant or local-linear manifold-response smoother.
- Parameters:
manifold (Any)
predictors (Any)
training_data (Any)
bandwidth (float)
degree (int)
kernel (Callable[[...], Any] | None)
maxiter (int)
tol (float)
diagnostics (Mapping[str, Any])
- predict(predictors)#
Solve the local Fréchet problem at each scalar query.
- Parameters:
predictors (Any)
- Return type:
Any
Inference#
- geojax.learning.frechet_anova(manifold, data, groups, *, method='asymptotic', n_permutations=999, key=None, maxiter=100, tol=1e-06, variance_floor=1e-12)#
Test equality of metric-space populations using Dubey-Mueller FANOVA.
- Parameters:
manifold (Any)
data (Any)
groups (Any)
method (str)
n_permutations (int)
key (Any | int | None)
maxiter (int)
tol (float)
variance_floor (float)
- Return type:
- geojax.learning.biswas_ghosh_two_sample_test(manifold, x, y, *, n_permutations=999, key)#
Run the metric-space modification of the Biswas-Ghosh two-sample test.
- Parameters:
manifold (Any)
x (Any)
y (Any)
n_permutations (int)
key (Any | int | None)
- Return type:
- geojax.learning.wasserstein_two_sample_test(manifold, x, y, *, p=2.0, n_permutations=999, key, tolerance=1e-10)#
Permutation test using exact empirical Wasserstein distance.
- Parameters:
manifold (Any)
x (Any)
y (Any)
p (float)
n_permutations (int)
key (Any | int | None)
tolerance (float)
- Return type:
- geojax.learning.bootstrap_frechet_mean(manifold, data, *, sample_weight=None, n_bootstrap=999, confidence_level=0.95, key, maxiter=100, tol=1e-06)#
Bootstrap an intrinsic mean and return a geodesic confidence ball.
- Parameters:
manifold (Any)
data (Any)
sample_weight (Any | None)
n_bootstrap (int)
confidence_level (float)
key (Any | int | None)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.energy_two_sample_test(manifold, x, y, *, n_permutations=999, key)#
Run a biased metric energy-statistic permutation test.
The statistic is guaranteed nonnegative at the population level only when the metric is of negative type. GeoJAX therefore preserves the signed finite-sample V-statistic instead of clipping it, which also preserves the exact permutation ordering for arbitrary manifold metrics.
- Parameters:
manifold (Any)
x (Any)
y (Any)
n_permutations (int)
key (Any | int | None)
- Return type:
- geojax.learning.kernel_mmd_two_sample_test(manifold, x, y, *, bandwidth=None, kernel=None, check_psd=True, psd_tolerance=1e-08, n_permutations=999, key)#
Run a finite-sample PSD-kernel maximum mean discrepancy test.
- Parameters:
manifold (Any)
x (Any)
y (Any)
bandwidth (float | None)
kernel (Callable[[Any], Any] | None)
check_psd (bool)
psd_tolerance (float)
n_permutations (int)
key (Any | int | None)
- Return type:
- geojax.learning.paired_frechet_test(manifold, x, y, *, n_permutations=999, key, maxiter=100, tol=1e-06)#
Test a zero mean paired displacement by within-pair random sign flips.
Pairwise exchangeability of
(x_i, y_i)under the null gives exact randomization calibration: swapping a pair leaves the pooled base fit unchanged and negates that pair’s tangent displacement. Central symmetry of the tangent displacements is a weaker modeling route to the same sign-flip invariance.- Parameters:
manifold (Any)
x (Any)
y (Any)
n_permutations (int)
key (Any | int | None)
maxiter (int)
tol (float)
- Return type:
- class geojax.learning.HypothesisTestResult(statistic, pvalue, null_distribution, method, diagnostics=<factory>)#
Bases:
objectObserved statistic, calibrated p-value, and simulated null statistics.
- Parameters:
statistic (Any)
pvalue (Any)
null_distribution (Any)
method (str)
diagnostics (Mapping[str, Any])
- class geojax.learning.BootstrapResult(estimate, replicates, confidence_radius, confidence_level, diagnostics=<factory>)#
Bases:
objectA point estimate, bootstrap replicates, and geodesic confidence radius.
- Parameters:
estimate (Any)
replicates (Any)
confidence_radius (Any)
confidence_level (float)
diagnostics (Mapping[str, Any])
Clustering#
- geojax.learning.kmeans(manifold, data, *, n_clusters, key=None, sample_weight=None, init='kmeans++', n_init=1, maxiter=100, tol=1e-06, center_maxiter=100)#
Run weighted intrinsic Lloyd clustering with deterministic-key initialization.
- Parameters:
manifold (Any)
data (Any)
n_clusters (int)
key (Any | int | None)
sample_weight (Any | None)
init (str | Any)
n_init (int)
maxiter (int)
tol (float)
center_maxiter (int)
- Return type:
- geojax.learning.lightweight_coreset(manifold, data, *, size, key, sample_weight=None)#
Sample the lightweight-coreset sensitivity heuristic on a manifold.
- Parameters:
manifold (Any)
data (Any)
size (int)
key (Any | int | None)
sample_weight (Any | None)
- Return type:
- geojax.learning.kmedoids(manifold, data, *, n_clusters, key, sample_weight=None, maxiter=100)#
Cluster using exact sample medoids and arbitrary manifold distances.
- Parameters:
manifold (Any)
data (Any)
n_clusters (int)
key (Any | int | None)
sample_weight (Any | None)
maxiter (int)
- Return type:
- geojax.learning.agglomerative_clustering(manifold, data, *, n_clusters=2, linkage='average')#
Perform dense single, complete, or average-linkage clustering.
- Parameters:
manifold (Any)
data (Any)
n_clusters (int)
linkage (str)
- Return type:
- geojax.learning.spectral_clustering(manifold, data, *, n_clusters, key, affinity='rbf', bandwidth=None, n_neighbors=7, laplacian='symmetric', maxiter=100)#
Cluster an exact-distance affinity graph through a Laplacian embedding.
- Parameters:
manifold (Any)
data (Any)
n_clusters (int)
key (Any | int | None)
affinity (str)
bandwidth (float | None)
n_neighbors (int)
laplacian (str)
maxiter (int)
- Return type:
- geojax.learning.mean_shift(manifold, data, *, bandwidth, sample_weight=None, maxiter=100, tol=1e-06, merge_tol=None)#
Find modes by Gaussian-kernel geodesic mean-shift updates.
- Parameters:
manifold (Any)
data (Any)
bandwidth (float)
sample_weight (Any | None)
maxiter (int)
tol (float)
merge_tol (float | None)
- Return type:
- geojax.learning.competitive_quantization(manifold, data, *, n_clusters, key, epochs=10, initial_gain=0.5, decay=0.01, tol=1e-06)#
Run competitive learning Riemannian quantization (CLRQ).
- Parameters:
manifold (Any)
data (Any)
n_clusters (int)
key (Any | int | None)
epochs (int)
initial_gain (float)
decay (float)
tol (float)
- Return type:
- class geojax.learning.ClusteringResult(labels, centers, objective, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectCluster assignments, representatives, objective, and convergence status.
- Parameters:
labels (Any)
centers (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- class geojax.learning.HierarchicalClusteringResult(labels, linkage, objective, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectFlat labels and the merge table from metric-compatible agglomeration.
- Parameters:
labels (Any)
linkage (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- class geojax.learning.CoresetResult(indices, points, weights, diagnostics=<factory>)#
Bases:
objectSampled observations and normalized importance weights for a coreset.
- Parameters:
indices (Any)
points (Any)
weights (Any)
diagnostics (Mapping[str, Any])
Scalable summaries#
- geojax.learning.streaming_frechet_mean(manifold, data, *, sample_weight=None, initial_point=None, initial_weight=0.0)#
Compute the one-pass inductive Fréchet mean in observation order.
initial_pointinfluences the estimate only wheninitial_weightis positive, making any prior mass explicit rather than silently counting the starting point as an extra observation.- Parameters:
manifold (Any)
data (Any)
sample_weight (Any | None)
initial_point (Any | None)
initial_weight (float)
- Return type:
- geojax.learning.minibatch_frechet_mean(manifold, data, *, batch_size=32, epochs=10, key, sample_weight=None, initial_point=None, learning_rate=1.0, decay=0.1, tol=1e-06)#
Approximate a Fréchet mean by shuffled mini-batch log-map updates.
- Parameters:
manifold (Any)
data (Any)
batch_size (int)
epochs (int)
key (Any | int | None)
sample_weight (Any | None)
initial_point (Any | None)
learning_rate (float)
decay (float)
tol (float)
- Return type:
- geojax.learning.minibatch_kmeans(manifold, data, *, n_clusters, batch_size=32, epochs=10, key, sample_weight=None, learning_rate=0.5, decay=0.01, tol=1e-06)#
Run shuffled mini-batch intrinsic k-means center updates.
- Parameters:
manifold (Any)
data (Any)
n_clusters (int)
batch_size (int)
epochs (int)
key (Any | int | None)
sample_weight (Any | None)
learning_rate (float)
decay (float)
tol (float)
- Return type:
Barycentric coding and dictionaries#
- geojax.learning.geodesic_barycentric_coding(manifold, data, atoms, *, ridge=1e-06, maxiter=200, tol=1e-07, reconstruction_maxiter=100)#
Code points by simplex weights minimizing a log-map barycentric residual.
- Parameters:
manifold (Any)
data (Any)
atoms (Any)
ridge (float)
maxiter (int)
tol (float)
reconstruction_maxiter (int)
- Return type:
- geojax.learning.manifold_dictionary_learning(manifold, data, *, n_atoms, key=None, initial_atoms=None, sample_weight=None, ridge=1e-06, maxiter=20, coding_maxiter=100, center_maxiter=100, tol=1e-05)#
Alternate intrinsic barycentric codes and weighted atom updates.
- Parameters:
manifold (Any)
data (Any)
n_atoms (int)
key (Any | int | None)
initial_atoms (Any | None)
sample_weight (Any | None)
ridge (float)
maxiter (int)
coding_maxiter (int)
center_maxiter (int)
tol (float)
- Return type:
- class geojax.learning.BarycentricCodingResult(codes, reconstructions, objective, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectSimplex codes, intrinsic reconstructions, and solver diagnostics.
- Parameters:
codes (Any)
reconstructions (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- class geojax.learning.DictionaryLearningResult(atoms, codes, reconstructions, objective, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectLearned manifold atoms, barycentric codes, and alternating-fit status.
- Parameters:
atoms (Any)
codes (Any)
reconstructions (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
Robust analysis#
- geojax.learning.trimmed_frechet_mean(manifold, data, *, trim_fraction=0.1, sample_weight=None, initial_point=None, maxiter=100, center_maxiter=100, tol=1e-06)#
Compute a least-trimmed-squares intrinsic location estimate.
- Parameters:
manifold (Any)
data (Any)
trim_fraction (float)
sample_weight (Any | None)
initial_point (Any | None)
maxiter (int)
center_maxiter (int)
tol (float)
- Return type:
- geojax.learning.geodesic_m_estimator(manifold, data, *, loss='huber', scale=None, sample_weight=None, initial_point=None, maxiter=100, center_maxiter=100, tol=1e-06)#
Compute a geodesic M-location by iteratively reweighted Fréchet means.
- Parameters:
manifold (Any)
data (Any)
loss (str)
scale (float | None)
sample_weight (Any | None)
initial_point (Any | None)
maxiter (int)
center_maxiter (int)
tol (float)
- Return type:
- geojax.learning.geodesic_spatial_depth(manifold, points, reference_data, *, sample_weight=None)#
Evaluate intrinsic spatial depth relative to a reference sample.
- Parameters:
manifold (Any)
points (Any)
reference_data (Any)
sample_weight (Any | None)
- Return type:
Any
- geojax.learning.metric_distance_ranks(manifold, data, *, center=None, sample_weight=None, maxiter=100, tol=1e-06)#
Rank observations by geodesic distance from an intrinsic median.
- Parameters:
manifold (Any)
data (Any)
center (Any | None)
sample_weight (Any | None)
maxiter (int)
tol (float)
- Return type:
- class geojax.learning.RobustLocationResult(point, objective, gradient_norm, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectA robust intrinsic location estimate and terminal stationarity residual.
- Parameters:
point (Any)
objective (Any)
gradient_norm (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- class geojax.learning.MetricRanksResult(ranks, scores, center, diagnostics=<factory>)#
Bases:
objectMidranks of center distances together with their underlying scores.
- Parameters:
ranks (Any)
scores (Any)
center (Any)
diagnostics (Mapping[str, Any])
Semi-supervised learning#
- geojax.learning.label_propagation(manifold, data, labels, *, unlabeled=-1, bandwidth=None, n_neighbors=None, alpha=0.95, maxiter=1000, tol=1e-07)#
Propagate categorical labels over a geodesic-distance affinity graph.
- Parameters:
manifold (Any)
data (Any)
labels (Any)
unlabeled (Any)
bandwidth (float | None)
n_neighbors (int | None)
alpha (float)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.manifold_regularized_regression(manifold, data, targets, *, labeled_mask=None, bandwidth=None, n_neighbors=None, ambient_regularization=0.001, intrinsic_regularization=1.0)#
Fit transductive squared-loss regression with graph-Laplacian regularization.
- Parameters:
manifold (Any)
data (Any)
targets (Any)
labeled_mask (Any | None)
bandwidth (float | None)
n_neighbors (int | None)
ambient_regularization (float)
intrinsic_regularization (float)
- Return type:
- class geojax.learning.SemiSupervisedResult(predictions, scores, objective, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectTransductive predictions, vertex scores, and graph-solver diagnostics.
- Parameters:
predictions (Any)
scores (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
Dimension reduction#
- geojax.learning.classical_mds(manifold, data, *, n_components=2)#
Classical scaling of the exact manifold distance matrix.
- Parameters:
manifold (Any)
data (Any)
n_components (int)
- Return type:
- geojax.learning.principal_geodesic_analysis(manifold, data, *, n_components=2, mean=None, maxiter=200, tol=1e-07)#
Perform tangent PCA using the Riemannian metric at a Fréchet mean.
- Parameters:
manifold (Any)
data (Any)
n_components (int)
mean (Any | None)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.kernel_pca(manifold, data, *, n_components=2, bandwidth=None, kernel=None, allow_indefinite=False)#
Kernel PCA using an RBF manifold-distance kernel or user callable.
A genuine kernel-PCA covariance operator requires a positive-semidefinite centered Gram matrix. Set
allow_indefinite=Trueto request the explicit positive-spectral-part approximation for an indefinite similarity.- Parameters:
manifold (Any)
data (Any)
n_components (int)
bandwidth (float | None)
kernel (Callable[[Any, float], Any] | None)
allow_indefinite (bool)
- Return type:
- geojax.learning.isomap(manifold, data, *, n_components=2, n_neighbors=7, mutual=True, disconnected='error')#
Isomap with a dense exact-distance neighbor graph.
- Parameters:
manifold (Any)
data (Any)
n_components (int)
n_neighbors (int)
mutual (bool)
disconnected (str)
- Return type:
- geojax.learning.sammon_mapping(manifold, data, *, n_components=2, maxiter=300, tol=1e-07)#
Optimize Sammon stress from a classical-MDS initialization.
- Parameters:
manifold (Any)
data (Any)
n_components (int)
maxiter (int)
tol (float)
- Return type:
- geojax.learning.tsne(manifold, data, *, n_components=2, perplexity=30.0, key, maxiter=1000, learning_rate=None, early_exaggeration=12.0, exaggeration_iterations=250, tol=1e-07)#
Dense t-SNE from exact manifold distances with explicit random state.
- Parameters:
manifold (Any)
data (Any)
n_components (int)
perplexity (float)
key (Any | int | None)
maxiter (int)
learning_rate (float | None)
early_exaggeration (float)
exaggeration_iterations (int)
tol (float)
- Return type:
- geojax.learning.phate(manifold, data, *, n_components=2, n_neighbors=5, decay=40.0, diffusion_time=None, max_diffusion_time=50, potential='log')#
Compute a dense PHATE-style diffusion-potential embedding.
This dependency-free implementation follows PHATE’s adaptive affinity and diffusion-potential construction, followed by classical scaling. The reference implementation instead uses metric MDS; classical scaling and the deterministic maximum-curvature diffusion-time rule are documented approximations in GeoJAX.
- Parameters:
manifold (Any)
data (Any)
n_components (int)
n_neighbors (int)
decay (float)
diffusion_time (int | None)
max_diffusion_time (int)
potential (str)
- Return type:
- class geojax.learning.EmbeddingResult(coordinates, objective, iterations, converged, reason, model=None, diagnostics=<factory>)#
Bases:
objectEuclidean coordinates and method-specific fit or spectral diagnostics.
- Parameters:
coordinates (Any)
objective (Any)
iterations (int)
converged (bool)
reason (str)
model (Any)
diagnostics (Mapping[str, Any])
Transport and metric learning#
- geojax.learning.empirical_wasserstein_distance(manifold, x, y, *, p=2.0, weights_x=None, weights_y=None, tolerance=1e-10, max_pivots=10000)#
Compute exact weighted empirical Wasserstein distance by transportation simplex.
- Parameters:
manifold (Any)
x (Any)
y (Any)
p (float)
weights_x (Any | None)
weights_y (Any | None)
tolerance (float)
max_pivots (int)
- Return type:
- geojax.learning.sinkhorn_divergence(manifold, x, y, *, epsilon=0.05, p=2.0, weights_x=None, weights_y=None)#
Return debiased entropic transport divergence through optional OTT-JAX.
- Parameters:
manifold (Any)
x (Any)
y (Any)
epsilon (float)
p (float)
weights_x (Any | None)
weights_y (Any | None)
- Return type:
Any
- class geojax.learning.TransportResult(distance, cost, plan, iterations, converged, reason, diagnostics=<factory>)#
Bases:
objectTransport distance, powered cost, coupling, and optimality diagnostics.
- Parameters:
distance (Any)
cost (Any)
plan (Any)
iterations (int)
converged (bool)
reason (str)
diagnostics (Mapping[str, Any])
- geojax.learning.riemannian_metric_learning(manifold, data, labels, *, regularization=0.1, balance=0.5, embedding=None, eigenvalue_floor=1e-10)#
Fit the regularized log-Euclidean RMML closed form.
For similar- and dissimilar-pair scatter matrices
SandD, this implements Equation (21) of Zhu et al. (2018),A = exp((-balance * log(S) + (1 - balance) * log(D)) / 2).The derivation is invariant only when
embeddingis an appropriate equivariant embedding for the supplied geometry. An arbitrary callable still defines a valid Euclidean-feature Mahalanobis model, but it does not inherit that Riemannian invariance automatically.- Parameters:
manifold (Any)
data (Any)
labels (Any)
regularization (float)
balance (float)
embedding (Callable[[Any], Any] | None)
eigenvalue_floor (float)
- Return type:
- class geojax.learning.MetricLearningModel(manifold, metric, embedding, regularization, diagnostics=<factory>)#
Bases:
objectAn equivariant embedding followed by a learned positive metric.
- Parameters:
manifold (Any)
metric (Any)
embedding (Callable[[Any], Any])
regularization (float)
diagnostics (Mapping[str, Any])
- transform(x)#
Return Euclidean coordinates whose norm realizes the learned metric.
- Parameters:
x (Any)
- Return type:
Any
- pairwise_distances(x, y=None)#
Return distances induced by the fitted embedding metric.
- Parameters:
x (Any)
y (Any | None)
- Return type:
Any