NN API¶
The pyrox_nn subpackage ships uncertainty-aware neural network layers in four families:
- Geographic / spherical encoders (re-exported from
geonnax) — degree/radian, lon/lat, cyclic, spherical-harmonic, and Slepian preprocessing for geophysical inputs. - Dense / Bayesian-linear layers (
pyrox_nn._dense) — reparameterization, Flipout, hierarchical, NCP, DVI, rank-1 ensemble, and variational-dropout variants ofWx + b. - Spectral / GP-flavoured layers — random-feature kernel maps, SNGP and VSSGP heads, deep random-feature expansions.
- Ensembles & output heads — BatchEnsemble layers, heteroscedastic Monte-Carlo output heads.
- Bayesian Neural Field stack (
pyrox_nn._bnf) — five layers that together implement the BNF architecture (Saad et al., Nat. Comms. 2024). - Pure-JAX feature helpers (re-exported from
geonnax.basis) — pandas-free building blocks the BNF layers wrap.
See also: Geo encoders for the longitude/latitude and spherical-harmonic API surface.
Dense / Bayesian-linear layers¶
DenseReparameterization
¶
Bases: PyroxModule
Bayesian dense layer via the reparameterization trick.
Samples weight and bias from learned Gaussian posteriors at every forward pass. Registers NumPyro sample sites so the KL between the variational posterior and the prior is tracked by the ELBO.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
out_features |
int
|
Output dimension. |
bias |
bool
|
Whether to include a bias term. |
prior_scale |
float
|
Scale of the isotropic Gaussian prior on weights and bias. The prior mean is zero. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
DenseFlipout
¶
Bases: PyroxModule
Bayesian dense layer with Flipout sign-flip structure.
Samples weight from the prior and applies per-example Rademacher sign flips to the weight perturbation (Wen et al., 2018). Under a NumPyro guide that learns the posterior mean, the sign flips decorrelate gradient estimates across minibatch examples.
In model mode (no guide) this is equivalent to
DenseReparameterization — the Flipout variance reduction
activates when a guide provides a posterior centered at a learned
mean.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
out_features |
int
|
Output dimension. |
bias |
bool
|
Whether to include a bias term. |
prior_scale |
float
|
Scale of the isotropic Gaussian prior. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
DenseVariational
¶
Bases: PyroxModule
Dense layer with a user-supplied prior factory.
Provides flexibility over the weight prior by accepting a callable
that builds the prior distribution given the layer shape. The
model samples from the prior; the posterior is handled by a NumPyro
guide (e.g., AutoNormal).
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
out_features |
int
|
Output dimension. |
make_prior |
Callable[..., Any]
|
Callable |
bias |
bool
|
Whether to include a bias term. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
DenseDVI
¶
Bases: PyroxModule
Deterministic Variational Inference dense layer (Wu et al., 2018).
Propagates a Gaussian distribution through the linear layer analytically — there is no Monte Carlo sampling. The input is a diagonal-covariance Gaussian \((\mu_x, \sigma_x^2)\), the output is the (still-diagonal) Gaussian \((\mu_y, \sigma_y^2)\) induced by an independent-Gaussian variational posterior \(q(W) = \mathcal{N}(M, S)\) over the weights and a separate diagonal posterior on the bias.
With weight posterior mean \(M\) (shape \(D_\mathrm{in}\times D_\mathrm{out}\)) and per-element posterior variance \(S\) (same shape):
plus the bias mean / variance if enabled. Compared to MC estimators, DVI gives zero-variance gradients of the ELBO at the cost of propagating second-order statistics layer by layer (so it only really pays off when all dense layers in a block are DVI; a single DVI layer in a sampling stack just adds bookkeeping).
The KL between the diagonal-Gaussian variational posterior and a
fixed isotropic Gaussian prior \(p(W) = \mathcal{N}(0, \pi^2)\)
is closed-form and is registered with numpyro.factor so
SVI's Trace_ELBO picks it up:
Plate semantics
Same as the rest of the pyrox Bayesian dense family — call
this layer outside numpyro.plate("data", ..., subsample_size=...).
The KL is a weight-prior term: it sums over the weight and
bias matrices, not over the batch, so it's a single scalar
per layer. numpyro.factor is still a sample-type site,
though, and putting it inside a subsampled plate would broadcast
the scalar to the plate dim and apply scale = N/B — the
same over-counting trap that affects every per-layer
numpyro.factor. Keep this layer at the top of the model
(or outside any data plate) and only plate the observation
likelihood:
def model(x, y=None):
mean, var = dvi(x_mean, x_var) # KL emitted here
with numpyro.plate("data", x.shape[0]):
numpyro.sample("obs",
dist.Normal(mean, jnp.sqrt(var)), obs=y)
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D_\mathrm{in}\). |
out_features |
int
|
Output dimension \(D_\mathrm{out}\). |
bias |
bool
|
Whether to include a diagonal-Gaussian bias. |
prior_scale |
float
|
Std \(\pi\) of the isotropic Gaussian prior. |
init_log_var |
float
|
Initial value for the log posterior variance (a small negative number keeps initial draws tight). |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> dvi = DenseDVI(in_features=3, out_features=2, pyrox_name="dvi")
>>> mean = jnp.ones((4, 3))
>>> var = 0.1 * jnp.ones((4, 3))
>>> with handlers.seed(rng_seed=0):
... out_mean, out_var = dvi(mean, var)
>>> out_mean.shape, out_var.shape
((4, 2), (4, 2))
References
Wu, A., Nowozin, S., Meeds, E., Turner, R. E., Hernández-Lobato, J. M., & Gaunt, A. L. (2018). Deterministic Variational Inference for Robust Bayesian Neural Networks. ICLR.
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 | |
DenseHierarchical
¶
Bases: PyroxModule
Hierarchical Bayesian dense layer with multiplicative shrinkage.
Decomposes the effective weight matrix into a deterministic base \(\theta \in \mathbb{R}^{D_\mathrm{in} \times D_\mathrm{out}}\) multiplied row-wise by a per-input-unit local scale \(z^{(\mathrm{loc})} \in \mathbb{R}^{D_\mathrm{in}}\) and an overall global scale \(z^{(\mathrm{glob})} \in \mathbb{R}\),
with isotropic Gaussian priors centred at one,
The local scale prunes individual input units (a column of
\(\theta\) whose z_loc posterior concentrates near zero is
effectively switched off) while the global scale modulates the
overall layer activation — the same hierarchical-shrinkage
structure used by horseshoe-style BNNs (Louizos et al., 2017).
Both scales are pyrox_sample sites so any standard NumPyro
guide (AutoNormal, etc.) drives the variational posterior; the
deterministic base \(\theta\) and bias are pyrox_param.
Plate semantics
Same as the rest of pyrox_nn's Bayesian dense layers — call
outside numpyro.plate("data", ..., subsample_size=...) and
only plate the observation likelihood, otherwise the
per-layer prior log-probabilities of z_loc and z_glob
get scaled by the subsample ratio.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D_\mathrm{in}\). |
out_features |
int
|
Output dimension \(D_\mathrm{out}\). |
bias |
bool
|
Whether to include a deterministic bias term. |
prior_local_scale |
float
|
Std \(\sigma_\mathrm{loc}\) of the local scale prior. |
prior_global_scale |
float
|
Std \(\sigma_\mathrm{glob}\) of the global scale prior. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> layer = DenseHierarchical(
... in_features=4, out_features=2, pyrox_name="hier"
... )
>>> x = jnp.ones((3, 4))
>>> with handlers.seed(rng_seed=0):
... y = layer(x)
>>> y.shape
(3, 2)
References
Louizos, C., Ullrich, K., & Welling, M. (2017). Bayesian Compression for Deep Learning. NeurIPS.
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 | |
DenseVariationalDropout
¶
Bases: PyroxModule
Sparse variational dropout dense layer.
Implements variational dropout (Kingma et al., 2015) extended by Molchanov et al. (2017) to a log-uniform prior that enables automatic sparsification via per-weight learnable dropout rates. The variational posterior on weights is
Forward passes use the local reparameterization trick — the pre-activation distribution is closed-form and the noise is sampled once per output unit per batch element rather than once per weight:
The KL between the posterior and the log-uniform prior is
approximated analytically (Molchanov et al., 2017) and added to the
NumPyro trace via numpyro.factor. SVI then optimizes
Weights with log_alpha > threshold (default 3.0, dropout rate
~0.95) are effectively pruned; inspect the trained pattern via
sparsity.
Plate semantics
The KL contribution is registered via numpyro.factor,
which is itself a sample-type site and therefore subject to
numpyro.plate scaling. To keep the per-layer KL counted
once (not scaled by the data-plate's subsample ratio), call
the layer outside any plate("data", ..., subsample_size=...)
block — the standard pyrox / NumPyro convention for global
Bayesian parameters. Plate only the observation likelihood.
Correct (forward outside the data plate):
def model(x, y=None):
layer = DenseVariationalDropout(in_features=D, out_features=1)
f = layer(x).squeeze(-1) # KL emitted here
with numpyro.plate("data", x.shape[0]):
numpyro.sample("obs", dist.Normal(f, 0.5), obs=y)
Incorrect (forward inside a subsampled data plate scales KL by
N / subsample_size):
def model(x, y=None):
with numpyro.plate("data", N, subsample_size=B) as idx:
f = layer(x[idx]).squeeze(-1) # ⚠ scales KL
numpyro.sample("obs", ...)
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
out_features |
int
|
Output dimension. |
bias |
bool
|
Whether to include a bias term. |
log_alpha_init |
float
|
Initial value for |
threshold |
float
|
|
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> layer = DenseVariationalDropout(
... in_features=4, out_features=2, pyrox_name="vd"
... )
>>> x = jnp.ones((3, 4))
>>> with handlers.seed(rng_seed=0):
... y = layer(x)
>>> y.shape
(3, 2)
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 | |
sparsity(log_alpha: Float[Array, 'D_in D_out']) -> Float[Array, '']
¶
Fraction of weights with log_alpha > threshold.
Pass the trained log_alpha parameter, typically retrieved
from the SVI param store under f"{pyrox_name}.log_alpha".
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
DenseNCP
¶
Bases: PyroxModule
Noise Contrastive Prior dense layer (Hafner et al., 2019).
Decomposes a dense layer into a prior-regularized backbone plus a scaled stochastic perturbation:
where all weights are pyrox_sample sites with Gaussian priors
and \(\sigma\) has a LogNormal prior. The backbone carries
the bulk of the signal; the perturbation branch adds calibrated
uncertainty that can be trained via a noise contrastive objective.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
out_features |
int
|
Output dimension. |
init_scale |
float
|
Initial value for the perturbation scale \(\sigma\). |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
NCPContinuousPerturb
¶
Bases: Module
Input perturbation for the Noise Contrastive Prior pattern.
Adds Gaussian noise scaled by a fixed positive scale to the input:
Place before a deterministic network to inject input uncertainty;
pair with a Bayesian DenseNCP head for the full NCP
architecture (Hafner et al., 2019).
Stochasticity comes from the explicit PRNG key argument.
Attributes:
| Name | Type | Description |
|---|---|---|
scale |
float | Float[Array, '']
|
Perturbation scale \(\sigma\). |
Examples:
>>> import jax.numpy as jnp
>>> import jax.random as jr
>>> perturb = NCPContinuousPerturb(scale=0.5)
>>> x = jnp.zeros(3)
>>> out = perturb(x, key=jr.PRNGKey(0)) # x̃ = x + σ·ε, ε ~ N(0, I)
>>> out.shape
(3,)
Source code in .venv/lib/python3.12/site-packages/geonnax/ncp.py
NCPNormalOutput
¶
Bases: PyroxModule
Output-side Noise Contrastive Prior layer (Hafner et al., 2018).
Completes the NCP pattern in pyrox_nn: pair with
NCPContinuousPerturb at the input and a heteroscedastic
network (e.g. an MLP terminating in a mean head and a positive-std
head — a softplus or exp of a learned log-scale) so the
network produces predictions for both the clean batch and the
input-perturbed noisy batch. Given the noisy batch's predictive
distribution \(\mathcal{N}(\hat{y}_n, \hat{\sigma}_n^2)\),
this layer adds the analytic NCP regulariser
to the model log density via numpyro.factor. Pulling the
noisy-input predictive distribution toward the fixed prior away
from the training distribution gives the network calibrated
out-of-distribution uncertainty, which is the central claim of NCP.
The closed-form Gaussian KL used here is
Plate semantics
Unlike pyrox's weight-prior KL terms, the NCP KL is
data-dependent — every input row contributes its own
\(\mathrm{KL}_n\) term. Internally the layer emits the
numpyro.factor site as a per-example vector
(shape (*batch,)) rather than a pre-summed scalar; that
lets NumPyro's plate machinery sum over the batch axis and
apply the subsample scaling automatically.
The canonical training pattern is to emit the layer inside
numpyro.plate("data", N, subsample_size=B):
def model(x_clean, y_clean, x_noisy):
clean_mean, _clean_std = network(x_clean)
noisy_mean, noisy_std = network(x_noisy)
ncp_out = NCPNormalOutput(prior_std=1.0)
with numpyro.plate("data", N, subsample_size=B):
ncp_out(noisy_mean, noisy_std) # scaled to N
numpyro.sample("obs",
dist.Normal(clean_mean, ...), obs=y_clean)
Inside the plate, NumPyro sums the per-example log-densities
over the batch dim and multiplies by scale = N / B,
producing the standard unbiased estimate of the full-dataset
NCP KL Σ_{n=1}^N KL_n. Outside any plate the layer's
contribution is just Σ_{n in batch} KL_n (i.e. the raw
batch sum), which is the correct full-dataset value when
you train on the whole dataset at once.
Attributes:
| Name | Type | Description |
|---|---|---|
prior_mean |
float
|
Prior predictive mean \(\mu_\mathrm{prior}\). |
prior_std |
float
|
Prior predictive std \(\sigma_\mathrm{prior}\) (must be positive). |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> ncp = NCPNormalOutput(
... prior_mean=0.0, prior_std=1.0, pyrox_name="ncp_out"
... )
>>> noisy_mean = jnp.zeros((4, 1))
>>> noisy_std = 0.5 * jnp.ones((4, 1))
>>> with handlers.seed(rng_seed=0):
... kl = ncp(noisy_mean, noisy_std)
>>> kl.shape
()
References
Hafner, D., Tran, D., Lillicrap, T., Irpan, A., & Davidson, J. (2018). Noise Contrastive Priors for Functional Uncertainty. UAI.
Source code in packages/pyrox-nn/src/pyrox_nn/_dense.py
254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 | |
RBFFourierFeatures
¶
Bases: PyroxModule
SSGP-style RFF layer with RBF spectral density.
Both the spectral frequencies \(W\) and the lengthscale
\(\ell\) are pyrox_sample sites — \(W\) has a
standard normal prior (the RBF spectral density) and \(\ell\)
has a LogNormal prior. Under SVI, the guide learns a posterior
over both; under a seed handler, they are drawn from the prior.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
n_features |
int
|
Number of frequency pairs (output dim
|
init_lengthscale |
float
|
Prior location for the lengthscale. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
RBFCosineFeatures
¶
Bases: PyroxModule
Cosine-bias variant of random Fourier features for the RBF kernel.
Uses the single-cosine feature map with a bias term:
where \(W \sim \mathcal{N}(0, I)\) and
\(b \sim \mathrm{Uniform}(0, 2\pi)\). This variant produces
n_features-dimensional output (half the dimension of the
[cos, sin] variant in RBFFourierFeatures) and is
commonly used in Random Kitchen Sinks implementations.
All parameters (\(W\), \(b\), \(\ell\)) are
pyrox_sample sites.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
n_features |
int
|
Number of random features (= output dimension). |
init_lengthscale |
float
|
Prior location for the lengthscale. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
MaternFourierFeatures
¶
Bases: PyroxModule
SSGP-style RFF layer with Matern spectral density.
Spectral frequencies \(W\) have a StudentT(df=2\nu) prior
(the Matern spectral density). The smoothness \(\nu\) controls
the regularity: nu=0.5 (Laplace), nu=1.5 (Matern-3/2),
nu=2.5 (Matern-5/2).
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
n_features |
int
|
Number of frequency pairs. |
nu |
float
|
Smoothness parameter \(\nu\). |
init_lengthscale |
float
|
Prior location for the lengthscale. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
MaternCosineFeatures
¶
Bases: PyroxModule
Cosine-bias variant of random Fourier features for the Matern kernel.
Single-cosine analogue of MaternFourierFeatures:
where \(W \sim \mathrm{StudentT}(2\nu)\) (the Matern spectral
density) and \(b \sim \mathrm{Uniform}(0, 2\pi)\). Output dim is
n_features (vs 2 * n_features for the [cos, sin]
variant). Approximates the same kernel as
MaternFourierFeatures in expectation but with higher
variance per draw — see Sutherland & Schneider (2015).
All parameters (\(W\), \(b\), \(\ell\)) are
pyrox_sample sites.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
n_features |
int
|
Number of random features (= output dimension). |
nu |
float
|
Smoothness parameter \(\nu\). |
init_lengthscale |
float
|
Prior location for the lengthscale. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
LaplaceFourierFeatures
¶
Bases: PyroxModule
SSGP-style RFF layer with Laplace (Matern-1/2) spectral density.
Spectral frequencies \(W\) have a Cauchy prior (Student-t
with df = 1).
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
n_features |
int
|
Number of frequency pairs. |
init_lengthscale |
float
|
Prior location for the lengthscale. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
LaplaceCosineFeatures
¶
Bases: PyroxModule
Cosine-bias variant of random Fourier features for the Laplace kernel.
Single-cosine analogue of LaplaceFourierFeatures (the
Matern-1/2 kernel):
where \(W \sim \mathrm{Cauchy}(0, 1)\) (Student-t with
df = 1) and \(b \sim \mathrm{Uniform}(0, 2\pi)\). Output
dim is n_features.
All parameters (\(W\), \(b\), \(\ell\)) are
pyrox_sample sites.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
n_features |
int
|
Number of random features (= output dimension). |
init_lengthscale |
float
|
Prior location for the lengthscale. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
ArcCosineFourierFeatures
¶
Bases: PyroxModule
Random features for the arc-cosine kernel (Cho & Saul, 2009).
The arc-cosine kernel of order \(p\) corresponds to an infinite-width single-layer ReLU network. The random feature map is:
where \(W \sim \mathcal{N}(0, I)\).
order=0 gives the Heaviside (step) feature; order=1 gives
the ReLU feature (the most common); order=2 gives the squared
ReLU feature.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension. |
n_features |
int
|
Number of random features (= output dimension). |
order |
int
|
Kernel order (0, 1, or 2). |
init_lengthscale |
float
|
Prior location for the lengthscale. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
RandomKitchenSinks
¶
Bases: PyroxModule
Random Kitchen Sinks: RFF + a learned linear head.
Composes any RFF layer (RBFFourierFeatures,
MaternFourierFeatures, LaplaceFourierFeatures)
with a trainable linear projection:
The linear head (beta, bias) is registered via
pyrox_sample with Normal priors.
Attributes:
| Name | Type | Description |
|---|---|---|
rff |
RBFFourierFeatures | MaternFourierFeatures | LaplaceFourierFeatures
|
The underlying RFF feature layer. |
init_beta |
Float[Array, 'D_rff D_out']
|
Initial linear weights. |
init_bias |
Float[Array, ' D_out']
|
Initial bias vector. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
init(rff: RBFFourierFeatures | MaternFourierFeatures | LaplaceFourierFeatures, out_features: int) -> RandomKitchenSinks
classmethod
¶
Construct from a pre-built RFF layer with zero-initialized head.
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
Wave-4 spectral layers (#41)¶
VariationalFourierFeatures
¶
Bases: PyroxModule
VSSGP — RFF with a learnable variational posterior over frequencies.
Standard RFF (e.g. RBFFourierFeatures) treats the spectral
frequencies \(W\) as a frozen prior draw; VSSGP (Gal & Turner,
2015) treats \(W\) as a latent with a learnable mean-field
posterior, recovering spectral uncertainty on top of the
feature-space uncertainty.
Prior: \(p(W) = \mathcal{N}(0, I)\) (RBF spectral density in
lengthscale-1 units). The lengthscale is itself a sampled site
(LogNormal(log init_lengthscale, 1)) so that frequencies are
rescaled to the physical kernel.
Under SVI, attach an AutoNormal to
learn the posterior on W; under prior-only seeds, behaves
identically to RBFFourierFeatures.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D\). |
n_features |
int
|
Number of frequency pairs (output dim |
init_lengthscale |
float
|
Prior location for the kernel lengthscale. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
OrthogonalRandomFeatures
¶
Bases: Module
Orthogonal Random Features (Yu et al., 2016) — variance-reduced RFF.
Frequencies are drawn from blocks of Haar-orthogonal matrices scaled by
independent chi-distributed magnitudes, giving the same RBF kernel
approximation as plain RBFFourierFeatures in expectation but
with provably lower variance for finite n_features.
Frozen at construction time — no priors, no SVI on W. The frequency
matrix is built once from a key and stored as a static array.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D\). |
n_features |
int
|
Number of feature pairs. Must satisfy
|
lengthscale |
Float[Array, '']
|
Fixed kernel lengthscale (no prior; pass a value). |
W |
Float[Array, 'D_in D_orf']
|
Pre-built frequency matrix of shape |
The feature map is the shared RFF map
\(\phi(x) = \sqrt{1/D}\,[\cos(W^\top x/\ell),\,\sin(W^\top x/\ell)]\),
so calling the module on a (in_features,) vector yields a
(2 * n_features,) feature vector.
Examples:
>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.randfeat import OrthogonalRandomFeatures
>>> orf = OrthogonalRandomFeatures.init(4, 8, key=jr.PRNGKey(0))
>>> orf(jnp.ones(4)).shape
(16,)
Source code in .venv/lib/python3.12/site-packages/geonnax/randfeat.py
107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 | |
init(in_features: int, n_features: int, *, key: jax.Array, lengthscale: float = 1.0) -> OrthogonalRandomFeatures
classmethod
¶
Build the frozen ORF frequency matrix and wrap it in the module.
n_features must be divisible by in_features so the
Haar-orthogonal blocks tile cleanly.
Examples:
>>> import jax.random as jr
>>> from geonnax.randfeat import OrthogonalRandomFeatures
>>> orf = OrthogonalRandomFeatures.init(4, 8, key=jr.PRNGKey(0))
>>> orf.W.shape
(4, 8)
Source code in .venv/lib/python3.12/site-packages/geonnax/randfeat.py
HSGPFeatures
¶
Bases: PyroxModule
Hilbert-Space Gaussian Process feature layer (Riutort-Mayol et al., 2023).
A deterministic Laplacian-eigenfunction basis on the bounded box \([-L, L]^D\) plus learnable per-basis amplitudes with a kernel-spectral-density prior:
This is the NN-side dual of pyrox_gp.FourierInducingFeatures
— same basis, different prior wiring. As M and L grow, the
induced GP converges to the kernel passed in.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D\). |
num_basis_per_dim |
tuple[int, ...]
|
Per-axis number of 1D eigenfunctions; total
basis count is |
L |
tuple[float, ...]
|
Per-axis box half-width. |
kernel |
Kernel
|
A stationary kernel from |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Source code in packages/pyrox-nn/src/pyrox_nn/_features.py
647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 | |
SIREN — Sinusoidal Representation Networks¶
SIREN (Sitzmann, Martel, Bergman, Lindell, Wetzstein — NeurIPS 2020) replaces
ReLU/GELU with sin and prescribes a three-regime initialisation scheme that
keeps pre-activation variance stable across depth.
Three-regime weight initialisation (Theorem 1)¶
| Layer | W init |
Activation |
|---|---|---|
"first" |
U(-1/d_in, 1/d_in) |
sin(ω₀ · (W x + b)) |
"hidden" |
U(-√(c/d_in)/ω, √(c/d_in)/ω) |
sin(ω · (W x + b)) |
"last" |
U(-√(c/d_in), √(c/d_in)) |
none (linear) — W x + b |
Bias b is initialised U(-1/√d_in, 1/√d_in) for every regime.
Typical choice: ω₀ = ω = 30 for image / high-frequency INR tasks.
Usage¶
import jax.random as jr, jax.numpy as jnp
from pyrox_nn import SirenDense, SIREN, BayesianSIREN
# Single layer
layer = SirenDense.init(3, 64, key=jr.PRNGKey(0), layer_type="first")
y = layer(jnp.ones((5, 3))) # (5, 64)
# Multi-layer network (depth=5 → first + 3 hidden + last)
net = SIREN.init(2, 64, 1, depth=5, key=jr.PRNGKey(0))
y = net(jnp.zeros((100, 2))) # (100, 1)
# Bayesian variant (no key needed — weights come from the prior)
from numpyro import handlers
bnet = BayesianSIREN.init(2, 32, 1, depth=3)
with handlers.seed(rng_seed=0):
y = bnet(jnp.zeros((10, 2))) # (10, 1)
Alternative INR backbone
SIREN and GaborNet / FourierNet (MFN, #87) are complementary INR
backbones: SIREN composes nonlinearities deeply, while MFN uses a product
of Gabor filters. Choose based on the signal's smoothness profile.
SirenDense
¶
Bases: Module
Sine-activated dense layer: y = sin(ω · (W x + b)) or y = W x + b.
Single primitive of a SIREN network with three init regimes (Sitzmann et al. 2020, Theorem 1):
+----------+-------------------------------------------+-------------+
| Regime | W init | Activation |
+==========+===========================================+=============+
| first | U(-1/d_in, 1/d_in) | sin(ω··)|
+----------+-------------------------------------------+-------------+
| hidden | U(-√(c/d_in)/ω, √(c/d_in)/ω) | sin(ω··)|
+----------+-------------------------------------------+-------------+
| last | U(-√(c/d_in), √(c/d_in)) | none |
+----------+-------------------------------------------+-------------+
Bias b is initialised U(-1/√d_in, 1/√d_in) for every regime.
Attributes:
| Name | Type | Description |
|---|---|---|
W |
Float[Array, 'in_features out_features']
|
Weight matrix of shape |
b |
Float[Array, ' out_features']
|
Bias vector of shape |
omega |
float
|
Frequency multiplier applied inside the sine. |
in_features |
int
|
Input dimension. |
out_features |
int
|
Output dimension. |
layer_type |
SirenLayerType
|
One of |
c |
float
|
Constant from Theorem 1 (default 6.0). |
Source code in .venv/lib/python3.12/site-packages/geonnax/siren.py
121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 | |
init(in_features: int, out_features: int, *, key: Array, omega: float = 30.0, layer_type: SirenLayerType = 'hidden', c: float = 6.0) -> SirenDense
classmethod
¶
Construct a SirenDense with Sitzmann-regime weight initialisation.
Examples:
>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.siren import SirenDense
>>> layer = SirenDense.init(
... 3, 8, key=jr.PRNGKey(0), layer_type="first"
... )
>>> layer(jnp.ones(3)).shape # (3,) -> (8,)
(8,)
Source code in .venv/lib/python3.12/site-packages/geonnax/siren.py
SIREN
¶
Bases: Module
Multi-layer sinusoidal representation network (Sitzmann et al., NeurIPS 2020).
Topology:
Each layer uses the corresponding Sitzmann Theorem 1 init regime
(SirenDense): "first" for layer 0, "hidden" for
intermediate layers, and "last" for the readout.
depth counts all layers including the readout; depth=2 gives
one first-layer + one last-layer (no hidden layers); depth=5 gives
first + 3 hidden + last. Must be ≥ 2.
Source code in .venv/lib/python3.12/site-packages/geonnax/siren.py
227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 | |
init(in_features: int, hidden_features: int, out_features: int, *, depth: int, key: Array, first_omega: float = 30.0, hidden_omega: float = 30.0, c: float = 6.0) -> SIREN
classmethod
¶
Construct a SIREN with the correct per-layer init regimes.
Examples:
>>> import jax.numpy as jnp, jax.random as jr
>>> from geonnax.siren import SIREN
>>> net = SIREN.init(2, 16, 1, depth=4, key=jr.PRNGKey(0))
>>> net(jnp.zeros(2)).shape # (2,) -> (1,)
(1,)
>>> len(net.layers) # first + 2 hidden + last
4
Source code in .venv/lib/python3.12/site-packages/geonnax/siren.py
BayesianSIREN
¶
Bases: PyroxModule
SIREN with regime-scaled Normal priors on all layer weights.
Replaces the deterministic weight matrices of SIREN with NumPyro
sample sites. For layer \(i\) with Sitzmann Theorem 1 half-width
\(a_i\) (the uniform bound used by SirenDense):
where \(\sigma_0\) is prior_std and \(d_i\) is the input
dimension of layer \(i\). The \(a_i / \sqrt{3}\) factor makes
\(\operatorname{Var}(W_i)\) equal to the variance of Sitzmann's
\(\mathcal{U}(-a_i, a_i)\) init exactly, so the Bayesian prior
preserves the activation variance prescribed by Theorem 1 — avoiding
the saturated-sine pathology that a flat \(\mathcal{N}(0, 1)\)
prior would cause.
Registered sites: {scope}.layer_0.W, {scope}.layer_0.b, …,
{scope}.layer_{depth-1}.W, {scope}.layer_{depth-1}.b
— exactly 2 · depth sites per forward call.
Attributes:
| Name | Type | Description |
|---|---|---|
specs |
tuple[SirenLayerSpec, ...]
|
Tuple of per-layer specs (static). Holds each layer's
|
in_features |
int
|
Input dimension. |
hidden_features |
int
|
Hidden dimension. |
out_features |
int
|
Output dimension. |
depth |
int
|
Total layers including readout. Must be ≥ 2. |
first_omega |
float
|
Frequency multiplier for the first layer. |
hidden_omega |
float
|
Frequency multiplier for hidden layers. |
prior_std |
float
|
Scale factor for the regime-scaled Normal prior (default 1.0). |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.random as jr, jax.numpy as jnp
>>> from numpyro import handlers
>>> net = BayesianSIREN.init(2, 32, 1, depth=3)
>>> with handlers.seed(rng_seed=0):
... y = net(jnp.zeros((4, 2)))
>>> y.shape
(4, 1)
Source code in packages/pyrox-nn/src/pyrox_nn/_siren.py
30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 | |
init(in_features: int, hidden_features: int, out_features: int, *, depth: int, first_omega: float = 30.0, hidden_omega: float = 30.0, c: float = 6.0, prior_std: float = 1.0, pyrox_name: str | None = None) -> BayesianSIREN
classmethod
¶
Construct a BayesianSIREN.
All weights come from the prior, so no PRNG key is needed at
construction time — the key enters when sampling inside a
numpyro handler (handlers.seed, SVI, etc.).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
in_features
|
int
|
Input dimension. |
required |
hidden_features
|
int
|
Hidden dimension. |
required |
out_features
|
int
|
Output dimension. |
required |
depth
|
int
|
Total layers including readout. Must be ≥ 2. |
required |
first_omega
|
float
|
Frequency for the first layer. |
30.0
|
hidden_omega
|
float
|
Frequency for hidden layers. |
30.0
|
c
|
float
|
Theorem-1 constant. |
6.0
|
prior_std
|
float
|
Scale factor for the Normal priors (default 1.0, must be > 0). |
1.0
|
pyrox_name
|
str | None
|
Optional explicit scope name for NumPyro. |
None
|
Returns:
| Type | Description |
|---|---|
BayesianSIREN
|
Initialised |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in packages/pyrox-nn/src/pyrox_nn/_siren.py
SNGP — spectral-normalised GP head¶
The SNGP output layer (Liu et al., 2020): a random-feature GP last layer
whose posterior covariance comes from a Laplace approximation
(LaplaceRandomFeatureCovariance, re-exported from geonnax), giving
distance-aware uncertainty from a single deterministic forward pass.
RandomFeatureGaussianProcess
¶
Bases: PyroxModule
SNGP output layer (Liu et al., 2020).
A random Fourier feature map followed by a learnable linear head, plus a Laplace-approximation covariance over the linear weights. The forward pass returns the mean prediction and (optionally) a per-input predictive variance summarising distance from the training distribution.
Forward (mean):
The frequencies \(W\) and bias \(b\) of the RFF map are
frozen (they implicitly define the kernel approximation): they
are registered as pyrox_param sites for substitution and
checkpointing, then guarded with jax.lax.stop_gradient
inside feature_map so SGD-style optimisers leave them
untouched. The lengthscale \(\ell\), the linear head
\(H, b_H\), and the Laplace precision are the trainable /
updated quantities.
Predictive variance — when \(\hat{\Lambda}\) is the current precision matrix:
Training pattern (one minibatch):
mean = layer(x)registers / reuses the trainable params and returns the mean prediction. Compute the loss, take a gradient step on the SVI parameter store as usual.- After the gradient step, call
new_layer = layer.update_precision(features)wherefeaturesis the result offeature_mapevaluated on the same minibatch using the updated parameters. This returns a new layer with the LRFC's precision EMA-updated.
At inference, mean, var = layer(x, return_cov=True) produces
the mean and the Laplace per-input predictive variance.
Plate semantics
Same as the rest of pyrox_nn's Bayesian / heteroscedastic
dense layers — call this layer outside
numpyro.plate("data", ..., subsample_size=...) and only
plate the observation likelihood.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D_\mathrm{in}\). |
num_features |
int
|
Number of random Fourier features \(D\). |
out_features |
int
|
Output dimension \(D_\mathrm{out}\). |
init_lengthscale |
float
|
Initial lengthscale \(\ell\). Optimised
during training as a positive |
W_init |
Float[Array, 'D_in D']
|
Frozen RFF frequencies, shape |
bias_init |
Float[Array, ' D']
|
Frozen RFF biases, shape |
output_linear_init |
Float[Array, 'D D_out']
|
Init for the linear head, shape |
covariance |
LaplaceRandomFeatureCovariance
|
The |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
References
Liu, J. Z., et al. (2020). Simple and Principled Uncertainty Estimation with Deterministic Deep Learning via Distance Awareness. NeurIPS.
Source code in packages/pyrox-nn/src/pyrox_nn/_sngp.py
31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 | |
init(key: PRNGKeyArray, in_features: int, num_features: int, out_features: int, *, init_lengthscale: float = 1.0, momentum: float = 0.999, ridge: float = 1.0, head_scale: float = 0.01, pyrox_name: str | None = None) -> RandomFeatureGaussianProcess
classmethod
¶
Construct an SNGP head with frozen RFF freqs and an empty precision.
Source code in packages/pyrox-nn/src/pyrox_nn/_sngp.py
feature_map(x: Float[Array, '*batch D_in']) -> Float[Array, '*batch D']
¶
Random Fourier feature map: \(\phi(x) = \sqrt{2/D}\,\cos(Wx/\ell + b)\).
Frequencies and bias are registered as pyrox_param sites for
substitution / checkpointing, but jax.lax.stop_gradient
is applied so SVI's gradient-based optimisers leave them
frozen at their init values. The lengthscale is the active
bandwidth control and is constrained positive.
Source code in packages/pyrox-nn/src/pyrox_nn/_sngp.py
update_precision(features: Float[Array, '*batch D']) -> RandomFeatureGaussianProcess
¶
Return a new layer with an EMA-updated Laplace precision.
Pure-functional: self is unchanged. Pass features computed
on the current minibatch (e.g. via feature_map) — the
update folds the empirical second moment into the EMA. Call
this once per training batch after the gradient step.
Source code in packages/pyrox-nn/src/pyrox_nn/_sngp.py
LaplaceRandomFeatureCovariance
¶
Bases: Module
Laplace-approximation precision for an SNGP output head.
Stores the precision matrix \(\hat{\Lambda} \in \mathbb{R}^{D \times D}\) over the linear weights of the output layer. Updated as an exponential moving average of feature outer products during training:
At test time the predictive variance for a feature vector \(\phi(x_*)\) is
computed stably via a Cholesky solve.
The container is pure-functional: update returns a new
instance with an updated precision rather than mutating self,
matching how Equinox composes immutable PyTrees with optimisers.
A small ridge \(\lambda\) initialises the precision at
\(\lambda I\) and is also added at solve-time inside
covariance and variance_at so the Cholesky stays
numerically well-conditioned even after many EMA steps with low
momentum (which would otherwise let the ridge contribution decay
geometrically and the precision approach singularity).
Equivalently \(\hat\Sigma = (\hat\Lambda + \lambda I)^{-1}\) —
the Bayesian-linear-regression interpretation of SNGP, where
\(\lambda I\) is a Gaussian prior precision on the head weights.
Attributes:
| Name | Type | Description |
|---|---|---|
precision |
Float[Array, 'D D']
|
Current precision matrix \(\hat{\Lambda}\). |
momentum |
float
|
EMA momentum \(m \in [0, 1]\). Higher values give
slower updates; |
ridge |
float
|
Diagonal ridge \(\lambda\). Used both as the init
value of |
Examples:
>>> import jax.numpy as jnp
>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4, ridge=1.0)
>>> cov.precision.shape
(4, 4)
>>> cov.variance_at(jnp.eye(4)).shape
(4,)
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 | |
init(num_features: int, *, momentum: float = 0.999, ridge: float = 1.0) -> LaplaceRandomFeatureCovariance
classmethod
¶
Construct a fresh covariance container with ridge * I precision.
Examples:
>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4, momentum=0.9)
>>> cov.precision.shape
(4, 4)
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
update(features: Float[Array, 'B D']) -> LaplaceRandomFeatureCovariance
¶
Return a new container with EMA-updated precision.
\(\hat\Lambda \leftarrow m\,\hat\Lambda + (1-m)\,\Phi^\top\Phi/B\).
Examples:
>>> import jax.numpy as jnp
>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4)
>>> new = cov.update(jnp.ones((8, 4)))
>>> new is cov
False
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
covariance() -> Float[Array, 'D D']
¶
Inverse of the precision matrix (one-shot Cholesky inversion).
\(\hat\Sigma = (\hat\Lambda + \lambda I)^{-1}\), shape (D, D).
Examples:
>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4)
>>> cov.covariance().shape
(4, 4)
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
variance_at(features: Float[Array, 'N D']) -> Float[Array, ' N']
¶
Per-row predictive variance \(\phi(x_n)^\top \hat{\Sigma}\,\phi(x_n)\).
Computed via a triangular solve to avoid materialising the full \(D \times D\) covariance:
Examples:
>>> import jax.numpy as jnp
>>> from geonnax.sngp import LaplaceRandomFeatureCovariance
>>> cov = LaplaceRandomFeatureCovariance.init(4)
>>> cov.variance_at(jnp.eye(4)).shape
(4,)
Source code in .venv/lib/python3.12/site-packages/geonnax/sngp.py
Deep spectral GPs¶
DeepVSSGP
¶
Bases: PyroxModule
Deep Random Feature Expansion for Variational SSGP (Cutajar et al. 2017).
A stack of \(L\) variational SSGP layers, each with random spectral frequencies \(\Omega_l\) and random projection weights \(W_l\):
Each layer registers three sample sites:
layer_{l}.W_freq— RFF frequencies, prior \(\mathcal{N}(0, 1)\) (RBF spectral density in lengthscale-1 units).layer_{l}.lengthscale— kernel lengthscale, prior \(\mathrm{LogNormal}(\log \ell_{\mathrm{init}}, 1)\).layer_{l}.W_proj— projection weights, prior \(\mathcal{N}(0, \sigma_W^2)\).
Under SVI an
AutoNormal learns mean-field
Gaussian posteriors over all \(3L\) sites — one MC sample per
forward pass gives the doubly-stochastic reparameterised ELBO of
Cutajar et al. (2017).
At depth=1 this reduces to a single VSSGP layer mapping
in_features -> out_features via the RFF basis (same model class
as VariationalFourierFeatures followed by a
DenseReparameterization head). Stacking adds
non-stationarity at the cost of a non-Gaussian aggregate likelihood
— the layer-wise marginalisation that makes single-layer SSGP
closed-form is no longer available, hence the variational
treatment.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D_{\mathrm{in}}\). |
hidden_features |
int
|
Inter-layer dimension \(D_h\) (constant across hidden layers). |
out_features |
int
|
Output dimension \(D_{\mathrm{out}}\). |
n_features |
int
|
Per-layer Fourier-feature pair count \(M\) (so each layer's hidden state is \(2M\)-dim before projection). |
depth |
int
|
Total number of stacked SSGP layers \(L\). Must be \(\ge 1\). |
init_lengthscale |
float
|
Prior location for each layer's lengthscale. |
prior_std |
float
|
Standard deviation of the per-layer projection prior \(\mathcal{N}(0, \sigma_W^2)\). |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.random as jr, jax.numpy as jnp
>>> from numpyro import handlers
>>> net = DeepVSSGP.init(in_features=2, hidden_features=4,
... out_features=1, depth=3, n_features=16)
>>> with handlers.seed(rng_seed=0):
... y = net(jnp.zeros((8, 2)))
>>> y.shape
(8, 1)
Source code in packages/pyrox-nn/src/pyrox_nn/_vssgp.py
22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 | |
init(in_features: int, hidden_features: int, out_features: int, *, depth: int, n_features: int = 64, lengthscale: float = 1.0, prior_std: float = 1.0, pyrox_name: str | None = None) -> DeepVSSGP
classmethod
¶
Construct a DeepVSSGP.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
in_features
|
int
|
Input dimension. Must be \(\ge 1\). |
required |
hidden_features
|
int
|
Hidden dimension. Must be \(\ge 1\). |
required |
out_features
|
int
|
Output dimension. Must be \(\ge 1\). |
required |
depth
|
int
|
Total stacked SSGP layers (including readout). Must be \(\ge 1\). |
required |
n_features
|
int
|
Per-layer Fourier-feature pair count. Must be \(\ge 1\). |
64
|
lengthscale
|
float
|
Prior location for each layer's lengthscale. Must be \(> 0\). |
1.0
|
prior_std
|
float
|
Per-layer projection prior standard deviation. Must be \(> 0\). |
1.0
|
pyrox_name
|
str | None
|
Optional explicit scope name for NumPyro site registration. |
None
|
Returns:
| Type | Description |
|---|---|
DeepVSSGP
|
Initialised |
Raises:
| Type | Description |
|---|---|
ValueError
|
If |
Source code in packages/pyrox-nn/src/pyrox_nn/_vssgp.py
Ensembles — BatchEnsemble / rank-1¶
Efficient deep ensembles that share one weight matrix and learn per-member rank-1 perturbations (Wen et al., 2020; Dusenberry et al., 2020).
DenseRank1
¶
Bases: PyroxModule
Rank-1 ensemble dense layer.
Implements the BatchEnsemble (Wen et al., 2020) / rank-1 BNN (Dusenberry et al., 2020) parameterization: a single shared kernel \(W \in \mathbb{R}^{D_\mathrm{in} \times D_\mathrm{out}}\) and per-member rank-1 multiplicative perturbations \(s_i \in \mathbb{R}^{D_\mathrm{in}}\), \(r_i \in \mathbb{R}^{D_\mathrm{out}}\) for \(i = 1, \ldots, M\). The per-member effective weight is
and the efficient forward pass avoids materialising \(W_i\):
Two modes via the bayesian flag:
bayesian=False(default) — BatchEnsemble. \(r, s, W, b\) are all deterministicpyrox_paramsites and per-member diversity comes purely from the random initialisation of \(r_i, s_i\). Use this for ensemble training under a single shared SGD trajectory.bayesian=True— rank-1 BNN. \(r, s\) arepyrox_samplesites with Normal priors centered at the per-member init values; \(W, b\) remain deterministic. Plug into NumPyro's SVI machinery (anAutoNormalguide onr, srecovers Dusenberry et al., 2020).
Plate semantics
Identical to other pyrox Bayesian dense layers — call this
layer outside numpyro.plate("data", ..., subsample_size=...)
and only plate the observation likelihood. The model log
density picks up \(\log p(r_i)\) and \(\log p(s_i)\)
once per layer (not once per example) under the canonical
pattern.
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D_\mathrm{in}\). |
out_features |
int
|
Output dimension \(D_\mathrm{out}\). |
ensemble_size |
int
|
Number of ensemble members \(M\). |
bias |
bool
|
Whether to include a per-member bias. |
bayesian |
bool
|
If |
prior_scale |
float
|
Std of the Bayesian priors on \(r, s\). Only
used when |
W_init |
float
|
Shared kernel init, shape |
r_init |
float
|
Per-member output-side init, shape |
s_init |
float
|
Per-member input-side init, shape |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.random as jr
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> layer = DenseRank1.init(
... jr.PRNGKey(0),
... in_features=4,
... out_features=2,
... ensemble_size=3,
... )
>>> x = jnp.ones((5, 4))
>>> with handlers.seed(rng_seed=0):
... y = layer(x)
>>> y.shape
(3, 5, 2)
References
Wen, Y., Tran, D., & Ba, J. (2020). BatchEnsemble: An Alternative Approach to Efficient Ensemble and Lifelong Learning. ICLR.
Dusenberry, M. W., et al. (2020). Efficient and Scalable Bayesian Neural Nets with Rank-1 Factors. ICML.
Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 | |
init(key: PRNGKeyArray, in_features: int, out_features: int, ensemble_size: int, *, bias: bool = True, bayesian: bool = False, init_scale: float = 0.5, prior_scale: float = 0.5, pyrox_name: str | None = None) -> DenseRank1
classmethod
¶
Construct a layer with random per-member init vectors.
Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
LayerNormEnsemble
¶
Bases: PyroxModule
Per-ensemble-member LayerNorm.
Drop-in replacement for LayerNorm inside BatchEnsemble / Rank1
architectures. Computes the standard LayerNorm normalisation over
the trailing feature dimension and applies a per-member affine
transform — each ensemble member \(i \in \{1, \ldots, M\}\)
gets its own learnable scale \(\gamma_i \in \mathbb{R}^D\)
and bias \(\beta_i \in \mathbb{R}^D\):
where \(\mu\) and \(\sigma^2\) are the empirical mean and
variance over the trailing feature axis (computed independently
for each member-batch slice). Without per-member scale/bias,
sharing a single LayerNorm across the ensemble would couple all
members and erase the diversity introduced by DenseRank1
or any other BatchEnsemble layer upstream.
Input is expected to carry a leading ensemble axis of size
ensemble_size and a trailing feature axis of size
feature_dim. Any number of intermediate batch / time axes
are supported and pass through unchanged.
Attributes:
| Name | Type | Description |
|---|---|---|
ensemble_size |
int
|
Number of ensemble members \(M\). |
feature_dim |
int
|
Trailing feature dimension \(D\) over which the normalisation is computed. |
eps |
float
|
Small positive constant added to the variance for numerical stability. |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> ln = LayerNormEnsemble(
... ensemble_size=3, feature_dim=4, pyrox_name="ln"
... )
>>> x = jnp.ones((3, 5, 4)) # (M, batch, D)
>>> with handlers.seed(rng_seed=0):
... y = ln(x)
>>> y.shape
(3, 5, 4)
Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 | |
MultiHeadAttentionBE
¶
Bases: PyroxModule
Multi-head attention with BatchEnsemble rank-1 projections.
Standard scaled-dot-product multi-head attention where each of the four linear projections — query, key, value, and output — uses a BatchEnsemble parameterisation: a shared full-rank kernel plus per-ensemble-member rank-1 multiplicative perturbations. So for member \(i \in \{1, \ldots, M\}\) and projection \(P \in \{Q, K, V, O\}\),
and the attention itself is the usual
The forward consumes un-ensembled inputs (query, key,
value of shape (T, D) / (S, D)), adds the ensemble
axis when projecting to Q, K, V, runs per-member
attention in parallel, and returns the per-member output of
shape (M, T, D). Equivalent to running M independent
attention heads with rank-1 weight perturbations and stacking
their outputs.
Plate semantics
Same convention as DenseRank1 and the rest of the
pyrox_nn ensemble / Bayesian dense family — call this
layer outside numpyro.plate("data", ..., subsample_size=...)
and only plate the observation likelihood. All four projections
register their parameters as pyrox_param sites; nothing
about the layer is data-dependent so plate-scaling does not
come into play unless the user puts the call inside a
subsampled plate.
Attributes:
| Name | Type | Description |
|---|---|---|
embed_dim |
int
|
Total feature dimension \(D\) of query / key /
value (must be divisible by |
num_heads |
int
|
Number of attention heads \(H\). Each head sees
|
ensemble_size |
int
|
Number of ensemble members \(M\). |
bias |
bool
|
Whether each of the four projections includes a
per-member bias. When |
q_init |
/ k_init / v_init / o_init
|
Per-projection
BatchEnsemble init arrays. Build via |
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.random as jr
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> mha = MultiHeadAttentionBE.init(
... jr.PRNGKey(0),
... embed_dim=8, num_heads=2, ensemble_size=3,
... )
>>> x = jnp.ones((5, 8))
>>> with handlers.seed(rng_seed=0):
... y = mha(x, x, x) # self-attention
>>> y.shape
(3, 5, 8)
Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 | |
init(key: PRNGKeyArray, embed_dim: int, num_heads: int, ensemble_size: int, *, bias: bool = True, init_scale: float = 0.5, pyrox_name: str | None = None) -> MultiHeadAttentionBE
classmethod
¶
Construct an MHA-BE layer with random Q/K/V/O projection inits.
Source code in packages/pyrox-nn/src/pyrox_nn/_ensemble.py
Heteroscedastic output heads¶
Monte-Carlo sigmoid / softmax output layers with factor-analysis noise (Collier et al., 2021) for input-dependent label noise.
MCSigmoidDenseFA
¶
Bases: _HeteroscedasticBase
Heteroscedastic multi-label output layer (FA noise + sigmoid).
Identical low-rank-plus-diagonal logit-noise model as
MCSoftmaxDenseFA, but the per-class outputs are
independent Bernoullis — final probabilities are the MC average of
element-wise sigmoids, not a softmax. Use this for multi-label
classification or independent binary heads.
See MCSoftmaxDenseFA for the noise model, plate semantics,
init API, and references.
Source code in packages/pyrox-nn/src/pyrox_nn/_heteroscedastic.py
MCSoftmaxDenseFA
¶
Bases: _HeteroscedasticBase
Heteroscedastic multi-class output layer (FA noise + softmax).
Implements Collier et al. (2021): the logit covariance is input-dependent low-rank-plus-diagonal,
where \(V(x) = \mathrm{reshape}(W_V x + b_V, [C, r])\) and \(\sigma(x) = \exp(W_\sigma x + b_\sigma)\). Output is the Monte Carlo average of softmaxed perturbed logits
All linear factors are deterministic pyrox_param sites — the
layer is heteroscedastic but not Bayesian over its weights. Use it
as a drop-in head for classification when label noise is known to
be input-dependent (label disagreement, fine-grained categories).
Plate semantics
Same as other pyrox_nn Bayesian dense layers — call
outside numpyro.plate("data", ..., subsample_size=...) so
the parameter sites are unscaled. The MC noise is drawn from
numpyro.prng_key().
Attributes:
| Name | Type | Description |
|---|---|---|
in_features |
int
|
Input dimension \(D_\mathrm{in}\). |
num_classes |
int
|
Number of classes \(C\). |
rank |
int
|
Rank \(r\) of the low-rank factor \(V(x)\). |
num_mc_samples |
int
|
Number of MC softmax samples \(S\) per forward call. |
diag_init_bias |
float
|
Initial value for the diagonal-scale bias
|
pyrox_name |
str | None
|
Explicit scope name for NumPyro site registration. |
Examples:
>>> import jax.random as jr
>>> import jax.numpy as jnp
>>> from numpyro import handlers
>>> layer = MCSoftmaxDenseFA.init(
... jr.PRNGKey(0), in_features=4, num_classes=3, rank=2,
... )
>>> x = jnp.ones((5, 4))
>>> with handlers.seed(rng_seed=0):
... probs = layer(x)
>>> probs.shape
(5, 3)
>>> bool(jnp.allclose(probs.sum(axis=-1), 1.0))
True
References
Collier, M., Mustafa, B., Kokiopoulou, E., Jenatton, R., & Berent, J. (2021). Correlated Input-Dependent Label Noise in Large-Scale Image Classification. CVPR.
Source code in packages/pyrox-nn/src/pyrox_nn/_heteroscedastic.py
Bayesian Neural Field stack¶
Standardization
¶
Bases: PyroxModule
Apply a fixed-coefficient affine standardization.
Both mu and std are static (fit-time) constants, not
learned. Use pyrox_nn.preprocessing.fit_standardization to
construct from a pandas DataFrame.
Attributes:
| Name | Type | Description |
|---|---|---|
mu |
Float[Array, ' D']
|
Per-feature mean, shape |
std |
Float[Array, ' D']
|
Per-feature standard deviation, shape |
pyrox_name |
str | None
|
Optional override for the per-instance scope name. |
Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
FourierFeatures
¶
Bases: PyroxModule
Per-input dyadic-frequency Fourier basis.
For each input column, evaluates 2 * degree Fourier features at
frequencies \(2\pi \cdot 2^d\) for \(d \in \{0, \dots,
\text{degree} - 1\}\). Concatenated across all columns.
Wraps pyrox_nn._features.fourier_features per input
dimension.
Attributes:
| Name | Type | Description |
|---|---|---|
degrees |
tuple[int, ...]
|
Number of dyadic frequencies per input column, as a
Python |
rescale |
bool
|
If |
pyrox_name |
str | None
|
Optional scope-name override. |
Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
SeasonalFeatures
¶
Bases: PyroxModule
Period-and-harmonic cos/sin basis on a scalar time axis.
For each period \(\tau_p\) with \(H_p\) harmonics, emits
2 * H_p cos/sin columns. Total output width is \(2 \sum_p
H_p\).
Wraps pyrox_nn._features.seasonal_features. Periods and
harmonics are kept as Python tuples (static) so the inner shape
structure is known at trace time.
Attributes:
| Name | Type | Description |
|---|---|---|
periods |
tuple[float, ...]
|
Period values, |
harmonics |
tuple[int, ...]
|
Harmonics per period, |
rescale |
bool
|
If |
pyrox_name |
str | None
|
Optional scope-name override. |
Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
InteractionFeatures
¶
Bases: PyroxModule
Element-wise products on selected pairs of input columns.
Wraps pyrox_nn._features.interaction_features.
Attributes:
| Name | Type | Description |
|---|---|---|
pairs |
tuple[tuple[int, int], ...]
|
Index pairs, |
pyrox_name |
str | None
|
Optional scope-name override. |
Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
BayesianNeuralField
¶
Bases: PyroxModule
The full Bayesian Neural Field architecture.
A spatiotemporal MLP with:
- A learned per-input log-scale adjustment (Logistic(0, 1) prior).
- Four feature blocks concatenated into
h_0: rescaled inputs, Fourier features, seasonal features, interaction products. - Per-block
softplus(feature_gain)modulation. - A depth-
LMLP whose layers are \(h_{\ell+1} = \sigma_\alpha\bigl(g_\ell \cdot W_\ell\, h_\ell / \sqrt{\lvert h_\ell \rvert}\bigr)\), where \(\sigma_\alpha = \mathrm{sig}(\beta) \cdot \mathrm{elu} + (1 - \mathrm{sig}(\beta)) \cdot \mathrm{tanh}\) is a learned mixed activation. - A final linear layer scaled by
softplus(output_gain).
All weights, biases, gains, scales, and the activation logit carry
independent \(\mathrm{Logistic}(0, 1)\) priors registered via
PyroxModule.pyrox_sample.
The \(1/\sqrt{\text{fan-in}}\) pre-normalization is the standard NTK-scaling trick — it makes the layer-wise prior predictive a fan-in-independent Gaussian process in the infinite-width limit (Lee et al., 2018).
Attributes:
| Name | Type | Description |
|---|---|---|
input_scales |
tuple[float, ...]
|
Per-input fixed scale (typically training-data
inter-quartile range). Static |
fourier_degrees |
tuple[int, ...]
|
Per-input number of dyadic Fourier
frequencies. Static |
interactions |
tuple[tuple[int, int], ...]
|
Pair-index list for interaction features. Static
|
seasonality_periods |
tuple[float, ...]
|
Periods for seasonal features. Static
|
num_seasonal_harmonics |
tuple[int, ...]
|
Harmonics per period. Static
|
width |
int
|
Hidden layer width. |
depth |
int
|
Number of hidden MLP layers. |
time_col |
int
|
Index of the time column inside |
pyrox_name |
str | None
|
Optional scope-name override. |
Source code in packages/pyrox-nn/src/pyrox_nn/_bnf.py
156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 | |
Pure-JAX feature helpers¶
fourier_features(x: Float[Array, ' N'], max_degree: int, *, rescale: bool = False) -> Float[Array, 'N two_max_degree']
¶
Cos/sin Fourier basis at dyadic frequencies.
For each input element and each degree \(d \in \{0, \dots, D-1\}\), evaluates
Returns the columns concatenated as [cos_0, ..., cos_{D-1},
sin_0, ..., sin_{D-1}], matching Google's bayesnf layout.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Float[Array, ' N']
|
Length- |
required |
max_degree
|
int
|
Number of dyadic frequencies |
required |
rescale
|
bool
|
If |
False
|
Returns:
| Type | Description |
|---|---|
Float[Array, 'N two_max_degree']
|
Array of shape |
Examples:
>>> import jax.numpy as jnp
>>> from geonnax.basis import fourier_features
>>> fourier_features(jnp.linspace(0.0, 1.0, 5), max_degree=3).shape
(5, 6)
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
seasonal_features(x: Float[Array, ' N'], periods: Sequence[float], harmonics: Sequence[int], *, rescale: bool = False) -> Float[Array, 'N two_F']
¶
Cos/sin features at multiples of \(2\pi / \tau_p\).
For each period \(\tau_p\) with \(H_p\) harmonics, evaluates
for \(h = 1, \dots, H_p\). Returns the cos columns concatenated with the sin columns, length \(F = \sum_p H_p\) each.
periods and harmonics are Python sequences (tuples,
lists, or 0-d JAX arrays wrapped at the call site). Keeping them as
Python values lets the function run cleanly under jax.jit and
lax.scan without triggering a concretization error.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Float[Array, ' N']
|
Time/index input, shape |
required |
periods
|
Sequence[float]
|
Period values. |
required |
harmonics
|
Sequence[int]
|
Harmonics per period. |
required |
rescale
|
bool
|
If |
False
|
Returns:
| Type | Description |
|---|---|
Float[Array, 'N two_F']
|
Array of shape |
Examples:
>>> import jax.numpy as jnp
>>> from geonnax.basis import seasonal_features
>>> x = jnp.linspace(0.0, 10.0, 4)
>>> # F = 1 + 2 = 3 frequencies -> 2 * F = 6 columns
>>> seasonal_features(x, [7.0, 365.0], [1, 2]).shape
(4, 6)
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
seasonal_frequencies(periods: Sequence[float], harmonics: Sequence[int]) -> tuple[list[int], list[float]]
¶
Flatten (period, harmonic_count) pairs into Python lists.
For each period \(\tau_p\) with \(H_p\) harmonics, emits frequencies \(f_{p, h} = h / \tau_p\) for \(h = 1, \dots, H_p\). The total length is \(F = \sum_p H_p\).
Inputs are Python sequences, not JAX arrays, so this helper
runs at trace time and never triggers a concretization error under
jax.jit. Most callers won't use it directly; it's exposed for
symmetry with seasonal_features.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
periods
|
Sequence[float]
|
Period values. |
required |
harmonics
|
Sequence[int]
|
Number of harmonics per period. |
required |
Returns:
| Type | Description |
|---|---|
list[int]
|
|
list[float]
|
\(F = \sum_p H_p\). |
Examples:
>>> from geonnax.basis import seasonal_frequencies
>>> # periods (7, 365) with (1, 2) harmonics -> F = 1 + 2 = 3 freqs
>>> idx, freqs = seasonal_frequencies([7.0, 365.0], [1, 2])
>>> idx
[0, 1, 1]
>>> len(freqs)
3
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
interaction_features(x: Float[Array, 'N D'], pairs: Int[Array, 'K 2']) -> Float[Array, 'N K']
¶
Element-wise products on selected pairs of input columns.
For each pair \((i, j)\) and each row \(n\), computes \(x_{n, i} \cdot x_{n, j}\).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
x
|
Float[Array, 'N D']
|
Input matrix, shape |
required |
pairs
|
Int[Array, 'K 2']
|
Index pairs, shape |
required |
Returns:
| Type | Description |
|---|---|
Float[Array, 'N K']
|
Array of shape |
Examples:
>>> import jax.numpy as jnp
>>> from geonnax.basis import interaction_features
>>> x = jnp.arange(6.0).reshape(2, 3) # (N=2, D=3)
>>> pairs = jnp.array([[0, 1], [0, 2]]) # (K=2, 2)
>>> interaction_features(x, pairs).shape
(2, 2)
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
standardize(x: Float[Array, '*shape'], mu: Float[Array, '*shape'], std: Float[Array, '*shape']) -> Float[Array, '*shape']
¶
Affine standardize: (x - mu) / std.
Broadcasts mu and std against x per the JAX broadcasting
rules. std is not clamped; pass a positive value or guard
upstream.
Examples:
>>> import jax.numpy as jnp
>>> from geonnax.basis import standardize
>>> x = jnp.array([1.0, 3.0, 5.0])
>>> bool(jnp.allclose(standardize(x, 3.0, 2.0), jnp.array([-1, 0, 1])))
True
Source code in .venv/lib/python3.12/site-packages/geonnax/basis.py
unstandardize(z: Float[Array, '*shape'], mu: Float[Array, '*shape'], std: Float[Array, '*shape']) -> Float[Array, '*shape']
¶
Inverse of standardize: z * std + mu.
Examples:
>>> import jax.numpy as jnp
>>> from geonnax.basis import standardize, unstandardize
>>> x = jnp.array([1.0, 3.0, 5.0])
>>> # unstandardize undoes standardize for the same (mu, std).
>>> z = standardize(x, 3.0, 2.0)
>>> bool(jnp.allclose(unstandardize(z, 3.0, 2.0), x))
True