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: object

Canonical 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: object

Eager 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: Protocol

Callable protocol for user-registered representation adapters.

class geojax.learning.EquivariantEmbeddingProtocol(*args, **kwargs)#

Bases: Protocol

Geometry 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: ValueError

Raised 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=None denotes the axis immediately before each geometry’s event dimensions. Product representations and axes may be pytrees matching manifold.factors. Python sequences of complete points require representation='point_sequence' so their interpretation is explicit.

Parameters:
  • manifold (Any)

  • values (Any)

  • sample_axis (Any)

  • representation (Any)

  • check (str)

  • repair (bool)

Return type:

ManifoldData

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:

DataValidationReport

geojax.learning.register_manifold_data_adapter(geometry_type, representation, adapter, *, overwrite=False)#

Register an explicit representation converter for a geometry class.

Parameters:
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_size limits 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:

NeighborsResult

class geojax.learning.NeighborsResult(distances, indices)#

Bases: object

Distances 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:

FrechetMeanResult

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)), where rho_s(r)=r for r >= s and rho_s(r)=r^2/(2s)+s/2 otherwise. Thus smoothing controls 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:

FrechetMedianResult

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:

EnclosingBallResult

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:

KernelRegressionModel

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:

KernelCVResult

class geojax.learning.FrechetMeanResult(point, objective, gradient_norm, iterations, converged, reason, diagnostics=<factory>)#

Bases: object

A 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: object

A 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: object

A 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: object

A 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: object

Selected kernel model, bandwidth, and validation scores.

Parameters:

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:

NearestCentroidModel

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:

KNearestNeighborsModel

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:

TangentSpaceClassifierModel

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:

TangentSpaceClassifierModel

class geojax.learning.NearestCentroidModel(manifold, classes, centers, converged=True, reason='all class centroids converged', diagnostics=<factory>)#

Bases: object

Intrinsic class centroids and their internal mean-fit status.

predict_proba returns 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: object

A 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: object

A 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: object

A 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 M inside 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:

GeodesicRegressionModel

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:

LocalPolynomialRegressionModel

class geojax.learning.GeodesicRegressionModel(manifold, intercept, slope, predictor_mean, objective, iterations, converged, reason, diagnostics=<factory>)#

Bases: object

A 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: object

A 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:

HypothesisTestResult

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:

HypothesisTestResult

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:

HypothesisTestResult

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:

BootstrapResult

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:

HypothesisTestResult

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:

HypothesisTestResult

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:

HypothesisTestResult

class geojax.learning.HypothesisTestResult(statistic, pvalue, null_distribution, method, diagnostics=<factory>)#

Bases: object

Observed 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: object

A 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:

ClusteringResult

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:

CoresetResult

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:

ClusteringResult

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:

HierarchicalClusteringResult

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:

ClusteringResult

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:

ClusteringResult

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:

ClusteringResult

class geojax.learning.ClusteringResult(labels, centers, objective, iterations, converged, reason, diagnostics=<factory>)#

Bases: object

Cluster 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: object

Flat 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: object

Sampled 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_point influences the estimate only when initial_weight is 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:

FrechetMeanResult

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:

FrechetMeanResult

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:

ClusteringResult

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:

BarycentricCodingResult

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:

DictionaryLearningResult

class geojax.learning.BarycentricCodingResult(codes, reconstructions, objective, iterations, converged, reason, diagnostics=<factory>)#

Bases: object

Simplex 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: object

Learned 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:

RobustLocationResult

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:

RobustLocationResult

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:

MetricRanksResult

class geojax.learning.RobustLocationResult(point, objective, gradient_norm, iterations, converged, reason, diagnostics=<factory>)#

Bases: object

A 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: object

Midranks 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:

SemiSupervisedResult

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:

SemiSupervisedResult

class geojax.learning.SemiSupervisedResult(predictions, scores, objective, iterations, converged, reason, diagnostics=<factory>)#

Bases: object

Transductive 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:

EmbeddingResult

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:

EmbeddingResult

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=True to 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:

EmbeddingResult

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:

EmbeddingResult

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:

EmbeddingResult

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:

EmbeddingResult

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:

EmbeddingResult

class geojax.learning.EmbeddingResult(coordinates, objective, iterations, converged, reason, model=None, diagnostics=<factory>)#

Bases: object

Euclidean 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:

TransportResult

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: object

Transport 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 S and D, 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 embedding is 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:

MetricLearningModel

class geojax.learning.MetricLearningModel(manifold, metric, embedding, regularization, diagnostics=<factory>)#

Bases: object

An 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