Tip
Semantic interoperability among independently trained AI-native communication devices requires aligning heterogeneous latent spaces without retraining the underlying models. Existing alignment methods typically rely on linear transformations, which may be insufficient to capture nonlinear relations between independently learned semantic representations. In this paper, we propose Residual Kernel Alignment (RKA), a novel semantic alignment method that combines a geometry-preserving Stiefel transformation with a residual in a reproducing kernel Hilbert space (RKHS) to capture nonlinear latent space mismatch. An orthogonality constraint separates the two components and, under the Stiefel isometry condition, exactly decouples their estimation. The linear component is obtained through standard Procrustes alignment, while the nonlinear residual admits a closed-form constrained kernel ridge-regression solution. The proposed alignment strategy is learned from paired latent representations, referred to as semantic pilots. We therefore also address the design of the pilot set and develop a kernel-herding selection strategy to identify informative calibration samples. Numerical results show that RKA outperforms purely linear alignment and that optimized pilot selection provides substantial gains in the low-pilot regime.
This project uses uv for Python dependency management and just as the task runner.
Install the required tools:
Follow the installation instructions from their official documentation.
From the project root, run:
just setupThe setup recipe will:
- Create the
.venvvirtual environment (if it does not exist) - Install all project dependencies using
uv
After the command completes, the development environment will be ready to use. π
Three studies, one script and one config each, plus an average of the
last two over encoder pairs. All are driven by Hydra: no experiment
parameter lives in the justfile, so a run is fully described by
config/hydra/ plus the overrides on its command line. The defaults in
those configs are not the paper's setting β the commands below state
every override the figures were produced with.
The setting: SEMASIA CIFAR-10, regnety_016.pycls_in1k (888-d)
transmitting into vit_large_patch16_224.augreg_in21k_ft_in1k (1024-d),
a truncated-whitening chart shared by every method, kernel-herded pilots
and a linear probe on the receiver.
scripts/lambda_sweep.py fits RKA over its regularisation grid and
Procrustes once, per cell. Figures (ii) and (iii) do not refit either
method: they read these CSVs back, matched on their columns, so every
cell they need has to exist first. Figure (ii) needs rank 32 at every
pilot budget; figure (iii) needs every rank at 8192 pilots.
# k = 32 at every budget (figure ii), including N = 8192 (figure iii)
just lambda-sweep 'data.models=[regnety_016.pycls_in1k,vit_large_patch16_224.augreg_in21k_ft_in1k]' 'charts=[{preprocess:whiten}]' 'ranks=[32]' 'pilots.counts=[128,256,512,1024,2048,4096,8192]' decoder.kind=linear
# the other ranks at N = 8192 (figure iii)
just lambda-sweep 'data.models=[regnety_016.pycls_in1k,vit_large_patch16_224.augreg_in21k_ft_in1k]' 'charts=[{preprocess:whiten}]' 'ranks=[16,64,128,256,512]' 'pilots.counts=[8192]' decoder.kind=linearKernel herding is greedy and seeded from the config, so a cell's pilots do not depend on which other budgets share its run: splitting the grid across commands, or running it as one, writes the same CSVs.
Figure (i) β RKA against Procrustes over lambda, with Procrustes as
its lambda -> infinity limit:
figures/lambda_sweep/cifar10/whiten-k32/lambda_cifar10_regnety_016-to-vit_large_p16_whiten-k32_k32_n8192_lam1e-10to100x25_{accuracy,mrr}.{pdf,png}
figures/lambda_sweep/cifar10/whiten-k128/lambda_cifar10_regnety_016-to-vit_large_p16_whiten-k128_k128_n8192_lam1e-10to100x25_{accuracy,mrr}.{pdf,png}
scripts/pilot_sweep.py: RKA, Procrustes, Direct MLP and Residual MLP
against the pilot budget at 32 symbols, for kernel herding and
class-stratified random pilots (round_robin), mean Β± sd over five seeds.
RKA's lambda is read back per budget from step 1; the two MLPs are
trained on the same 32-d truncated-whitened pilots.
just pilot-sweep 'data.models=[regnety_016.pycls_in1k,vit_large_patch16_224.augreg_in21k_ft_in1k]' 'charts=[{preprocess:whiten}]' 'ranks=[32]' 'pilots.counts=[128,256,512,1024,2048,4096,8192]' 'pilots.strategies=[herding,round_robin]' decoder.kind=linearranks=[32] rather than symbols=32: the lambda lookup matches the chart
tag whiten-k32 that step 1 wrote.
figures/pilot_sweep/cifar10/whiten-k32/pilots_cifar10_regnety_016-to-vit_large_p16_whiten-k32_k32_N128to8192x7_herding-round_robin.{pdf,png}
scripts/dimension_sweep.py: the field against the number of
transmitted symbols at 8192 pilots. RKA and Procrustes are read back from
step 1; CCA, SVCCA and Proto-PFE are fitted on the same herded pilots,
with their rate set by the canonical rank or the anchor count rather than
by a truncation. Before drawing, the
run refits Procrustes and stops if it disagrees with the CSVs.
just dimension-sweep 'data.models=[regnety_016.pycls_in1k,vit_large_patch16_224.augreg_in21k_ft_in1k]' 'charts=[{preprocess:whiten}]' 'ranks=[16,32,64,128,256,512]' pilots.n_pilots=8192 decoder.kind=linearfigures/dimension_sweep/cifar10/whiten/dims_cifar10_regnety_016-to-vit_large_p16_whiten_n8192_herding_k16to512x6_dec-linear.{pdf,png}
scripts/pair_average.py fits nothing: it pools steps 2 and 3 for
every pair in config/hydra/pair_average.yaml and redraws both figures
with the band over pairs instead of seeds (seeds are averaged within a
pair first). The five pairs are CNN transmitters into differently
pre-trained ViT receivers, no encoder repeated:
| transmitter | receiver |
|---|---|
regnety_016.pycls_in1k |
vit_large_patch16_224.augreg_in21k_ft_in1k |
repvgg_b0.rvgg_in1k |
aimv2_large_patch14_224.apple_pt |
mobilenetv3_large_100.ra_in1k |
vit_base_patch16_224.augreg_in21k |
efficientvit_b0.r224_in1k |
beit_base_patch16_224.in22k_ft_in22k |
ghostnet_100.in1k |
eva02_base_patch14_224.mim_in22k |
Run steps 1β3 for each pair with the receiver pinned and one pilot-sweep seed, e.g. for the second pair:
just lambda-sweep 'data.models=[repvgg_b0.rvgg_in1k,aimv2_large_patch14_224.apple_pt]' receiver=aimv2_large_patch14_224.apple_pt 'charts=[{preprocess:whiten}]' 'ranks=[32]' 'pilots.counts=[128,256,512,1024,2048,4096,8192]' decoder.kind=linear
just lambda-sweep 'data.models=[repvgg_b0.rvgg_in1k,aimv2_large_patch14_224.apple_pt]' receiver=aimv2_large_patch14_224.apple_pt 'charts=[{preprocess:whiten}]' 'ranks=[16,64,128,256,512]' 'pilots.counts=[8192]' decoder.kind=linear
just pilot-sweep 'data.models=[repvgg_b0.rvgg_in1k,aimv2_large_patch14_224.apple_pt]' receiver=aimv2_large_patch14_224.apple_pt 'charts=[{preprocess:whiten}]' 'ranks=[32]' 'pilots.counts=[128,256,512,1024,2048,4096,8192]' 'pilots.strategies=[herding,round_robin]' 'seeds=[0]' decoder.kind=linear
just dimension-sweep 'data.models=[repvgg_b0.rvgg_in1k,aimv2_large_patch14_224.apple_pt]' receiver=aimv2_large_patch14_224.apple_pt 'charts=[{preprocess:whiten}]' 'ranks=[16,32,64,128,256,512]' pilots.n_pilots=8192 decoder.kind=linearthen pool them:
just pair-averageYou do not have to type the other pairs out: for any pair not yet on
disk, just pair-average stops and prints its four commands, already
filled in. A pair takes roughly an hour, most of it the pilot sweep's MLP
baselines.
figures/pair_average/cifar10/whiten/pairavg-pilots_cifar10_5pairs_whiten-k32_N128to8192x7_herding-round_robin.{pdf,png}
figures/pair_average/cifar10/whiten/pairavg-dims_cifar10_5pairs_whiten_n8192_herding_k16to512x6.{pdf,png}
Figures and the data behind them live in parallel trees, and every filename carries the whole configuration that produced it β a figure loses its path the moment it is copied into a paper.
figures/<study>/<dataset>/<chart>/<stem>[_<metric>].{pdf,png}
results/<study>/<dataset>/<stem>.csv
Every script writes each unit of work β a (chart, budget) cell, a
seed, a fitted method at one rank β as it finishes and skips what is
already on disk, so an interrupted run resumes rather than restarts, and
just retries automatically. Append to any command above:
resume=false # refit everything
plot_only=true # redraw from the CSVs, fit nothingIf you find this code useful for your research, please consider citing the following paper: