squarenet.squarenet
Classes
|
Grid straightening algorithm for D-dimensional point clouds. |
Exceptions
Raised when optimization does not converge. |
- exception squarenet.squarenet.ConvergenceWarning[source]
Bases:
UserWarningRaised when optimization does not converge.
- class squarenet.squarenet.SquareNet(gridshape, max_iter=1000, warnings_=True, verbose=2, backend='numpy', cuda=True, plot_config=None)[source]
Bases:
objectGrid straightening algorithm for D-dimensional point clouds.
SquareNet reorganizes an unordered set of points into a structured grid by iteratively sorting indices along each spatial axis. The goal is to enforce a topology where neighboring points in the original space are also close on the grid.
The grid is internally represented as an array of indices of shape
(N1, ..., ND), where D is the dimensionality of the embedding space.- Parameters:
gridshape (tuple of int) – Shape of the target grid. The product of these dimensions must exactly match the total number of input points (bijective gridification).
max_iter (int, default=100) – Maximum number of iterations allowed for the sorting procedure.
warnings_ – If True, emits a
ConvergenceWarningif the grid is not perfectly ordered at the end of the process.backend ({"numpy", "torch" or "jax"}, default="numpy") – input and output backend.
cuda (bool, default = True) – If True, SquareNet will search for a cuda GPU and a pytorch installation. If both are found, SquareNet will fit on GPU with pytorch, else SquareNet will fit on CPU with numpy. Note that the cuda flag affect ONLY the fit method, any other method is performed natively on the backend specified at init i.e. numpy torch or jax.
plot_config – python fonction that return a valid configuration for artist.sqplot many plot options are available.
- grid
The structured grid of indices, with shape
gridshape.- Type:
ndarray
- invert_grid
Inverse mapping from point indices to grid coordinates.
- Type:
ndarray, jax or numpy
- invgridflat
Flattened version of the inverse mapping for fast indexing.
- Type:
ndarray, jax or numpy
- learning_curve
History of the disorder metric across iterations.
- Type:
list of float
- 1. Run ``fit()`` with your preferred backend.
- 2. Call ``SquareNet.map()`` on any feature indexed like the points (N, \*C)
- To build a tensor version of it (*G, \*C) based on the points learned multi-indexes
- fit(points, method='fast')[source]
Fit the grid to a point cloud.
- Parameters:
points (ndarray of shape (N, D) or str) – Input point cloud or sampling method name.
method (fast, robust or ultimate.) – robust can be up to 5 time slower, but will probably give a better grid. ultimate can be up to 30 time slower but will give the best results among the three methods.
- Returns:
The method updates the internal state of the grid in-place.
- Return type:
None
Note
If a cuda GPU and a pytorch installation are available, fit is automatically performed on the GPU, which is much faster.
see
squarenet.core
- mapidx(index)[source]
Cloud index -> grid index.
- Returns:
directly usable for indexing.
- Return type:
tuple of arrays
- neighbormap(max_sample_size=20000000, max_window_size=31, criterion='rank', thresholdcut=1, projection=(0, 1), log2=False)[source]
Compute and display a neighborhood map from gridded points.
For each point X (or a subsample of sample_size if the dataset is too large), scans a square window of radius wr = ws//2 in the grid centered on X, and adds 1 to each cell if the point Y found in the cell matches the criterion.
- Parameters:
log2 (bool, default False) – If True, applies log2 scaling to counts.
max_sample_size (int, default 20 million) – sample size for the number of pairs (X, Y)
max_window_size (int, default=31) –
Note that
wr, the window radius, is defined aswr = window_size // 2. Search window size, i.e. the size of the window in which the neighbor map is computed, such that: gridindex(X) - gridindex(Y) <= wr` where <= is relative to the L∞ norm on grid indices.Warning
You will get no information, and thus no garantee at all, on what happens outside the search window. Conversely, using a very large search window dilutes the information, since the probability of sampling pairs for a given grid offset decreases.
Choosing the window size is therefore a tradeoff between macroscopic information and microscopic information
criterion (str, optional) – Method used to define neighborhoods (“rank” or “value”).
thresholdcut (int or float, optional Cutoff threshold for connections, interpreted according to the criterion.) –
If criterion is “rank”: thresholdcut = k means Y matches X
if Y is among the k nearest points to X. - If criterion is “value”: thresholdcut is a distance threshold, and Y matches X if the distance(X, Y) <= thresholdcut.
projection (tuple of int, optional) – Axes of the grid used for 2D projection.
- Returns:
None – Prints the resulting neighborhood matrix.
see
squarenet.neighbormap
- pad(points)[source]
Pad input points to the number of slots of the grid. Used if n_points < n_grid for injective gridification.
- Parameters:
points (ndarray of shape (n_points, d)) – The input data to be padded.
- Returns:
The populated array containing the input data at allowed positions and signed infinities elsewhere.
- Return type:
ndarray of shape (n_grid, d)
Note
Padding is performed such that points in the middle of the grid are true points and points in the corner are fictive points with +- inf coordinates, using the stencil.
- plot(style='checkerboard', animate=False, save=True, save_path='sqrnet/plot', **kwargs)[source]
Display the mapped grid as a static figure or animation.
Many rendering options (scales, colors, figsize, DS, frames …) can be passed as keyword arguments or configured once via
self.plot_config[key] = value. print(self.plot_config) to see all available parameters).- Parameters:
style (str) – Rendering style:
"checkerboard","mesh"or"scatter".animate (bool) – If True, produce a morphing animation (grid → identity). If False, render a static plot.
save_path (str or None) – None → display with
plt.show(). str → save to disk
- Returns:
fig (matplotlib.figure.Figure)
ani (matplotlib.animation.FuncAnimation or None)
see
squarenet.artist
- search_sorted(X, n_iter=2, side='left')[source]
X must be a SINGLE point, search_sorted will return a multiindex IX in the grid. IX is the grid index of PROBABLY one of the closest points to X in the first (side = ‘left’) or last (side = ‘right’) quadrant (relative to X) of the space. The search is a greedy coordinate descent: each dimension is refined by applying 1D searchsorted on row, then columns… Augmenting n_iter increases accuracy.
- squarenet.squarenet.cartesian_sort(gridmap, points, max_iter=100, method='fast', backend='numpy', loop=None, loopseq='decreasing', verbose=2)[source]
Supported:
method [fast, robust, ultimate] backend [numpy, torch]
Goal
Given an unordered point cloud:
points.shape = (N, D)
rearrange point indices into a structured cartesian grid:
grid.shape = (N1, N2, …, ND)
such that neighbouring cells of the grid contain spatially coherent points.
GRID INTERPRETATION
A grid index map bijectively to a point index:
point index: n <=> grid index: f(n) = (n1, n2, …, nD)
Each grid axis defines a local neighbour relation.
For axis d, let define:
next(n, d) = f-1(n1,…, nd+1,…,nD)
such that P(next(n,d)) is the neighbour obtained by incrementing the d-th grid coordinate.
The objective is therefore:
nearby euclidean points -> nearby grid cells
Note that the reciprocal property
nearby grid cells -> nearby euclidean points
is NOT guaranteed.
Indeed, datasets may contain cracks, holes, disconnected clusters, folds, or other topological singularities.
Gridification will naturally tend to close cracks, stitch nearby boundaries together and overlap the folds.
This algorithm is primarily designed for speed and scalability on large datasets. It is NOT an optimal transport solver.
Therefore, users should always inspect the resulting grid, especially near boundaries where geometric distortions are more likely to appear.
HEURISTIC ORDERING PRINCIPLE
For each spatial dimension d, define d spatial heuristics:
H_d(point) -> scalar
Heuristics must be orthogonal and monotonic Typical choice: cartesian heuristics
H_0 = x H_1 = y H_2 = z …
One could think on an improved version of the algorithm which would learn heuristics that best fit the dataset but cartesian are already pretty good
The desired property is:
H_d(P(n)) <= H_d(P(next(n, d)))
for every point index n and every axis d.
In words:
values of heuristic d should increase along grid axis d.
DISORDER METRIC
A local inversion occurs when:
H_d(P(n)) > H_d(P(next(n, d)))
The total disorder is simply the number of such violations.
Pseudo-definition:
- disorder =
sum over all axes d sum over all point pairs (current, next) of inversion count
FAST METHOD
The simplest strategy consists in repeatedly sorting each grid axis independently.
Pseudo-code
initialize gridmap
for iteration in range(max_iter):
for axis d in range(D):
sort grid indexes along axis d using heuristic H_d
compute disorder
- if disorder == 0:
stop
Properties
- Advantages:
extremely fast
fully vectorized
memory efficient
surprisingly effective
- Limitations:
for loop on the axes
-> early axes dominate later ones -> result in weak axes
only adjacent comparison is a weak accuracy criterion
-> disorder is blind on what happens on diagonals -> may converge toward local minima
ROBUST METHOD
To reduce axis domination and improve stability, sorting is performed only on random independent subgrids at each step, thus learning is much more progressiv
At every iteration:
each axis is randomly partitioned
only selected cartesian sub-blocks are sorted
different blocks are used for different axes
Result:
smoother convergence
better axis symmetry
fewer local minima
improved isotropy
Pseudo-code
for iteration:
generate random cartesian subgrids
for axis d:
for subgrid in selected_subgrids[d]:
sort only inside this subgrid
compute disorder
ULTIMATE METHOD
Even robust cartesian sorting still treats grid lines as parallel and therefore almost independent.
This can produce:
stratification
layered artifacts
weak coupling between parallel slices
To solve this, a second refinement stage is introduced.
The grid is embedded into a “hash table” containing shifted/sheared copies of the grid.
Repeated cyclic shears produce strong cross-line coupling.
This allows information to propagate between previously independent cartesian lines.
High-level idea
repeat:
shear grid into staggered hash table
sort along one axis
project back to grid
Effect
The repeated shearing progressively destroys artificial parallel structures and greatly reduces stratification. ============================================================ SUMMARY ============================================================
- FAST:
deterministic global line sorting (cartesian sort)
- ROBUST:
stochastic partial sorting preserving axis symmetry
- ULTIMATE:
sheared multi-pass relaxation reducing stratification artifacts
COMPLEXITY
A few hundreds iterations will be largely enough for convergence in most of the cases.
Let:
N = total number of points
Each iteration performs approximately:
O(N (log N +D)) operations
with highly vectorized NumPy operations.
- squarenet.squarenet.default_config()[source]
I. MAIN ARGUMENTS — passed directly to sqplot()
- stylerendering style. One of:
‘checkerboard’ — points coloured in a 2-colour tile pattern ‘mesh’ — grid lines drawn at 4 levels of density, ‘scatter’ — plain point cloud; depth-coloured (cmap) in
3-D, black in 2-D.
- animateFalse → single static PNG.
- True → GIF that morphs continuously from the input grid
to the identity grid and back.
save : whether to write the output to disk. save_path : destination path.
II. EXTRA ARGUMENTS — passed as cfg = {…}
All keys are optional. Only specify what you want to override; everything else is filled from the defaults shown here.
- Layout
figsize : (width, heigth) size of the base figure
- Projection
- projection3-tuple of axis indices, e.g. (0, 1, 2).
Used if ndim >= 3. this selects both which axes to keep AND in what order
- Scale & density
scale_factor : positive integer. bigger -> finer details pointsize : base size for a single point linewidth : base width for a mesh lines. mesh_long_edge : length ratio (relative to the median) above wich an edge will not be ploted in ‘mesh’ style
- Colours
colors_checkerboard : 2 colors for the ‘checkerboard’ style. cmap_scatter : colormap e.g. ‘plasma’, ‘coolwarm’ used in ‘scatter’ to encode the depth.
- Animation
frames : number of frames in the animation interval : inter-frame delay in milliseconds (for interactive sessions). fps : frames per second written (for the saved GIF).
- Display
- showif True, calls plt.show().
Will fail if session is not interactive
- squarenet.squarenet.make_stencil(gridshape, dtype=<class 'float'>)[source]
Create a grid stencil ordered by distance to the center of the grid. This is a helper for embedding an arbitrary set of n_target points into a smooth convex subset of a hyperrectangular lattice.
Invalid positions are initialized with signed infinities, while valid positions can later be filled with the input data according to fill_rank.
- Parameters:
gridshape (tuple of int) – The shape of the target grid.
dtype (data-type, optional) – The desired data-type for the stencil array (default is float).
- Returns:
stencil (ndarray of shape (n, d)) – A stencil prefilled with signed infinity at forbidden positions.
fill_rank (ndarray of shape (n,)) – Ranking of the lattice points by increasing distance to the grid center. The positions satisfying
fill_rank < n_targetare the locations where the firstn_targetpoints should be inserted.
- squarenet.squarenet.neighbormap(grided_points, mapidx, criterion='rank', thresholdcut=1, kernel=<function distkernel>, windowradius=10, projection=(0, 1), sample_size=1000000)[source]
- squarenet.squarenet.samplepoints(method='gaussian', size=(1000000, 2), plot_points=True)[source]
Generate point clouds using various sampling strategies.
- Parameters:
method (str) –
Sampling method. Supported values:
General methods (any dimension D): - “square” : uniform in [-1, 1]^D - “ball” : uniform in the unit ball - “ring” : Gaussian with radial offset - “gaussian” : standard normal distribution - “spiky” : L^alpha ball (alpha < 1) - “holy” : hypercube with spherical holes
Dataset-based methods (fixed dimensions): - “barbara”, “lena”, “france”, “germany” : 2D datasets - “everest” : 3D dataset (elevation-based sampling)
size (tuple of int) – (N, D) where N is the number of points and D the dimension. Note: D must match the dataset dimension for dataset-based methods.
plot_points (bool, optional) – If True, display the generated points.
- Returns:
Array of shape (N, D) containing sampled points.
- Return type:
numpy.ndarray
- squarenet.squarenet.sqplot(grid, verbose, style='checkerboard', animate=False, save=True, save_path='sqrnet/plot', cfg=None)[source]
Render a structured grid as either a static figure or a morphing animation.
The input grid is displayed using one of several rendering styles and may optionally be animated by continuously interpolating between the input grid and the corresponding identity grid.
- Parameters:
grid (ndarray) – Structured grid of shape
(..., D), whereDis either 2 or 3.verbose (bool) – If
True, prints progress information while rendering or exporting.style ({"checkerboard", "mesh", "scatter"}, default="checkerboard") –
Rendering style.
"checkerboard": alternating colours on adjacent grid cells."mesh": draw the grid connectivity as a wireframe."scatter": draw only the grid points. In 3-D, points are coloured
according to their depth after projection; in 2-D they are drawn in black.
animate (bool, default=False) – If
True, generate an animation interpolating between the input grid and the identity grid. Otherwise, produce a single static figure.save (bool, default=True) – If
True, save the generated figure or animation to disk.save_path (str, default="sqrnet/plot") – Output path used when
save=True.cfg (dict, optional) – Rendering configuration. Any omitted entries are filled with the default values returned by
default_config(). Seehelp(default_config)for the complete list of supported options.
- Returns:
fig (matplotlib.figure.Figure) – The generated figure.
ani (matplotlib.animation.FuncAnimation or None) – The animation object if
animate=True; otherwiseNone.
See also
default_configDefault rendering options and their documentation.
- squarenet.squarenet.to_backend(x, backend='numpy')[source]
Convert arrays between NumPy, Torch and JAX.
Philosophy
If backend already match, return x unchanged.
Otherwise always convert through NumPy.
- squarenet.squarenet.warn()
Issue a warning, or maybe ignore it or raise an exception.
- message
Text of the warning message.
- category
The Warning category subclass. Defaults to UserWarning.
- stacklevel
How far up the call stack to make this warning appear. A value of 2 for example attributes the warning to the caller of the code calling warn().
- source
If supplied, the destroyed object which emitted a ResourceWarning
- skip_file_prefixes
An optional tuple of module filename prefixes indicating frames to skip during stacklevel computations for stack frame attribution.