class AMICAMLXNG:
"""MLX natural-gradient EM backend, full torch-equivalent surface (#76/#81, epic #278).
Parameters are :class:`AMICATorchNG`'s, with the same names, defaults and
validation, except ``device`` and ``dtype``: MLX runs on its default
device (the Apple GPU) in float32 only. The same ``seed`` produces the same
initial parameters as the PyTorch/NumPy backends, so cross-backend
equivalence is testable.
The convergence-stop parameters (issue #248) carry AMICATorchNG's names,
defaults and semantics exactly, so the two backends stop on the same
iteration for the same reason:
``use_min_dll`` (True) / ``min_dll`` (1e-9) / ``maxincs`` (5)
Stop once the per-sample-per-channel log-likelihood gain
``ll_history[-1] - ll_history[-2]`` stays below ``min_dll`` for more than
``maxincs`` *consecutive* iterations (Fortran amica15.f90:1078-1090);
``stop_reason="min_dll"``. The counter resets on any larger gain and a
likelihood decrease counts as a small gain, as in Fortran.
``use_grad_norm`` (True) / ``min_nd`` (1e-7)
Stop once the weight-gradient RMS norm ``ndtmpsum`` falls to or below
``min_nd`` (Fortran amica15.f90:1091-1097); ``stop_reason="grad_norm"``.
The same threshold is also the second half of the likelihood-decrease
stop (amica15.f90:1058, ``stop_reason="grad_norm_floor"``), which runs
regardless of ``use_grad_norm``. Under the shipped defaults
``"grad_norm"`` shadows ``"grad_norm_floor"`` -- see AMICATorchNG's
``use_grad_norm`` docstring, whose precedence note applies verbatim.
``min_nd`` is not reachable on small recordings in any backend
(issue #218); the Fortran-faithful default is kept rather than retuned.
All three checks require two log-likelihood values, so none can fire on the
first iteration (Fortran's ``if (iter > 1)``, amica15.f90:1051).
The component-sharing parameters (issue #263) likewise carry AMICATorchNG's
names, defaults, validation and semantics:
``share_comps`` (False)
Enable multi-model component sharing (Fortran ``share_comps`` /
``identify_shared_comps``, amica15.f90:1916): components that are
near-collinear across different models are merged so they share one
mixing vector (one row of ``A``) and one density. Requires
``n_models >= 2`` (a model cannot share with itself); accepted but
inert otherwise. OFF by default, so default fits are unchanged. The
reference's similarity metric is never initialized (like the dead
``do_choose_pdfs``, #26), so it has no bit-exact oracle; this implements
the intended algorithm on the component rows (issue #334), validated
against the PyTorch backend, whose update from a merged state matches
the reference. A
merge that fires on the LAST fit iteration is reflected in the
returned model but trails in ``final_ll_``; see that attribute's
comment (issue #269).
``share_start`` (100) / ``share_iter`` (100)
Sharing schedule: first iteration to attempt merges (counted from 1)
and the interval between attempts. They also set the reference's
A-freeze, which applies to every fit, sharing on or off (issue #345):
from iteration ``share_start`` on, the update of ``A`` (with its lrate
ramp) is held on every iteration whose number, counted from 1, has a
remainder of 0 to 5 modulo ``share_iter`` (amica15.f90:1803), so with
the defaults on iterations 100-105, 200-205, and so on.
Both are validated whether or not ``share_comps`` is on:
``share_start`` must be an integer >= 1, and ``share_iter`` an
integer >= 7, or the freeze would hold A permanently from
``share_start`` on.
``comp_thresh`` (0.99)
Cosine-similarity cutoff, in the de-sphered (sensor-space) metric, above
which two components' mixing vectors are identified and merged. Must be
in
``(0, 1]``. The de-sphering uses ``pinv(sphere)``, so sharing also works
on rank-reduced and rank-deficient fits (issues #253, #221); see
:meth:`_identify_shared_comps`.
The Newton parameters (issue #264) likewise carry AMICATorchNG's names,
defaults and semantics:
``do_newton`` (False)
Precondition the ``A``/``W`` natural gradient with the approximate
Hessian from iteration ``newt_start`` on (Fortran ``do_newton``: the
2x2 solve and its positive-definiteness guard at amica15.f90:1718-1741,
the ramp and fallback at :1803-1816), which converges faster near the
optimum. OFF by default, as in the reference's compiled default (the
bundled reference ``input.param`` turns it on), and every accumulator
it needs is gated on it, so a
default fit is bit-for-bit what it was before #264.
``newt_start`` (20)
Iteration at which the Newton step switches on (natural gradient runs
before it, letting the mixture parameters settle first). Counted from
1, as the reference counts it (``iter .ge. newt_start``): the Newton
step is taken on the ``newt_start``-th iteration, i.e. once
``iteration + 1 >= newt_start`` for the 0-based ``iteration``
attribute (issue #335). ``newt_start=0`` fits exactly as ``1`` does,
Newton from the first iteration. It also gates the ``maxdecs`` ratchet
of the ``rholrate`` ceiling, independently of ``do_newton``, and under
Newton that of ``newtrate``, to iterations after ``newt_start``
(Fortran amica15.f90:1067/1070), so it is validated (an integer
>= 0) whether or not ``do_newton`` is on.
``newtrate`` (0.5)
Maximum learning rate the ramp climbs to while Newton is active and
positive definite; the natural-gradient phase (and any fallback
iteration) is capped at ``lrate_cap`` instead. A ceiling itself: under
Newton it ratchets down by ``lratefact`` at each ``maxdecs`` cycle
completed after iteration ``newt_start`` (amica15.f90:1070), and is
reset to its constructor value at the start of every fit.
``newt_ramp`` (10)
Denominator of the per-iteration learning-rate ramp toward the current
ceiling: ``lrate = min(ceiling, lrate + min(1/newt_ramp, lrate))``.
Whenever any source pair fails the positive-definiteness guard the whole
model falls back to the natural gradient for that iteration and
``n_newton_fallbacks`` counts it (as AMICATorchNG does), so an all-fallback
run is visible without re-instrumenting.
The rho learning rate carries AMICATorchNG's names and semantics too:
``rholrate`` (0.05) / ``rholratefact`` (0.1)
Learning rate of the generalized-Gaussian shape ``rho``. As in the
reference (``rholrate``/``rholrate0``, amica15.f90:1063-1068), it has
a working value (``rholrate``), which each likelihood decrease
multiplies by ``rholratefact``, and a ceiling (``rholrate_cap``),
which the ``maxdecs`` ratchet multiplies by ``rholratefact`` after
``newt_start``; every update of ``A`` resets the working value to the
ceiling before ``rho`` moves (:1806/:1813), so the decrease scaling
reaches ``rho`` only on an iteration on which ``A`` is held (see
``share_iter``). Both are reset to the constructor value at the start
of every fit and saved with the model (issue #339).
The source-density family parameters (issue #265, porting AMICATorchNG's
issue #26) likewise carry AMICATorchNG's names, defaults and semantics:
``pdftype`` (0)
Per-source density family (Fortran ``amica15.f90`` ``pdtype`` codes): 0
generalized Gaussian (default; rho adapts), 2 Gaussian, 3 logistic, 4
sub-Gaussian cosh+. ``pdftype=1`` enables the extended-Infomax adaptive
switcher, which flips each source between the super-Gaussian (code 1)
and sub-Gaussian (code 4) cosh densities by kurtosis sign (see
:meth:`_choose_pdfs`). The GG shape update is frozen for every non-GG
family (Fortran ``dorho=.false.``, ``self.dorho = pdftype == 0``); the
single-component families 1/4 (and the adaptive mode) require
``n_mix=1``. ``pdftype=0`` stays byte-for-byte the pre-#265
implementation (the ``_pdtype_h`` ``None`` fast path, policy 2 --
verified by a before/after bit-identity check, see
``.context/issue-265/pdf_family_findings.md``). ``rho`` does not
describe the fitted density for codes 1-4: it stays frozen at ``rho0``
and is only ever meaningful for the generalized-Gaussian family
(code 0).
``kurt_start`` (3) / ``num_kurt`` (5) / ``kurt_int`` (1)
Adaptive-switch schedule (only used when ``pdftype=1``): first
iteration to re-estimate kurtosis (counted from 1), number of switch
passes, and the iteration interval between them. ``num_kurt=0``
disables switching (the family stays at its super-Gaussian init). No
bit-exact oracle -- the reference's own switch is dead code
(``do_choose_pdfs`` is set but ``m2sum``/``m4sum`` are never
accumulated, amica15.f90:608-615) -- so this is behavior-validated on
real data (ADR 0002).
The block-size search parameters (issue #232) likewise carry
AMICATorchNG's names, defaults and semantics:
``do_opt_block`` (False)
Time candidate block sizes on the real data and GPU at the start of
``fit`` and keep the fastest, instead of using ``block_size`` as given
(Fortran ``do_opt_block``). MLX is the least block-size-sensitive
backend measured (2.6x from 512 to a single block, against 32x for
PyTorch-MPS, issue #216), so there is less here to win than on the
other backends. The choice is timing-based and therefore
machine-dependent, so it is OFF by default and a run compared against
the reference binary must leave it off and pin ``block_size``.
``blk_min`` (4096) / ``blk_max`` (32768) / ``blk_step`` (4096)
Candidate sweep, Fortran's arithmetic stepping, clamped to
``n_samples`` and to a conservative estimate of what fits in the GPU's
recommended working set. A candidate MLX cannot allocate raises a
catchable ``RuntimeError`` (``[metal::malloc] ...`` -- not the process
abort MLX's LU takes on singular input, issue #274), so it is skipped
and the fit continues at the largest size that ran.
The best-of-N restart parameters (issue #198) likewise carry
AMICATorchNG's names, defaults and semantics:
``n_restarts`` (1)
Number of independent fits to run from different seeds, keeping the one
with the highest ``final_ll_``. ``1`` (the default) bypasses the restart
machinery entirely, so a default fit is bit-for-bit what it was
before #198. ``n_restarts > 1`` requires a base ``seed`` (or explicit
``restart_seeds``) so the winner can be reproduced, and costs
``n_restarts`` times as long (restarts run serially). This is a pamica
extension: Fortran has no search over seeds. See
:mod:`pamica.restarts` and ``docs/guides/amica-differences.md``.
``restart_seeds`` (None)
Explicit per-restart seeds, exactly ``n_restarts`` of them; otherwise
``seed, seed + 1, ...``.
The outlier-rejection parameters (issue #123's AMICATorchNG mechanism,
epic #278 Phase 3/#289) likewise carry AMICATorchNG's names, defaults and
validation:
``do_reject`` (False)
Permanently drop samples whose per-sample log-likelihood is a low
outlier on the ``rejstart``/``rejint``/``maxrej`` schedule (Fortran
``reject_data``, amica15.f90:2380-2464). The rejection statistic is
read FROM the LLt stash (``_llt_ll``, issue #157) rather than a
second forward pass over the good set -- the NumPy backend's design,
pre-empting AMICATorchNG's own open follow-up (issue #298) to drop
its separate ``_sample_ll`` pass; the statistic is mathematically
identical either way. OFF by default, so a default fit is unaffected.
``rejsig`` (3.0)
Reject a sample when its log-likelihood is below ``mean - rejsig *
std`` (population std) over the current good set.
``rejstart`` (2) / ``rejint`` (3) / ``maxrej`` (1)
Rejection schedule: first iteration to reject (counted from 1, as the
reference counts it; issue #335), interval between subsequent passes,
and the maximum number of passes. ``rejstart`` must be an integer
>= 1 when ``do_reject`` is on: ``rejstart <= 0`` would silently skip
the reference's unconditional first pass.
The best-iterate safeguard (issue #51, epic #278 Phase 2/#288) likewise
carries AMICATorchNG's name, default and semantics:
``keep_best`` (True)
Return the highest-log-likelihood iterate instead of the last one. The
lrate schedule is non-monotone (it anneals only after an LL
*decrease*), so a late Newton-fallback overshoot can leave the final
iterate below a peak the run already reached. When the final LL falls
more than a small tolerance below that peak, ``fit`` restores the
peak's parameters. A monotone single-model run (issue #24 parity) is a
bit-exact no-op. Automatically inactive under ``share_comps`` (a merge
changes the parameter count, so pre- and post-merge LLs are not
comparable and reverting to an earlier snapshot would silently undo
the merge; issue #269) and under ``do_reject`` (the good-sample set,
and so the LL normalization, changes across iterations, so
per-iteration LLs are not comparable), matching AMICATorchNG.
The explicit PCA-reduction parameters (issue #323) likewise carry
AMICATorchNG's names, defaults, validation and semantics. The validation
and the precedence rule are the shared :mod:`pamica.rank` policy, so the
three array backends cannot disagree on them:
``pcakeep`` (None)
Keep this many principal dimensions (Fortran ``pcakeep``). Must be an
integer >= 1, or the constructor raises ``ValueError``. Capped by the
detected numerical rank, matching Fortran's ``numeigs = min(pcakeep,
count(eigs > mineig))`` (amica15.f90:413), so ``pcakeep >=
n_channels`` keeps every dimension the data have.
``pcadb`` (None)
Keep the dimensions whose covariance eigenvalue lies within ``pcadb``
dB of the largest. Must be a finite number > 0, or the constructor
raises ``ValueError``. A pamica extension: the reference parses
``pcadb`` (amica15.f90:3459-3461) but never uses it. When both are
set, ``pcakeep`` takes precedence and ``pcadb`` is ignored (one INFO
log line), as in the reference.
Both are ignored, with one WARNING, when ``do_sphere=False``: reduction
happens only while sphering, as in the reference (``numeigs = nx`` in its
no-sphere branch, amica15.f90:527).
With both ``None`` (the default) only automatic ``mineig``/``mineig_rel``
rank detection sizes the model, exactly as before #323. A reduced fit has
a non-square ``(n_channels, n_channels_in)`` sphere; map its components
back to the input channels with :meth:`get_sensor_mixing_matrix`.
``fit(mir_step > 0)`` rejects an explicit reduction request up front, as
AMICATorchNG does (see :meth:`_fit_once`).
The rescale parameters (issue #333) likewise carry AMICATorchNG's names,
defaults, validation and semantics:
``doscaling`` (True)
Rescale each component's mixing vector (a row of the stored ``A``) to
unit norm, with the matching ``mu``/``beta`` rescale,
an exact change of scale (see :meth:`_rescale_components`).
``scalestep`` (1)
Run the rescale on iterations ``scalestep``, ``2*scalestep``, ...
counted from 1; the default 1 rescales every iteration, as the
reference always does (it ignores ``scalestep``). Validated only when
``doscaling`` is on (an integer >= 1, or the constructor raises
``ValueError``); with ``doscaling`` off it is inert, never read.
"""
def __init__(
self,
n_channels: int,
n_models: int = 1,
n_mix: int = 3,
block_size: int = 8192,
do_opt_block: bool = False,
blk_min: int = blocktune.DEFAULT_BLK_MIN,
blk_max: int = blocktune.DEFAULT_BLK_MAX,
blk_step: int = blocktune.DEFAULT_BLK_STEP,
lrate: float = 0.1,
minlrate: float = 1e-12,
lratefact: float = 0.5,
maxdecs: int = 5,
use_min_dll: bool = True,
min_dll: float = 1e-9,
maxincs: int = 5,
use_grad_norm: bool = True,
min_nd: float = 1e-7,
newt_ramp: int = 10,
newt_start: int = 20,
newtrate: float = 0.5,
do_newton: bool = False,
do_reject: bool = False,
rejsig: float = 3.0,
rejstart: int = 2,
rejint: int = 3,
maxrej: int = 1,
rho0: float = 1.5,
minrho: float = 1.0,
maxrho: float = 2.0,
rholrate: float = 0.05,
rholratefact: float = 0.1,
pdftype: int = 0,
kurt_start: int = 3,
num_kurt: int = 5,
kurt_int: int = 1,
invsigmin: float = 1e-4,
invsigmax: float = 1000.0,
doscaling: bool = True,
scalestep: int = 1,
share_comps: bool = False,
share_start: int = 100,
share_iter: int = 100,
comp_thresh: float = 0.99,
do_mean: bool = True,
do_sphere: bool = True,
do_approx_sphere: bool = True,
pcakeep: Optional[int] = None,
pcadb: Optional[float] = None,
mineig: float = MINEIG,
mineig_rel: Optional[float] = MINEIG_REL,
seed: Optional[int] = None,
n_restarts: int = restarts.DEFAULT_N_RESTARTS,
restart_seeds: Optional[Sequence[int]] = None,
keep_best: bool = True,
):
self.n_channels = n_channels
# The input channel count, kept apart from n_channels, which
# _preprocess shrinks to the kept rank on any rank reduction; same
# role as AMICATorchNG._n_input_channels (see _fit_once).
self._n_input_channels = n_channels
self.n_models = n_models # multi-model (#81) + component sharing (#263)
self.n_mix = n_mix
self.n_comps = n_channels * n_models
self.block_size = block_size
self.do_opt_block = do_opt_block
self.blk_min = blk_min
self.blk_max = blk_max
self.blk_step = blk_step
if do_opt_block:
# Validated only when the search is on, matching share_comps and
# AMICATorchNG: inert otherwise, and a Fortran input.param carrying
# these alongside do_opt_block=0 must stay loadable (issue #232).
blocktune.validate_block_tune_params(blk_min, blk_max, blk_step)
self.lrate0 = lrate
self.lrate = lrate
self.lrate_cap = lrate
self.minlrate = minlrate
self.lratefact = lratefact
self.maxdecs = maxdecs
# Convergence stops (issue #248), same names/defaults/semantics as
# AMICATorchNG: the small-likelihood-gain stop (use_min_dll/min_dll/
# maxincs) and the weight-gradient-norm stop (use_grad_norm/min_nd). See
# fit() for the per-iteration checks and _update_parameters for the
# ndtmpsum computation they both read.
if maxincs < 0:
raise ValueError(f"maxincs must be >= 0, got {maxincs}")
self.use_min_dll = use_min_dll
self.min_dll = min_dll
self.maxincs = maxincs
self.use_grad_norm = use_grad_norm
self.min_nd = min_nd
self.newt_ramp = newt_ramp
# Newton schedule (issue #264), same names/defaults/semantics as
# AMICATorchNG. newt_start doubles as the gate on the rholrate ceiling
# ratchet, which Fortran conditions on iter > newt_start independently of
# do_newton (amica15.f90:1067) -- so it stays meaningful for a
# natural-gradient fit too. newtrate is a CEILING that ratchets down at
# maxdecs during fit, so keep the constructor value for the per-fit reset
# in _initialize_parameters (as lrate0/rholrate0 do). Validated whether
# or not do_newton is on, for that same ratchet gate.
schedule.validate_iteration_setting("newt_start", newt_start, 0)
self.newt_start = newt_start
self.newtrate = newtrate
self.newtrate0 = newtrate
self.do_newton = do_newton
# Iterations on which the Newton direction was rejected as not positive
# definite and the natural gradient used instead (AMICATorchNG's counter
# of the same name); reset per fit in _initialize_parameters.
self.n_newton_fallbacks = 0
# Outlier rejection (issue #123's AMICATorchNG mechanism, epic #278
# Phase 3/#289), same names/defaults/validation as AMICATorchNG
# (``AMICATorchNG.__init__``; Fortran do_reject/rejsig/
# rejstart/rejint/maxrej, amica15_header.f90). numrej/good_idx are set
# up per fit in _fit_once (good_idx = None until then, matching the
# do_reject=False no-op path).
self.do_reject = do_reject
self.rejsig = rejsig
self.rejstart = rejstart
self.rejint = rejint
self.maxrej = maxrej
if do_reject:
if rejint < 1:
raise ValueError(f"rejint must be >= 1, got {rejint}")
if rejsig <= 0:
raise ValueError(f"rejsig must be > 0, got {rejsig}")
if maxrej < 0:
raise ValueError(f"maxrej must be >= 0, got {maxrej}")
# Counted from 1, so rejstart <= 0 would silently disable the
# reference's unconditional ``iter == rejstart`` pass.
schedule.validate_iteration_setting("rejstart", rejstart, 1)
self.numrej = 0
self.good_idx: Optional[mx.array] = None
self.rho0 = rho0
self.minrho = minrho
self.maxrho = maxrho
# Working rho rate and its ceiling, as AMICATorchNG keeps them (the
# reference's rholrate/rholrate0, amica15.f90:1063-1068, 1806/1813):
# a decrease scales the working rate, a maxdecs ratchet the ceiling, and
# every A update resets the working rate to the ceiling.
self.rholrate0 = rholrate
self.rholrate = rholrate
self.rholrate_cap = rholrate
self.rholratefact = rholratefact
# Source-density family selection (issue #265, porting AMICATorchNG's
# issue #26 -- see ``AMICATorchNG.__init__`` for the identical block).
# Values match Fortran's per-source pdtype codes: 0 generalized Gaussian
# (the default, GG-mixture with adaptive rho), 2 Gaussian mixture, 3
# logistic (sech^2) mixture, 4 sub-Gaussian cosh+ (single component).
# pdftype=1 enables the extended-Infomax adaptive switcher (Fortran's
# do_choose_pdfs trigger), which flips each source between the
# super-Gaussian (code 1) and sub-Gaussian (code 4) cosh densities by
# kurtosis sign on the kurt_start/num_kurt/kurt_int schedule. Families 1
# and 4 are single-component (no alpha mixture).
if pdftype not in (0, 1, 2, 3, 4):
raise ValueError(f"pdftype must be one of 0,1,2,3,4; got {pdftype}")
self.pdftype = pdftype
# Fortran freezes the GG shape update for every non-GG family
# (amica15.f90: `if (pdftype /= 0) dorho = .false.`, lines 3704-3705).
self.dorho = pdftype == 0
# pdftype==1 is Fortran's adaptive trigger (amica15.f90:612).
self.do_choose_pdfs = pdftype == 1
self.kurt_start = kurt_start
self.num_kurt = num_kurt
self.kurt_int = kurt_int
# Families 1/4 (and the adaptive mode, which uses only codes 1 and 4) are
# single-component densities: Fortran's z0 references only mixture
# component j=1 and omits log(alpha). They are meaningful only with
# n_mix == 1.
if pdftype in (1, 4) and n_mix != 1:
raise ValueError(
f"pdftype={pdftype} is a single-component density (adaptive mode "
f"uses codes 1 and 4); it requires n_mix=1, got n_mix={n_mix}."
)
# Validate the adaptive-switch schedule up front (mirrors the
# share_comps checks below): kurt_int==0 would otherwise raise a bare
# ZeroDivisionError deep in fit(), and a negative kurt_int silently
# changes the schedule.
if self.do_choose_pdfs:
if kurt_int < 1:
raise ValueError(f"kurt_int must be >= 1, got {kurt_int}")
if kurt_start < 1:
raise ValueError(f"kurt_start must be >= 1, got {kurt_start}")
if num_kurt < 0:
raise ValueError(f"num_kurt must be >= 0, got {num_kurt}")
# Adaptive-switch counter (Fortran-1-indexed schedule check in fit());
# reset per fit in _initialize_parameters.
self.n_kurt_done = 0
self.invsigmin = invsigmin
self.invsigmax = invsigmax
self.doscaling = doscaling
self.scalestep = scalestep
if doscaling:
# A zero cadence divides by zero mid-fit; a fractional one fires on
# no meaningful schedule. The reference never reads scalestep.
schedule.validate_iteration_setting("scalestep", scalestep, 1)
# Component sharing (issue #263), same names/defaults/validation as
# AMICATorchNG (torch_impl/core.py). OFF by default and inert for
# n_models=1 (a model cannot share a component with itself), so the
# default trajectory is untouched.
self.share_comps = share_comps
self.share_start = share_start
self.share_iter = share_iter
self.comp_thresh = comp_thresh
# The A-freeze schedule reads share_start/share_iter whether or not
# share_comps is on (issue #345), so both are validated always, with
# AMICATorchNG's messages.
schedule.validate_share_start(share_start)
schedule.validate_share_iter(share_iter)
if share_comps:
if not 0.0 < comp_thresh <= 1.0:
raise ValueError(f"comp_thresh must be in (0, 1], got {comp_thresh}")
self.do_mean = do_mean
self.do_sphere = do_sphere
self.do_approx_sphere = do_approx_sphere
# Explicit PCA reduction (issue #323), validated by the policy shared
# with the PyTorch and NumPy backends (pamica/rank.py) so a bad value
# fails here, before any data is touched, exactly as it does there;
# then one log line for any part of it a fit will ignore.
validate_pca_reduction(pcakeep, pcadb)
log_ignored_pca_request(pcakeep, pcadb, do_sphere)
self.pcakeep = pcakeep
self.pcadb = pcadb
# Numerical-rank floors (issue #223); see pamica/rank.py and ADR 0004.
self.mineig = mineig
self.mineig_rel = mineig_rel
# Best-iterate safeguard (issue #51, epic #278 Phase 2/#288), same
# name/default/semantics as AMICATorchNG. See _fit_once for the
# per-iteration tracking and the end-of-fit restore decision.
self.keep_best = keep_best
self.seed = seed
# Best-of-N restarts (issue #198), a pamica extension: Fortran has no
# search over seeds. Resolved here so a bad configuration fails before
# any data is touched, and derived from the CONSTRUCTOR seed so that a
# second fit() on the same instance repeats the same seeds even though
# fit() leaves self.seed on the winning restart.
self._restart_seeds = restarts.resolve_seeds(n_restarts, restart_seeds, seed)
self.n_restarts = int(n_restarts)
self.restart_seeds = None if restart_seeds is None else list(restart_seeds)
# Per-restart records, set by fit(): index-aligned lists of the seed each
# restart ran from, the log-likelihood it returned (NaN for a degenerate
# restart) and why it stopped. A degenerate restart is excluded from
# selection but kept here -- it is a fact about that seed.
self.restart_seeds_: List[Optional[int]] = []
self.restart_lls_: List[float] = []
self.restart_stop_reasons_: List[Optional[str]] = []
self.iteration = 0
self.ll_history: list[float] = []
# Mutual Information Reduction (MIR) waypoint trajectory (issue
# #137, epic #278 Phase 3/#289), populated by fit() when mir_step >
# 0: (iteration, mir_nats, variance) tuples from the CURRENT
# (mid-fit) W/sphere. Like ll_history, this is a true trajectory
# that a keep_best restore does NOT rewrite -- the fit-end MIR is
# mir() on the returned parameters, not mir_history_[-1]. Not part
# of state_dict(): it's a diagnostic, not a fitted parameter. Not
# index-aligned with ll_history: the entry for iteration i is
# computed AFTER that iteration's _update_parameters, while
# ll_history[i] is the likelihood of the parameters BEFORE it, so
# the two describe states one update apart (issue #161). An
# iteration that ends the fit on a stop records no waypoint.
self.mir_history_: list[tuple[int, float, float]] = []
# Log-likelihood of the returned parameters, set by fit() to
# ll_history[-1], or to the best iterate's LL if the keep_best
# safeguard (issue #51) restores it -- see _fit_once, which also says
# when that is exact (a convergence stop) and when it trails the
# returned parameters by one update (max_iter, issue #339). Under
# share_comps, if a merge fires on the LAST fit iteration, the
# returned A/W/comp_list are already post-merge but final_ll_ still
# reports the pre-merge log-likelihood -- the merge runs after that
# iteration's LL is recorded, so its effect on the LL only shows up
# in the next iteration's E-step, which never runs. This matches the
# reference ordering (Fortran identify_shared_comps runs after the
# iteration's LL accumulation, amica15.f90:1856-1858) and
# AMICATorchNG's ordering, so it is documented behavior, not a bug
# (issue #269).
self.final_ll_: Optional[float] = None
self.stop_reason: Optional[str] = None
# Populated by fit()/_initialize_parameters().
self.A: Optional[mx.array] = None
self.W: Optional[mx.array] = None
self.mu: Optional[mx.array] = None
self.alpha: Optional[mx.array] = None
self.beta: Optional[mx.array] = None
self.rho: Optional[mx.array] = None
self.gm: Optional[mx.array] = None
self.c: Optional[mx.array] = None
self.comp_list: Optional[mx.array] = None
# Per-source density-family codes (issue #265), (n_channels, n_models)
# int32, allocated in _initialize_parameters. None only before the first
# fit -- once allocated it is always a full pdftype-filled tensor, even
# for pdftype=0 (the _pdtype_h None fast path is a SEPARATE, cheaper
# check on the scalar self.pdftype, not on this being None).
self.pdtype: Optional[mx.array] = None
self.mean: Optional[mx.array] = None
self.sphere: Optional[mx.array] = None
# float64 host copy of the sphere, kept because the sharing metric's
# back-map pinv(sphere) must be computed at the precision the sphere was
# built with, not from the float32 GPU cast (issue #263). Set alongside
# self.sphere in _preprocess, and _sphere_pinv is invalidated there too.
# _load_params (issue #287) is the other point that assigns a sphere --
# loading from a persisted float32 sphere rather than fitting one, so
# _sphere_np there is that sphere upcast to float64 rather than
# _preprocess's higher-precision original (see _load_params for the
# consequence) -- and it invalidates _sphere_pinv the same way.
self._sphere_np: Optional[np.ndarray] = None
self._sphere_pinv: Optional[np.ndarray] = None
# comp_used mask, CACHED here rather than derived per call; see the
# comp_used property.
self._comp_used_arr: Optional[mx.array] = None
self.sldet = 0.0
self._lgamma_table: Optional[mx.array] = (
None # (n_mix, n_comps): lgamma(1+1/rho)
)
self._logdet_W: Optional[mx.array] = (
None # scalar: log|det W|, refreshed per iter
)
# Weight-gradient norm (Fortran ndtmpsum), recomputed every iteration by
# _update_direction and read by fit()'s two grad-norm checks. Held as
# the unevaluated MLX scalar rather than a Python float (AMICATorchNG's
# eager ``_ndtmpsum`` float) so materializing it joins fit()'s
# per-iteration likelihood sync instead of adding one; the
# ``_ndtmpsum`` property below is the float view the checks and
# cross-backend tests read.
self._nd_arr: Optional[mx.array] = None
# Full-dataset per-sample/per-model log-likelihood (Fortran's LLt,
# issue #155), STASHED as the training E-step computes it rather than
# recomputed by a separate forward pass at write time (issue #157;
# epic #278 Phase 3, issue #289 -- port of AMICATorchNG's identical
# mechanism in ``AMICATorchNG.__init__``). ``_llt_logv``/``_llt_ll``
# are the live per-fit buffers (Fortran's ``modloglik``/``loglik``),
# zero-filled so a ``do_reject`` sample keeps Fortran's zero
# sentinel. ``fit`` converts them into the compact numpy
# ``_llt_lht``/``_llt_lt`` that :meth:`write_amica_output` consumes,
# then drops the mx buffers. Not fitted parameters (absent from
# ``state_dict()``/``_PARAM_ARRAYS``): a model restored via
# :meth:`from_state_dict` has none, so ``write_amica_output`` writes
# no LLt for it.
self._llt_logv: Optional[mx.array] = None
self._llt_ll: Optional[mx.array] = None
self._llt_lht: Optional[np.ndarray] = None
self._llt_lt: Optional[np.ndarray] = None
# Stop reasons that mark a fit as degenerate, the same set as
# AMICATorchNG's: a non-finite log-likelihood ("nan_ll"/"singular_ll"), a
# non-finite update direction ("nan_direction"), non-finite parameters after
# an update ("nan_params"). The last entry is only reachable under best-of-N
# restarts (issue #198): a restart whose fit raised rather than stopping,
# recorded as degenerate so a search in which every restart crashed still
# reports an unusable model.
_DEGENERATE_STOP_REASONS = (
"nan_ll",
"singular_ll",
"nan_direction",
"nan_params",
restarts.ERROR_STOP_REASON,
)
@property
def _ndtmpsum(self) -> Optional[float]:
"""Latest weight-gradient norm as a host float (AMICATorchNG's
``_ndtmpsum``), or None before the first step is built. Cheap after
fit()'s mx.eval has materialized it; forces evaluation otherwise, so a
direct ``_update_direction``/``_update_parameters`` call still reads the
current iteration's value."""
if self._nd_arr is None:
return None
return float(self._nd_arr.item())
# ------------------------------------------------------------------
# Preprocessing (host / numpy; mirrors AMICATorchNG._preprocess in float64)
# ------------------------------------------------------------------
def _preprocess(self, X: np.ndarray) -> mx.array:
"""Mean-removal + sphering in float64 on the host, then handed to MLX as
float32. Done in numpy (not MLX) because it reuses the exact float64
preprocessing AMICATorchNG already validates, so the sphere/sldet match
the PyTorch backend. (MLX's CPU-stream ``eigh`` is full float64 -- only
the GPU stream is unsupported -- so this is a code-sharing choice, not a
precision workaround.)"""
Xc = np.ascontiguousarray(X).astype(np.float64)
data_dim = Xc.shape[0]
if self.do_mean:
mean = Xc.mean(axis=1, keepdims=True)
Xc = Xc - mean
else:
mean = np.zeros((data_dim, 1))
if self.do_sphere:
# Population covariance (/N), matching Fortran's DSYRK scatter, not
# numpy's default sample covariance (/(N-1)) -- the same choice, and
# the reasoning for it, in ``AMICATorchNG._preprocess``.
cov = np.cov(Xc, bias=True)
evals, evecs = np.linalg.eigh(cov)
order = np.argsort(evals)[::-1]
evals = evals[order]
evecs = evecs[:, order]
# Numerical rank plus explicit PCA reduction, decided by the policy
# shared with the PyTorch and NumPy backends (pamica/rank.py,
# issues #223/#323) so the three cannot disagree. Fortran:
# numeigs = min(pcakeep, count(eigs > mineig)).
n_comp = numerical_rank(
evals,
mineig=self.mineig,
mineig_rel=self.mineig_rel,
pcakeep=self.pcakeep,
pcadb=self.pcadb,
)
evals = evals[:n_comp]
V = evecs[:, :n_comp]
inv_sqrt = np.diag(1.0 / np.sqrt(evals))
if n_comp < data_dim:
# Rank-reduced sphere (n_comp, data_dim), so the sphered data
# come out at the kept rank (Fortran nw = numeigs,
# amica15.f90:563).
w_pca = inv_sqrt @ V.T
if self.do_approx_sphere:
# Fortran's orthogonal polar-factor symmetrization of the
# reduced whitening (amica15.f90:501-508).
U_b, _, Vt_b = np.linalg.svd(evecs.T[:n_comp, :n_comp])
sphere = (Vt_b.T @ U_b.T) @ w_pca
else:
sphere = w_pca
self.n_channels = n_comp
self.n_comps = n_comp * self.n_models
elif self.do_approx_sphere:
# Symmetric ZCA sphere V diag(1/sqrt) V^T (Fortran default).
sphere = V @ inv_sqrt @ V.T
else:
sphere = inv_sqrt @ V.T
Xc = sphere @ Xc
sldet = float(-0.5 * np.log(evals).sum())
else:
sphere = np.eye(data_dim)
sldet = 0.0
self.mean = mx.array(mean.astype(np.float32))
self.sphere = mx.array(sphere.astype(np.float32))
# Keep the float64 sphere for _pinv_sphere and invalidate its cached
# pseudo-inverse here, so the cache can never describe a sphere other
# than the current one (AMICATorchNG._preprocess does the same).
self._sphere_np = sphere
self._sphere_pinv = None
self.sldet = sldet
return mx.array(Xc.astype(np.float32))
# ------------------------------------------------------------------
# Initialization (identical RNG draws to AMICATorchNG for cross-backend test)
# ------------------------------------------------------------------
def _initialize_parameters(self):
"""Initialize parameters with the *same* ``np.random.RandomState`` draw
order as ``AMICATorchNG._initialize_parameters``/AMICA_NumPy, so a shared seed
gives a bit-identical (float32-cast) starting point."""
rng = np.random.RandomState(self.seed)
n, m, ncomp, nmix = self.n_channels, self.n_models, self.n_comps, self.n_mix
# Per-model mixing blocks, drawn and normalized to unit-norm components
# as the reference does (issue #341, pamica.initialization), in float64
# before the float32 cast; comp_list maps each (source, model) to its
# component, a row of A (issue #334). Identical RNG draw order to
# AMICATorchNG.
A_np = initial_mixing(rng, n, m)
comp_list_np = np.zeros((n, m), dtype=np.int64)
for h in range(m):
comp_list_np[:, h] = np.arange(h * n, (h + 1) * n)
mu_np = np.zeros((nmix, ncomp))
for k in range(ncomp):
mu_np[:, k] = np.linspace(-1, 1, nmix)
mu_np[:, k] += 0.05 * (1 - 2 * rng.rand(nmix))
alpha_np = np.ones((nmix, ncomp)) / nmix
beta_np = np.ones((nmix, ncomp)) + 0.1 * (0.5 - rng.rand(nmix, ncomp))
rho_np = self.rho0 * np.ones((nmix, ncomp))
self.A = mx.array(A_np.astype(np.float32))
self.comp_list = mx.array(comp_list_np) # (n_channels, n_models) int
# Every component is referenced by the default block comp_list; reset
# here (not only in __init__) so a re-fit cannot inherit a merged mask.
self._comp_used_arr = mx.array(np.ones(ncomp, dtype=bool))
self.mu = mx.array(mu_np.astype(np.float32))
self.alpha = mx.array(alpha_np.astype(np.float32))
self.beta = mx.array(beta_np.astype(np.float32))
self.rho = mx.array(rho_np.astype(np.float32))
self.gm = mx.array((np.ones(m) / m).astype(np.float32))
self.c = mx.array(np.zeros((n, m), dtype=np.float32))
# Per-source density-family codes, Fortran `pdtype = pdftype`
# (amica15.f90:611; ``AMICATorchNG._initialize_parameters``). In adaptive mode
# (pdftype==1) every source starts as the super-Gaussian code (1),
# since self.pdftype IS 1 there -- no special-case fill needed.
self.pdtype = mx.array(np.full((n, m), self.pdftype, dtype=np.int32))
self.n_kurt_done = 0
# Reset the mutable optimization state to the pristine constructor values
# (lrate/lrate_cap, newtrate and rholrate/rholrate_cap are annealed or
# ratcheted down during fit, and n_newton_fallbacks counts one fit), so
# a re-fit starts fresh -- AMICATorchNG does the same in
# ``_initialize_parameters``/``_fit_once``.
self.lrate = self.lrate0
self.lrate_cap = self.lrate0
self.newtrate = self.newtrate0
self.rholrate = self.rholrate0
self.rholrate_cap = self.rholrate0
self.n_newton_fallbacks = 0
self.iteration = 0
self._refresh_lgamma_table()
self._update_unmixing_matrices()
def _refresh_lgamma_table(self):
"""Recompute ``lgamma(1+1/rho)`` host-side (MLX has no lgamma). Called at
init and after every rho update. Cheap: rho is ``(n_mix, n_comps)``.
``_log_pdf`` subtracts ``LOG2 + table``. At ``rho == 2`` the reference
takes its exact-Gaussian branch, whose normalizer is its own
single-precision literal (amica15.f90:1313, issue #344), so the entry
there is ``LOG_SQRT_PI - LOG2`` rather than ``lgamma(1.5)``: the same
normalizer the PyTorch and NumPy backends subtract, cast to float32. At
``rho == 1`` the Laplace branch's ``log(2)`` is exact and
``lgamma(2) == 0``, so the general form already matches it."""
rho_np = np.array(self.rho, dtype=np.float64)
table = gammaln(1.0 + 1.0 / rho_np)
table = np.where(rho_np == 2.0, LOG_SQRT_PI - LOG2, table)
self._lgamma_table = mx.array(table.astype(np.float32))
def _update_unmixing_matrices(self):
"""Per-model ``W_h = inv(A[comp_list[:, h], :])`` and the LL Jacobian
``log|det W_h|``, on the CPU stream (MLX linalg is CPU-only), hoisted to
once per iteration. ``W`` is ``(n_models, n, n)`` and ``_logdet_W`` is
``(n_models,)``. For n_models=1 this is ``inv(A)`` unchanged.
These build lazy graph nodes, so on a HEALTHY matrix the ``inv``/
``slogdet`` calls below do not themselves materialize -- the graph is
realized later where ``mx.eval`` runs (in ``fit``). A singular ``A``
is different: MLX 0.32's CPU-stream ``mx.linalg.inv`` does not raise a
catchable Python exception on one. LAPACK's LU failure aborts the
whole process (``libc++abi: ... [Inverse::eval_cpu] LU factorization
failed``), which no ``try``/``except`` around ``fit`` can catch
(issue #274). Near-singular-but-finite float32 matrices either invert
to large finite values or hit this same abort -- there is no route to
a catchable float32 overflow instead, which is why ``fit``'s
``nan_params`` guard treats a non-finite ``W``/``_logdet_W`` as
defense in depth rather than a reachable case on its own.
Consequently, each per-model matrix is condition-checked host-side,
eagerly, immediately before its ``inv`` call: see
``_INV_COND_THRESHOLD``. This eager host read (via ``np.array``)
forces the SAME materialization the CPU-stream ``inv``/``slogdet``
below would have forced anyway, so the guard adds no new cross-stream
handoff -- only a small ``np.linalg.cond`` on a matrix already on the
host. Measured on the bundled sample: negligible relative to
per-iteration time (see the issue #274 PR body for the number). A
condition number above the threshold raises ``RuntimeError`` naming
the model index, iteration, and value, in place of the uncatchable
abort.
A matrix with non-finite entries needs its own handling, verified by
isolated-subprocess reproduction: a matrix that is ONLY non-finite
(no other defect) never aborts -- ``inv`` propagates NaN/inf into
``W``, caught downstream by ``nan_params``. But a matrix that is
BOTH non-finite AND structurally singular elsewhere (an exact
duplicate column plus one unrelated NaN, confirmed to reach and
abort the process) is not covered by "only non-finite" reasoning --
skipping the check on any non-finite entry, as an earlier version of
this guard did, lets that combination through unguarded. The check
below therefore 0-fills non-finite entries (a no-op when already
finite) before computing the condition number, UNLESS every entry is
non-finite (the observed shape of a dead-model corruption -- a
zero-responsibility model dividing by ``dgm==0`` -- which carries no
structural signal to check and is left to flow to ``inv``/
``nan_params`` exactly as before).
Caveat, carried from ``_INV_COND_THRESHOLD``: no scalar condition
number, on the sanitized matrix or otherwise, can guarantee catching
every conceivable abort-capable matrix -- the empirically observed
LU-abort onset (cond~9e8 to beyond cond~5e10, matrix-dependent) sits
below the 1e12 threshold, so a believed-rare residual window remains
between "passes this check" and "would actually abort".
"""
assert self.A is not None and self.comp_list is not None
ws, logdets = [], []
for h in range(self.n_models):
A_h = self.A[self.comp_list[:, h], :]
a_h_np = np.array(A_h, dtype=np.float32, copy=False)
finite_mask = np.isfinite(a_h_np)
# A matrix with ZERO finite entries carries no signal to check --
# this is exactly the observed shape of a dead-model corruption
# (a zero-responsibility model dividing by dgm==0 propagates
# NaN/inf through the WHOLE per-model direction matrix, so all
# of that model's component rows go non-finite together, not
# just one entry -- confirmed on a real fitted 2-model dead-model
# state). Skipping here reproduces the pre-guard behavior exactly:
# mx.linalg.inv on a wholly non-finite A does not abort -- it
# propagates NaN/inf into W, which fit()'s existing nan_params
# guard already catches.
#
# Otherwise (fully finite, OR a MINORITY of entries non-finite),
# sanitize any non-finite entries to 0.0 before computing cond.
# This is a no-op when already fully finite. When partially
# non-finite, 0.0 is a neutral fill at the same natural scale as
# A's entries (near-unit-norm columns) -- unlike a huge/extreme
# sentinel, which was tried and rejected: it makes ANY non-finite
# entry look "infinitely far" from the rest of the matrix via pure
# scale mismatch, so a well-conditioned matrix with one stray NaN
# and a genuinely singular one both come back cond=inf, which
# cannot distinguish them. The neutral fill can (verified: a
# well-conditioned real matrix with one injected NaN reads
# cond~35 after 0-fill; the same matrix with an EXACT DUPLICATE
# column plus that same stray NaN -- the reviewer-reported killer
# combination, confirmed by isolated-subprocess reproduction to
# abort the process under the OLD skip-on-any-non-finite logic --
# reads cond~3.7e16, comfortably over the threshold). A raise here
# is unconditional on the underlying non-finite entries (the
# sanitized cond is a probe of the surrounding structure, not a
# claim about what the unknown entries "really" are), so it also
# still catches an exact-duplicate-column A with no non-finite
# entries at all, unchanged from before.
if np.any(finite_mask):
sentinel = np.where(finite_mask, a_h_np, np.float32(0.0)).astype(
np.float32
)
cond = float(np.linalg.cond(sentinel))
if not math.isfinite(cond) or cond > _INV_COND_THRESHOLD:
n_bad = int(a_h_np.size - finite_mask.sum())
nonfinite_note = (
f" ({n_bad} of {a_h_np.size} entries were already "
"non-finite and were 0-filled for this check.)"
if n_bad
else ""
)
# Set the degenerate marker BEFORE raising (PR #318
# review): this fires mid-_update_parameters, after A/mu/
# beta/rho/alpha/gm/c have already been reassigned to the
# new iterate but before self.W/_logdet_W are (the stack
# below never runs). The single-restart fit() path has no
# try/except, so this RuntimeError propagates straight to
# the caller with the instance left holding that
# inconsistent state -- stop_reason must already say so by
# the time it does, or every state_dict()/write_amica_
# output() degenerate check downstream (which key off
# stop_reason, not "did an exception fire once") would
# accept it. The multi-restart path's except block also
# sets this -- now redundant there, but harmless, and kept
# so that path does not depend on every raise site
# upstream doing this correctly.
self.stop_reason = restarts.ERROR_STOP_REASON
raise RuntimeError(
f"Singular unmixing matrix for model {h} at iteration "
f"{self.iteration}: cond(A[comp_list[:, {h}], :]) = "
f"{cond:.3e} exceeds the float32 threshold "
f"{_INV_COND_THRESHOLD:.1e} (MLX's CPU-stream inv "
"would otherwise abort the process instead of "
"raising; #274). Likely a component collapse -- "
"consider re-seeding or a lower lrate."
f"{nonfinite_note}"
)
wh = mx.linalg.inv(A_h, stream=_CPU)
ws.append(wh)
logdets.append(mx.linalg.slogdet(wh, stream=_CPU)[1])
self.W = mx.stack(ws, axis=0) # (n_models, n, n)
self._logdet_W = mx.stack(logdets) # (n_models,)
def _pdtype_h(self, h: int) -> Optional[mx.array]:
"""Per-source density-family codes for model ``h``, shaped for
broadcasting against ``(batch, n_channels, n_mix)`` arrays (AMICATorchNG
``_pdtype_h``), or ``None`` on the default
``pdftype=0`` (GG-only) fast path so the E-step stays bit-identical to
the pre-#265 implementation (policy 2).
"""
if self.pdftype == 0:
return None
assert self.pdtype is not None
return self.pdtype[:, h][None, :, None] # (1, n_channels, 1)
# ------------------------------------------------------------------
# E-step
# ------------------------------------------------------------------
def _forward(self, Xb: mx.array):
"""E-step forward pass for one block, per model (``AMICATorchNG._forward``).
``Xb`` is ``(n_channels, batch)``. Returns ``logV``
``(batch, n_models)`` and per-model lists ``(b, z, y, az_rho)``. For
n_models=1 (c=0, gm=1, comp_list=identity) this is numerically identical
to the single-model path. Each model's ``az_rho`` entry is ``None`` when
that model's family is non-GG (``_pdtype_h`` not None; policy 5) -- see
``_log_pdf``."""
assert (
self.comp_list is not None
and self.c is not None
and self.W is not None
and self.mu is not None
and self.beta is not None
and self.rho is not None
and self.alpha is not None
and self._lgamma_table is not None
and self.gm is not None
and self._logdet_W is not None
)
b_list, z_list, y_list, azrho_list, logv_cols = [], [], [], [], []
for h in range(self.n_models):
idx = self.comp_list[:, h]
b = (Xb - self.c[:, h][:, None]).T @ self.W[h] # (batch, n_channels)
mu_h = self.mu[:, idx].T[None] # (1, n_channels, n_mix)
beta_h = self.beta[:, idx].T[None]
rho_h = self.rho[:, idx].T[None]
alpha_h = self.alpha[:, idx].T[None]
lgamma_h = self._lgamma_table[:, idx].T[None]
y = beta_h * (b[..., None] - mu_h) # (batch, n_channels, n_mix)
log_pdf, az_rho = _log_pdf(y, rho_h, lgamma_h, self._pdtype_h(h))
z0 = mx.log(alpha_h) + mx.log(beta_h) + log_pdf
ll_i = mx.logsumexp(z0, axis=-1) # (batch, n_channels)
z = mx.softmax(z0, axis=-1)
logv_cols.append(
mx.log(self.gm[h]) + self._logdet_W[h] + self.sldet + ll_i.sum(axis=-1)
)
b_list.append(b)
z_list.append(z)
y_list.append(y)
azrho_list.append(az_rho)
logV = mx.stack(logv_cols, axis=1) # (batch, n_models)
return logV, b_list, z_list, y_list, azrho_list
def _get_block_updates(self, Xb: mx.array) -> dict:
"""Exact-EM sufficient statistics for one block
(``AMICATorchNG._get_block_updates``). Mixture stats are scattered into
their ``comp_list`` columns; ``dWtmp``/``dgm``/``dc_numer`` are
per-model. For n_models=1 (v==1, identity comp_list) this reproduces the
single-model accumulators exactly.
Under ``do_newton`` the three Newton curvature accumulators
(``dsigma2_numer``, ``dkappa_numer``, ``dlambda_numer``; see
:meth:`_finalize_newton_stats`) are emitted as well. They are gated on
``do_newton`` rather than always computed, matching AMICATorchNG: the key
is then either present in every block of a fit or absent from all of
them, so ``_accumulate_blocks``' generic key loop sums them with no
special case, and a natural-gradient fit does none of the work.
``drho_n`` (the rho digamma-update numerator) is likewise gated on
``self.dorho`` (issue #265, policy 5): rho is frozen for every non-GG
family (Fortran ``dorho=.false.``), so accumulating it there is dead
work AMICATorchNG still pays (it always accumulates and discards) --
this is a deliberate WORK-only divergence, not a numeric one. Gating
uniformly on ``self.dorho`` (fixed for the whole model, not per-block)
keeps the key consistently present-or-absent across every block of a
fit, exactly like the Newton keys above.
Also returns two per-sample (non-summable) entries consumed and
removed by :meth:`_accumulate_blocks`: ``logV`` (batch, n_models) and
``ll_samples`` (batch,) -- Fortran's ``modloglik``/``loglik`` columns
for this block, stashed for the LLt output (issue #157).
"""
logV, b_list, z_list, y_list, azrho_list = self._forward(Xb)
# Per-sample total log-likelihood (Fortran ``P``/``loglik``,
# amica15.f90:1402), kept as a vector rather than folded straight into
# the scalar sum so the LLt stash can reuse it (issue #157, epic #278
# Phase 3/#289, porting ``AMICATorchNG._get_block_updates``); ``block_ll``
# is the same summation as before, bit for bit.
block_ll_samples = mx.logsumexp(logV, axis=1)
block_ll = block_ll_samples.sum()
v = mx.softmax(logV, axis=1) # (batch, n_models) model responsibilities
nmix, ncomp = self.n_mix, self.n_comps
tiny = float(np.finfo(np.float32).tiny)
def zeros():
return mx.zeros((nmix, ncomp), dtype=mx.float32)
dalpha_n, dmu_n, dmu_d = zeros(), zeros(), zeros()
dbeta_n, dbeta_d = zeros(), zeros()
if self.dorho:
drho_n = zeros()
dgm_cols, dwtmp_mods, dc_cols = [], [], []
# Newton curvature, model-major like dWtmp: one entry per model, stacked
# below (the MLX convention -- MLX has no in-place slice assignment, so
# per-model lists + mx.stack replace torch's `dsigma2_numer[:, h] = ...`).
dsigma2_mods, dkappa_mods, dlambda_mods = [], [], []
assert (
self.comp_list is not None
and self.beta is not None
and self.rho is not None
)
for h in range(self.n_models):
idx = self.comp_list[:, h]
b, zr, y, az_rho = b_list[h], z_list[h], y_list[h], azrho_list[h]
v_h = v[:, h]
beta_h = self.beta[:, idx].T[None] # (1, n_channels, n_mix)
rho_h = self.rho[:, idx].T[None]
pdtype_h = self._pdtype_h(h)
fp = _score(y, rho_h, pdtype_h)
u = v_h[:, None, None] * zr # u = v*z, (batch, n_channels, n_mix)
ufp = u * fp
dgm_cols.append(v_h.sum())
dalpha_n = dalpha_n.at[:, idx].add(u.sum(0).T)
dmu_n = dmu_n.at[:, idx].add(ufp.sum(0).T)
# Phase A guard: float32 can round y to exactly 0 (fp(0)=0 => ufp=0),
# so ufp/y is 0/0=NaN; where y==0, 0/1 contributes 0 (issue #75).
# torch's safe_y substitution (``_get_block_updates``) is mirrored exactly
# for every family here, even though the true fp/y limit at y->0 is a
# finite nonzero constant for codes 2/3/1 (fp'(0): 1 Gaussian
# (fp=y), 0.5 logistic (fp=tanh(y/2)), 2 super-Gaussian
# (fp=y+tanh(y))) and 0 for code 4 (fp=y-tanh(y), whose Taylor
# expansion is O(y^3)) -- parity with AMICATorchNG is the spec, not
# the true limit; see the PR body for the measured per-family
# y==0 frequency.
safe_y = mx.where(y == 0, mx.ones_like(y), y)
dmu_d = dmu_d.at[:, idx].add((beta_h[0] * (ufp / safe_y).sum(0)).T)
dbeta_n = dbeta_n.at[:, idx].add(u.sum(0).T)
dbeta_d = dbeta_d.at[:, idx].add((ufp * y).sum(0).T)
if self.dorho:
logab = rho_h * mx.log(mx.maximum(mx.abs(y), tiny))
logab = mx.where(az_rho < EPSDBLE, mx.zeros_like(logab), logab)
drho_n = drho_n.at[:, idx].add((u * (az_rho * logab)).sum(0).T)
g = (beta_h * ufp).sum(-1) # (batch, n_channels)
dwtmp_mods.append(g.T @ b) # (n_channels, n_channels)
dc_cols.append(Xb @ v_h) # data-space bias numerator sum_t v_h*x
if self.do_newton:
# Newton curvature accumulators (Fortran amica15.f90:1439-1446,
# 1496-1513), in terms of the score fp -- not the density
# derivative dpdf -- and reusing this block's live E-step locals,
# so Newton adds no extra pass over the data.
dsigma2_mods.append((v_h[:, None] * b**2).sum(0)) # (n_ch,)
dkappa_mods.append(
((u * fp**2).sum(0) * beta_h[0] ** 2).T
) # (n_mix, n_ch)
dlambda_mods.append((u * (fp * y - 1.0) ** 2).sum(0).T) # (n_mix, n_ch)
updates = {
"dgm": mx.stack(dgm_cols), # (n_models,)
"dalpha_n": dalpha_n,
"dmu_n": dmu_n,
"dmu_d": dmu_d,
"dbeta_n": dbeta_n,
"dbeta_d": dbeta_d,
"dWtmp": mx.stack(dwtmp_mods, axis=0), # (n_models, n_ch, n_ch)
"dc_numer": mx.stack(dc_cols, axis=1), # (n_channels, n_models)
"ll": block_ll,
# Per-sample E-step outputs for the LLt stash (issue #157). NOT
# summable accumulators -- _accumulate_blocks pops them before
# folding the rest -- and cost nothing extra: both are already
# computed above for ``ll``/``v``.
"logV": logV,
"ll_samples": block_ll_samples,
}
if self.dorho:
updates["drho_n"] = drho_n
if self.do_newton:
updates["dsigma2_numer"] = mx.stack(dsigma2_mods) # (n_models, n_ch)
updates["dkappa_numer"] = mx.stack(dkappa_mods) # (n_models, n_mix, n_ch)
updates["dlambda_numer"] = mx.stack(dlambda_mods) # (n_models, n_mix, n_ch)
return updates
def _accumulate_blocks(self, X: mx.array, stash_llt: bool = False) -> dict:
"""Sum sufficient statistics over all blocks as one lazy graph (no
per-block ``mx.eval`` -- that over-syncs 2.6x).
Parameters
----------
X : mx.array
The (sphered) data to accumulate over, already restricted to the
good set under ``do_reject``.
stash_llt : bool, default=False
Scatter each block's per-sample ``logV``/``ll_samples`` into the
``_llt_logv``/``_llt_ll`` buffers as it goes, so the LLt output
never needs a second pass over the data (issue #157, epic #278
Phase 3/#289 -- port of ``AMICATorchNG._accumulate_blocks``). The scatter is
itself just another lazy MLX op (``self._llt_logv[rows] = logv``), so it
joins the same per-iteration graph ``fit`` already evaluates once -- no
extra sync is added here. Only the training loop sets this; the
``_tune_block_size`` probes leave the buffers untouched, so the
tuner still leaves no state behind. The per-sample values are
dropped from the returned dict either way -- they are per-block
quantities, not accumulators, and summing them would be
meaningless.
"""
n_samples = X.shape[1]
acc: Optional[dict] = None
for start in range(0, n_samples, self.block_size):
end = min(start + self.block_size, n_samples)
block_acc = self._get_block_updates(X[:, start:end])
logv = block_acc.pop("logV")
ll_samples = block_acc.pop("ll_samples")
if stash_llt:
assert self._llt_logv is not None and self._llt_ll is not None
# Map this block's rows back onto the full-dataset index.
# Under do_reject the caller passed X_t[:, good_idx], so block
# [start:end] of X is good_idx[start:end] of the dataset;
# otherwise the block index is the sample index.
rows = (
self.good_idx[start:end]
if self.good_idx is not None
else slice(start, end)
)
self._llt_logv[rows] = logv
self._llt_ll[rows] = ll_samples
if acc is None:
acc = block_acc
else:
for key in acc:
acc[key] = acc[key] + block_acc[key]
assert acc is not None
return acc
@staticmethod
def _available_memory_bytes() -> Optional[int]:
"""What the Apple GPU reports it can comfortably work with.
``max_recommended_working_set_size`` rather than the raw memory size:
MLX will happily allocate past the recommended set and start paging,
which shows up as a mysteriously slow candidate rather than a failure,
so the cap is applied against the number the driver actually
recommends. ``None`` (no cap) if MLX cannot report it.
Both keys report CAPACITY, not currently-free memory (unlike CUDA's
``mem_get_info``): neither subtracts what is already allocated, here or
by anything else sharing the unified memory. The cap is therefore an
upper bound on what the GPU could ever give, which is why it is only a
first filter -- catching the real ``[metal::malloc]`` failure is what
makes the search safe.
"""
try:
info = mx.device_info()
except (AttributeError, RuntimeError) as exc:
logger.debug(
"could not query MLX device memory (%s: %s); block-size "
"search runs without a memory cap",
type(exc).__name__,
exc,
)
return None
size = info.get("max_recommended_working_set_size") or info.get("memory_size")
return int(size) if size else None
def _tune_block_size(self, X: mx.array) -> None:
"""Set ``self.block_size`` to the fastest timed candidate (issue #232).
The probe is one ``_accumulate_blocks`` pass, evaluated in full: MLX
builds a lazy graph, so without ``mx.eval`` over every accumulator the
clock would measure graph *construction* and pick the block size that
builds fastest rather than the one that runs fastest. The pass only
reads model state and consumes no RNG, and ``block_size`` is restored
around every probe, so the fit that follows is bit-identical to one
started directly at the chosen size.
"""
saved = self.block_size
def probe(size: int) -> float:
self.block_size = size
try:
start = time.perf_counter()
acc = self._accumulate_blocks(X)
mx.eval(list(acc.values()))
return time.perf_counter() - start
finally:
# Never leave the model holding a candidate -- least of all one
# that just failed to allocate (issue #232).
self.block_size = saved
self.block_size = blocktune.search(
probe=probe,
fallback=saved,
blk_min=self.blk_min,
blk_max=self.blk_max,
blk_step=self.blk_step,
n_samples=int(X.shape[1]),
n_channels=self.n_channels,
n_mix=self.n_mix,
n_models=self.n_models,
# The MLX backend is float32 throughout (Apple GPUs have no
# float64); see the module docstring.
itemsize=4,
available_bytes=self._available_memory_bytes(),
log=logger,
)
# ------------------------------------------------------------------
# M-step
# ------------------------------------------------------------------
def _finalize_newton_stats(self, acc: dict):
"""Reduce the Newton block accumulators into ``(sigma2, lambda_, kappa)``
(``AMICATorchNG._finalize_newton_stats``; Fortran
amica15.f90:1666-1680).
The Fortran ``baralpha``/``dkappa_denom``/``dlambda_denom``
responsibility masses all cancel algebraically against the per-mixture
``dalpha`` weighting, leaving (with ``dgm = sum_t v_h`` the raw model
mass):
sigma2[h,i] = dsigma2_numer[h,i] / dgm[h]
kappa[h,i] = sum_j dkappa_numer[h,j,i] / dgm[h]
lambda[h,i] = sum_j (dlambda_numer[h,j,i]
+ dkappa_numer[h,j,i] * mu[j,comp(i,h)]^2) / dgm[h]
MUST be called with the PRE-update ``mu``: lambda folds ``mu^2`` in, and
Fortran does that during E-step accumulation, before the M-step moves mu
(see the call site in :meth:`_update_parameters`).
``dgm`` is ``(n_models,)`` and everything else is model-major, so the
model mass broadcasts on axis 0 (``dgm[:, None]``). Getting that axis
wrong is the NumPy backend's issue #267 crash -- there the layout is
model-MINOR, so the same reduction needs ``dgm[None, :]``.
Returns ``(sigma2, lambda_, kappa)``, each ``(n_models, n_channels)``.
"""
assert self.mu is not None and self.comp_list is not None
dgm = acc["dgm"][:, None] # (n_models, 1)
sigma2 = acc["dsigma2_numer"] / dgm
kappa = acc["dkappa_numer"].sum(axis=1) / dgm
# mu at each source's component: mu[j, comp_list[i,h]]. MLX 2-D advanced
# indexing matches NumPy's, giving (n_mix, n_ch, n_models); transpose to
# the model-major layout the accumulators use.
mu_at = self.mu[:, self.comp_list].transpose(2, 0, 1)
lambda_ = (acc["dlambda_numer"] + acc["dkappa_numer"] * mu_at**2).sum(
axis=1
) / dgm
return sigma2, lambda_, kappa
def _newton_direction(self, dA_h, sigma2_h, lambda_h, kappa_h):
"""Per-model Newton direction ``H`` from the natural gradient ``dA_h``
(``AMICATorchNG._newton_direction``).
Vectorized port of the per-source-pair 2x2 solve (Fortran
amica15.f90:1718-1741):
H[i,i] = dA_h[i,i] / lambda[i]
sk1 = sigma2[i]*kappa[k]; sk2 = sigma2[k]*kappa[i] (i != k)
H[i,k] = (sk1*dA_h[i,k] - dA_h[k,i]) / (sk1*sk2 - 1) if sk1*sk2 > 1
Closed form, so this needs no linear algebra and stays on the GPU stream
(unlike ``inv``/``slogdet``, which MLX runs CPU-only). The ``prod > 1.0``
test and the ``diagonal(dA_h)/lambda_h`` divide are deliberately raw --
no epsilon margin, no guard -- matching the PyTorch and NumPy backends;
a non-finite result is contained downstream by the ``nan_params`` abort
in :meth:`fit` and by the masked ``ndtmpsum`` reduction.
Returns ``(H, posdef)``. ``posdef`` is False if any off-diagonal pair
fails ``sk1*sk2 > 1`` (the positive-definiteness guard); the caller then
falls back to the natural gradient for this model. Reading it costs one
host sync of a scalar per Newton iteration per model -- accepted (it
selects the lrate ramp target and drives the fallback counter, neither of
which can stay on the device), alongside the M-step's existing
dead-model and rho-NaN scalar syncs.
"""
n = self.n_channels
sk1 = sigma2_h[:, None] * kappa_h[None, :] # [i,k] = sigma2[i]*kappa[k]
sk2 = sigma2_h[None, :] * kappa_h[:, None] # [i,k] = sigma2[k]*kappa[i]
prod = sk1 * sk2
valid = prod > 1.0
denom = mx.where(valid, prod - 1.0, mx.ones_like(prod))
h_off = (sk1 * dA_h - dA_h.T) / denom
H = mx.where(valid, h_off, mx.zeros_like(h_off))
# Diagonal overrides (uses lambda, not the off-diagonal formula).
diag = mx.diagonal(dA_h) / lambda_h
H = H - mx.diag(mx.diagonal(H)) + mx.diag(diag)
# Positive-definite iff every OFF-diagonal pair passed the guard. MLX has
# no boolean-mask indexing (torch's ``valid[offdiag].all()``), so force
# the diagonal True instead -- same reduction, no gather.
eye_bool = mx.eye(n, dtype=mx.bool_)
posdef = bool(mx.all(mx.logical_or(valid, eye_bool)).item())
return H, posdef
def _update_direction(self, acc: dict) -> _UpdateStep:
"""The natural-gradient or Newton step for ``A`` from this iteration's
sufficient statistics, and its norm, without changing any parameter
(``AMICATorchNG._update_direction``).
The reference computes both in ``accum_updates_and_likelihood``
(amica15.f90:1666-1761), with ``LL(iter)``, before the likelihood-
decrease response and the stopping checks read ``ndtmpsum`` and before
``update_params`` applies the step; :meth:`fit` calls this first and
hands the result to :meth:`_update_parameters` once the checks have
run. Sets ``self._nd_arr`` (lazily; read through ``_ndtmpsum``).
Everything here reads the parameters as the E-step saw them: the Newton
curvature folds in the pre-update ``mu``, and ``dAk`` weights the models
by the pre-update ``gm`` (the reference does not reassign ``gm`` until
``update_params``, :1788; issue #219). MLX arrays are immutable and
every parameter is only ever rebound, so nothing here needs a copy.
"""
assert (
self.mu is not None
and self.A is not None
and self.comp_list is not None
and self.gm is not None
and self._comp_used_arr is not None
)
tiny = float(np.finfo(np.float32).tiny)
# Finalize the Newton curvature with the PRE-update mu. Fortran folds the
# mu^2 term into lambda during E-step accumulation, before the M-step
# moves mu (amica15.f90:1666-1680); doing it here, before
# _update_parameters moves anything, is what reproduces that. Finalize
# after the mu update instead and lambda silently uses the updated mu: no
# error, no NaN, just a subtly wrong Hessian (the torch backend's issue
# #24 bug, pinned there and here by
# test_newton_finalize_uses_preupdate_mu).
newton_active = schedule.newton_active(
self.do_newton, self.iteration, self.newt_start
)
if newton_active:
sigma2, lambda_, kappa = self._finalize_newton_stats(acc)
# Natural-gradient A-update. A is stored as Fortran's A^T, one component
# per row (issue #334), so the update is a LEFT-multiply by the
# transposed direction (as in ``AMICATorchNG._update_parameters``, #24
# root cause). Each model's step is scattered into its component rows as
# a gm-weighted average (Fortran dAk/zeta) using the PRE-update gm (see
# the docstring): for the default disjoint
# comp_list every component has one contributor, so gm cancels and
# n_models=1 is byte-for-byte the old `A - lrate*(dA.T@A)`; a SHARED
# component (#263) takes Fortran's responsibility-weighted average of the
# models' steps for its one mixing vector, NOT a raw sum (a raw sum would
# over-step by the contributor count). A merged-away component needs no
# special case: nothing scatters into its row, so its zeta is 0 and its
# dAk is 0/tiny = 0, i.e. it takes no step. The rescale in
# _update_parameters does not touch it either: it rescales the rows of
# each model's block, and a merged-away row is in no model's block.
#
# The direction/dAk/gradient-norm computation below runs
# UNCONDITIONALLY, not gated on _a_frozen(): Fortran computes dAk and
# ndtmpsum every iteration in accum_updates_and_likelihood
# (amica15.f90:1749-1761), strictly before the separate, freeze-guarded
# update_A block (:1803) that steps A. Only the step itself -- and the
# lrate ramp and rho-rate reset Fortran nests inside that same guarded
# block -- are conditional (issue #207: the grad-norm stop must see the
# true gradient magnitude every iteration, not only when A moves).
# Newton only swaps out the per-model DIRECTION; the dAk/zeta scatter,
# the gradient norm and the freeze are untouched by it
# (``AMICATorchNG._update_parameters``). A model whose curvature fails the
# positive-definiteness guard falls back to its natural gradient for this
# iteration, and -- as in Fortran -- ANY model falling back also sends
# the lrate ramp to lrate_cap instead of newtrate.
eye = mx.eye(self.n_channels)
directions = []
no_newt = False
for h in range(self.n_models):
dA_h = -acc["dWtmp"][h] / acc["dgm"][h] + eye # I - <g b^T>/dgm
if newton_active:
H, posdef = self._newton_direction(
dA_h, sigma2[h], lambda_[h], kappa[h]
)
if posdef:
directions.append(H)
else:
no_newt = True
directions.append(dA_h) # fall back to natural gradient
else:
directions.append(dA_h)
dAk = mx.zeros_like(self.A)
zeta = mx.zeros((self.n_comps,), dtype=mx.float32)
for h in range(self.n_models):
idx = self.comp_list[:, h]
dAk = dAk.at[idx, :].add(self.gm[h] * (directions[h].T @ self.A[idx, :]))
zeta = zeta.at[idx].add(self.gm[h] + mx.zeros((self.n_channels,)))
dAk = dAk / mx.maximum(zeta, tiny)[:, None]
# Weight-gradient norm (Fortran ndtmpsum, amica15.f90:1760-1761):
# ``sqrt(sum(dAk**2, mask=comp_used) / (nw*count(comp_used)))``, built
# from the step direction BEFORE the lrate scaling and before the A step
# applies it, exactly as Fortran does in accum_updates_and_likelihood
# (:1749-1761) ahead of update_params' A step (:1803-1815). Read by
# fit()'s two grad-norm checks (AMICATorchNG._update_parameters computes
# the same quantity). One squared norm per component row, as the
# reference sums each component's column. The comp_used mask matters
# only once share_comps has merged components away: without sharing it
# is all-True, so the select returns nd unchanged and the count is
# n_comps, leaving this the plain RMS over dAk. Kept as a lazy scalar
# (no .item() here) so it rides fit()'s single per-iteration mx.eval.
#
# SELECT, not multiply-by-mask: 0*NaN is NaN, so a single non-finite
# row would poison the whole reduction, and a NaN ndtmpsum would
# disable BOTH grad-norm stops (NaN <= min_nd is False) if fit() did not
# stop on it ("nan_direction"); a masked-away row must not trigger that
# stop. mx.where drops the masked lanes structurally instead. Unreachable today (a merged-away row's
# dAk is exactly 0), but Phase 3's Newton direction feeds this same dAk.
used_f = self._comp_used_arr.astype(mx.float32)
nd = (dAk**2).sum(axis=1) # (n_comps,)
nd = mx.where(self._comp_used_arr, nd, mx.zeros_like(nd))
self._nd_arr = mx.sqrt(
nd.sum() / (self.n_channels * mx.maximum(used_f.sum(), 1.0))
)
return _UpdateStep(dAk, newton_active, no_newt)
def _update_parameters(
self, acc: dict, n_samples: int, step: Optional[_UpdateStep] = None
):
"""Exact-EM mixture updates + natural-gradient A-update, optionally
Newton-preconditioned (``AMICATorchNG._update_parameters``).
``step`` is the mixing-matrix step :meth:`_update_direction` computed
from the same ``acc``: :meth:`fit` passes the one its stopping checks
already read, so nothing is computed twice. A direct call may omit it,
and the step is then computed here first, from the parameters as they
stand."""
if step is None:
step = self._update_direction(acc)
# The step was built with the pre-update gm (_update_direction), so gm
# can be rebound now.
self.gm = acc["dgm"] / n_samples # (n_models,); == 1 for single model
tiny = float(np.finfo(np.float32).tiny)
# Per-model data-space bias c[i,h] = sum_t v_h*x / sum_t v_h (Fortran
# update_c, as in ``AMICATorchNG._update_parameters``). Skipped for n_models=1
# (v==1 => c is the zero data mean; the update would add a float-sum residual
# and break the #24 bit-exact single-model path). A dead model (dgm[h]==0) keeps
# its prior c rather than writing 0/0, and is surfaced (matching AMICATorchNG).
if self.n_models > 1:
dgm = acc["dgm"]
live = dgm > 0.0
new_c = acc["dc_numer"] / mx.maximum(dgm, tiny)[None, :]
self.c = mx.where(live[None, :], new_c, self.c)
if not bool(mx.all(live).item()):
logger.warning(
"Zero-responsibility model(s) at iter %d; kept their prior "
"bias c (dead-model guard).",
self.iteration,
)
# Component sharing (#263): a component merged away by
# _identify_shared_comps is no longer referenced by comp_list, so no
# sufficient statistic accumulates into its column (dalpha_n/dmu_d/
# dbeta_d == 0) and the divisions below would be 0/0 = NaN -- which the
# nan_params guard in fit() would (correctly) abort on. Update only USED
# columns and freeze the rest at their last finite value (Fortran carries
# NaN there behind its comp_used mask; keeping them finite matches
# AMICATorchNG and the NumPy backend, docs/guides/amica-differences.md
# row 8). With the default full comp_list every column is used, so
# ``used`` is all-True and every update below is bit-for-bit unchanged.
assert self._comp_used_arr is not None
used = self._comp_used_arr[None, :] # (1, n_comps)
self.alpha = mx.where(
used,
acc["dalpha_n"] / acc["dalpha_n"].sum(axis=0, keepdims=True),
self.alpha,
)
# The A branch, where the reference has it: after gm/alpha/c and before
# mu/sbeta/rho (amica15.f90:1803-1816; ``AMICATorchNG._update_parameters``).
# On an iteration the reference holds A
# (:func:`pamica.schedule.share_freeze`), everything inside it is skipped
# together: the Newton-fallback bookkeeping (so a discarded Newton
# direction cannot pollute the fallback counter), the lrate ramp, the
# reset of the working rho rate to its ceiling, and the step itself. The
# step was built by _update_direction from the parameters the E-step
# saw, so taking it before the mixture updates changes nothing they read.
if not self._a_frozen():
if step.newton_active and step.no_newt:
# Fortran prints "Hessian not positive definite, using natural
# gradient" (amica15.f90:1809-1811). Surface the same signal so
# an all-fallback run is visible without re-instrumenting.
self.n_newton_fallbacks += 1
logger.warning(
"Newton not positive definite at iter %d; using natural gradient.",
self.iteration,
)
# Learning-rate ramp: toward newtrate while Newton is active and
# stable, otherwise toward lrate_cap (Fortran amica15.f90:1803-1816),
# from this iteration's lrate, which a likelihood decrease has
# already halved (fit runs the response first, issue #339).
if step.newton_active and not step.no_newt:
self.lrate = min(
self.newtrate, self.lrate + min(1.0 / self.newt_ramp, self.lrate)
)
else:
self.lrate = min(
self.lrate_cap, self.lrate + min(1.0 / self.newt_ramp, self.lrate)
)
# rho moves at the ceiling (``rholrate = rholrate0``, :1806/:1813),
# so a decrease's scaling of the working rate reaches it only on a
# held iteration.
self.rholrate = self.rholrate_cap
self.A = self.A - self.lrate * step.dAk
self.mu = mx.where(used, self.mu + acc["dmu_n"] / acc["dmu_d"], self.mu)
self.beta = mx.where(
used,
mx.clip(
self.beta * mx.sqrt(acc["dbeta_n"] / acc["dbeta_d"]),
self.invsigmin,
self.invsigmax,
),
self.beta,
)
# GG shape update with the 1/psi(1+1/rho) digamma factor (Fortran
# :2013-2014); digamma is computed host-side (MLX has none). A NaN here
# (e.g. from an upstream mu/beta blow-up) is reset to rho0 and surfaced,
# matching ``AMICATorchNG._update_parameters``, so it does not silently
# poison the lgamma table and every subsequent E-step.
# Deliberate divergence from ``AMICATorchNG._update_parameters``, which also
# skips the update when rho is pinned to a boundary (all 1.0 or all 2.0):
# that early-exit needs a host sync on a (n_mix, n_comps) reduction over
# rho every iteration. This backend does make a few scalar host syncs per
# iteration -- the dead-model check above, the rho-NaN canary below, and
# the Newton posdef flag -- each because a Python branch genuinely
# depends on the value; the rho-boundary exit is not one of those, since
# mx.clip clamps straight back to the boundary and the results are
# identical either way. So it would buy nothing but a sync -- unrelated to
# (and not fixed by) issue #265: pdftype != 0 is now supported, and for
# those families this whole block is skipped by the outer ``self.dorho``
# gate below (rho is frozen, Fortran ``dorho=.false.``), which is the
# real reason the digamma work does not run there -- not an MLX
# limitation. See ``_get_block_updates`` for the matching ``drho_n``
# accumulator gate (a work-only divergence from AMICATorchNG, which
# always accumulates and discards it).
if self.dorho:
drho = acc["drho_n"] / mx.maximum(acc["dalpha_n"], 1e-8)
rho_np = np.array(self.rho, dtype=np.float64)
psi = mx.array(digamma(1.0 + 1.0 / rho_np).astype(np.float32))
new_rho = self.rho + self.rholrate * (1.0 - (self.rho / psi) * drho)
nan_mask = mx.isnan(new_rho)
if bool(mx.any(nan_mask).item()):
logger.warning(
"NaN in rho update at iter %d; resetting to rho0=%g.",
self.iteration,
self.rho0,
)
new_rho = mx.where(nan_mask, self.rho0, new_rho)
# ``used`` freezes a merged-away column's rho, as for mu/alpha/beta.
self.rho = mx.where(
used, mx.clip(new_rho, self.minrho, self.maxrho), self.rho
)
# The reference rescales every iteration (it parses ``scalestep`` but
# never reads it, amica15.f90:1843/3686); pamica keeps ``scalestep`` as
# an extension counted from 1 like the reference's other cadences
# (iterations s, 2s, ...; ``schedule.every``), so the default 1 is the
# reference.
if self.doscaling and schedule.every(self.iteration, self.scalestep):
self._rescale_components()
# rho is frozen for every non-GG family (self.dorho is False), so the
# table it feeds (used only on the GG _log_pdf path) cannot have
# changed; refreshing it every iteration there would be dead host-CPU
# work (issue #265, policy 5 -- gated like the drho_n accumulation and
# the digamma pull above). Still built unconditionally once at init
# (_initialize_parameters): _log_pdf's non-GG branches structurally
# need a valid lgamma_table for their (dead-code) GG-fallthrough term,
# even though its value is never selected there.
if self.dorho:
self._refresh_lgamma_table()
self._update_unmixing_matrices()
def _rescale_components(self) -> None:
"""Rescale every component to a unit-norm mixing vector (Fortran
``doscaling``, amica15.f90:1843-1851), an exact change of scale; the
same rule, norms, order and zero-norm guard as
``AMICATorchNG._rescale_components``.
Component ``k`` is row ``k`` of ``A`` (issue #334, ADR 0007), so that
row is divided by its norm while ``mu[:, k]``/``beta[:, k]`` are
multiplied/divided by it, leaving the log-likelihood unchanged (the
column rule issue #333 replaced was not a change of scale). A zero-norm
(or NaN-norm) row keeps its values: the scale is 1 there, so
``A``/``mu``/``beta`` are unchanged rather than ``mu`` being zeroed.
Each component is rescaled exactly once, including one that
``share_comps`` merged into several models, and a merged-away row
(in no model's block) keeps its scale of 1. The norms come from each
model's block ``A[comp_list[:, h], :]``, the same gathers as before
issue #334, and every attribute is rebound to a new array as the rest
of the M-step does, never mutated in place.
"""
assert (
self.A is not None
and self.mu is not None
and self.beta is not None
and self.comp_list is not None
)
scale = mx.ones((self.n_comps,), dtype=mx.float32)
for h in range(self.n_models):
idx = self.comp_list[:, h]
norm = mx.sqrt((self.A[idx, :] ** 2).sum(axis=1)) # (n_channels,)
# A shared row gets the same norm from every block that holds it.
scale[idx] = mx.where(norm > 0, norm, mx.ones_like(norm))
self.A = self.A / scale[:, None]
self.mu = self.mu * scale
self.beta = self.beta / scale
# ------------------------------------------------------------------
# Adaptive PDF switch (issue #265; AMICATorchNG's #26 port)
# ------------------------------------------------------------------
def _choose_pdfs(self, X: mx.array) -> None:
"""Extended-Infomax adaptive PDF switch (Fortran ``do_choose_pdfs``,
``AMICATorchNG._choose_pdfs``).
Re-estimates each source's kurtosis from the current model activations
and sets its density family to the super-Gaussian (code 1) or
sub-Gaussian (code 4) cosh density by kurtosis sign. The reference
binary declares this (``pdftype==1`` sets ``do_choose_pdfs``,
amica15.f90:612) but never runs the switch (``m2sum``/``m4sum`` are
never accumulated, :608-609), so there is no bit-exact oracle;
validated by real-data log-likelihood (must not decrease vs the fixed
GG default).
Mechanism difference from AMICATorchNG (MLX-motivated, not a decision
difference): the second moments are accumulated block-by-block on the
GPU in float32, but each block's small ``(n_channels,)``/scalar partial
sums are pulled to the host and accumulated in numpy float64, and the
kurtosis + validity guard + 1/4 decision run in numpy rather than on
the MLX graph. The kurt>0 sign test is a knife-edge decision with no
oracle, and ``m4`` loses float32 precision long before it would
overflow, so the accumulation itself is done at float64 host precision;
the host pulls happen on at most ``num_kurt`` iterations of a fit, so
this costs nothing on the hot per-block path. The decision semantics
(kurtosis formula, validity guard, super/sub-Gaussian mapping) are
identical to AMICATorchNG's -- see :meth:`_pdtype_from_kurtosis`.
"""
n_ch, n_models = self.n_channels, self.n_models
m2 = np.zeros((n_ch, n_models), dtype=np.float64)
m4 = np.zeros((n_ch, n_models), dtype=np.float64)
nsub = np.zeros(n_models, dtype=np.float64)
n_samples = X.shape[1]
for start in range(0, n_samples, self.block_size):
block = X[:, start : start + self.block_size]
logV, b_list, *_ = self._forward(block)
v = mx.softmax(logV, axis=1) # (batch, n_models)
for h in range(n_models):
b = b_list[h] # (batch, n_ch)
vh = v[:, h][:, None]
m2[:, h] += np.array((vh * b**2).sum(0), dtype=np.float64)
m4[:, h] += np.array((vh * b**4).sum(0), dtype=np.float64)
nsub[h] += float(v[:, h].sum().item())
# Kurtosis = E[b^4]/E[b^2]^2 - 3 = nsub * m4 / m2^2 - 3, per (source,
# model), in numpy float64 (policy 6).
tiny = np.finfo(np.float64).tiny
kurt = nsub[None, :] * m4 / np.maximum(m2**2, tiny) - 3.0
new_pdtype = self._pdtype_from_kurtosis(kurt, nsub)
# Silent-failure guard: the adaptive switcher only ever assigns codes 1
# (super-Gaussian) or 4 (sub-Gaussian), or keeps the prior value (which
# started as the ctor's all-1 fill and can therefore only ever BE 1 or
# 4 itself). Currently unreachable -- a bug elsewhere would have to
# write a different code first -- but an out-of-range code here would
# otherwise fall through the _score/_log_pdf mx.where chain to a stale
# GG density evaluated against a rho frozen by self.dorho, silently.
bad = set(np.unique(new_pdtype).tolist()) - {1, 4}
if bad:
# Mid-loop invariant raise (PR #318 review): _choose_pdfs runs
# from inside _fit_once's iteration body, after this iteration's
# _update_parameters already reassigned A/mu/beta/rho/alpha/gm/c
# -- so leaving this uncaught in the single-restart fit() path
# (no try/except there) would strand the instance holding a new
# iterate's parameters with a stale self.pdtype (the assignment
# below never runs) and stop_reason still "max_iter". Set the
# degenerate marker before raising so every downstream
# state_dict()/write_amica_output() refusal check catches it
# regardless of whether the caller catches this exception.
self.stop_reason = restarts.ERROR_STOP_REASON
raise RuntimeError(
f"_choose_pdfs produced pdtype code(s) outside {{1, 4}}: "
f"{sorted(bad)} (adaptive-switcher invariant violated)."
)
self.pdtype = mx.array(new_pdtype.astype(np.int32))
def _pdtype_from_kurtosis(self, kurt: np.ndarray, nsub: np.ndarray) -> np.ndarray:
"""Map per-source excess kurtosis to a density-family code (pure numpy;
``AMICATorchNG._pdtype_from_kurtosis``).
Super-Gaussian (positive kurtosis) -> code 1; sub-Gaussian -> code 4.
Only sources with a meaningful signal switch: a dead model
(``nsub[h]==0`` => ``kurt==-3.0``, finite) or a numerically blown-up
source (``kurt`` NaN, and ``NaN>0`` is False) would otherwise be
silently assigned code 4 with no diagnostic, so those keep their prior
``self.pdtype`` and are logged -- mirroring the dead-model / non-finite
guards in ``_update_parameters``. Split out from ``_choose_pdfs`` so
the decision (including the sub-Gaussian branch, which real EEG rarely
triggers) is unit-testable on constructed ``kurt``/``nsub`` arrays,
matching AMICATorchNG's ``test_pdtype_from_kurtosis_decision``.
"""
assert self.pdtype is not None
prior = np.array(self.pdtype, dtype=np.int64)
new_pdtype = np.where(kurt > 0.0, 1, 4)
valid = np.isfinite(kurt) & (nsub[None, :] > 0.0)
result = np.where(valid, new_pdtype, prior)
if not bool(valid.all()):
logger.warning(
"Non-finite or zero-mass kurtosis for %d source/model pair(s) "
"at iter %d; kept their prior pdtype (adaptive-switch guard).",
int((~valid).sum()),
self.iteration,
)
return result
# ------------------------------------------------------------------
# Best-iterate safeguard (issue #51, epic #278 Phase 2/#288; port of
# AMICATorchNG._snapshot_params/_restore_params)
# ------------------------------------------------------------------
def _snapshot_params(self) -> dict:
"""Snapshot the fitted state for the best-iterate safeguard (issue #51).
Copies each ``_PARAM_ARRAYS`` array (``mx.array(x)``, not an alias) so
a later in-fit rebind of ``self.A``/``self.W``/etc. does not roll the
snapshot forward -- MLX's M-step rebinds these attributes to new
arrays each iteration rather than mutating in place, but this clones
regardless so the invariant does not depend on that implementation
detail holding forever. Also captures the scalar ``n_kurt_done`` (the
adaptive-PDF switch counter that gates ``pdtype``) so a restored
model's switch count stays consistent with its rolled-back ``pdtype``
-- otherwise a switch applied after the peak iterate would leave the
two out of sync in a returned model (mirrors the torch silent-failure
fix this ports).
Also captures the LLt stash (issue #157, epic #278 Phase 3/#289; port
of ``AMICATorchNG._snapshot_params``).
``fit`` snapshots immediately after the E-step that produced both the
candidate ``best_ll`` and this iteration's ``_llt_logv``/``_llt_ll``,
so a restore rolls the on-disk LLt back to the E-step of the restored
iterate rather than leaving the last (discarded) iterate's per-sample
values behind -- what keeps the exported LLt the one belonging to the
exported parameters.
Also captures ``_logdet_W``/``_lgamma_table`` (PR #318 review): unlike
AMICATorchNG, which recomputes ``log|det W|``/``lgamma(1+1/rho)``
inline on every call, this backend hoists them to once-per-iteration
cached arrays (module docstring) -- and neither is in
``_PARAM_ARRAYS``, so without this they would silently stay at
whatever the LAST (discarded) iterate left them at after a restore,
corrupting every subsequent ``_forward``/``model_loglik``/
``model_probability``/``mir`` call on the "restored" model even
though ``W``/``rho`` themselves rolled back correctly. Captured
unconditionally (both always exist once :meth:`_initialize_parameters`
has run): at the point ``fit`` calls this, ``_logdet_W``/
``_lgamma_table`` still hold the values the just-computed E-step
(``best_ll``) actually used -- they are only refreshed at the END of
``_update_parameters``, i.e. AFTER this snapshot is taken -- so this
is precisely the state consistent with ``best_ll``, the same
E-step-not-yet-M-stepped point the rest of the snapshot captures.
Chosen over rebuilding them post-restore (:meth:`_load_params`'s
approach) because it keeps the snapshot a single self-contained
point-in-time capture with no follow-up step a future caller of
:meth:`_restore_params` could forget -- the generic ``setattr`` loop
there already restores whatever this dict holds.
"""
snap: dict = {
name: mx.array(getattr(self, name)) for name in self._PARAM_ARRAYS
}
snap["n_kurt_done"] = self.n_kurt_done
if self._logdet_W is not None:
snap["_logdet_W"] = mx.array(self._logdet_W)
if self._lgamma_table is not None:
snap["_lgamma_table"] = mx.array(self._lgamma_table)
# Present for every in-fit call (fit allocates the buffers before the
# loop and frees them only after the restore); absent only for a
# direct call on an already-returned model, where there is nothing to
# roll back.
if self._llt_logv is not None and self._llt_ll is not None:
snap["_llt_logv"] = mx.array(self._llt_logv)
snap["_llt_ll"] = mx.array(self._llt_ll)
return snap
def _restore_params(self, snapshot: dict) -> None:
"""Restore the state captured by :meth:`_snapshot_params`."""
for name, value in snapshot.items():
setattr(self, name, value)
# ------------------------------------------------------------------
# Outlier rejection (issue #123's AMICATorchNG mechanism; epic #278
# Phase 3/#289)
# ------------------------------------------------------------------
def _reject_outliers(self, ll_vec: np.ndarray) -> None:
"""Permanently drop samples whose (pre-update) log-likelihood is a low
outlier.
Fortran ``reject_data`` (amica15.f90:2380-2464): reject any
currently-good sample with ``loglik < mean - rejsig*std`` (population
std). The rejection is one-directional; ``good_idx`` only ever
shrinks, and the good-sample count drives the ``gm``/LL normalization
thereafter. ``ll_vec`` is the per-sample log-likelihood over the
current good set, in ``good_idx`` order, read from the LLt stash by
the caller (the design decision documented in :meth:`_fit_once`).
Runs on the host in numpy: the keep-mask COMPRESSION below
(``good[keep]``, dropping an arbitrary subset of entries) has no MLX
equivalent -- MLX has no boolean-mask gather, the same limitation
noted in :meth:`_newton_direction` -- so this mirrors
``_identify_shared_comps``' host-side control flow rather than
AMICATorchNG's/the NumPy backend's in-place tensor masking.
"""
assert self.good_idx is not None
good = np.array(self.good_idx)
mean = float(ll_vec.mean())
std = float(np.sqrt(max(float(np.mean(ll_vec**2) - mean**2), 0.0)))
keep = ll_vec >= (mean - self.rejsig * std)
if not bool(keep.any()):
# For finite log-likelihoods the max sample is always >= mean >=
# mean - rejsig*std (rejsig>0 is validated at construction), so it
# is always kept; the only way every sample is dropped is a
# non-finite per-sample LL (one NaN poisons mean/std, making every
# comparison False). Report that accurately instead of blaming
# rejsig (issue #127), which a user cannot fix by tuning rejsig.
n_bad = int(np.count_nonzero(~np.isfinite(ll_vec)))
if n_bad:
raise ValueError(
f"{n_bad} of {ll_vec.size} samples have a non-finite "
"log-likelihood; this indicates numerical instability "
"upstream (singular W / overflow), not a rejsig "
"miscalibration. Check for rank-deficient or "
"average-referenced data, or reduce the learning rate."
)
raise ValueError( # defensive: unreachable for finite LL, rejsig>0
f"Outlier rejection removed all {good.size} samples "
f"(rejsig={self.rejsig} too aggressive for this data)."
)
# Zero the LLt stash for the samples being dropped, exactly as
# Fortran's reject_data does (amica15.f90:2231-2234): they are never
# scored again, so without this they would keep the log-likelihood
# from the last iteration that still considered them good, and
# load_rej's ``sum(modloglik(:,i)) == 0`` sentinel would not see them
# as rejected.
dropped = good[~keep]
if self._llt_logv is not None and self._llt_ll is not None and dropped.size:
dropped_mx = mx.array(dropped)
self._llt_logv[dropped_mx] = mx.zeros(
(dropped.size, self.n_models), dtype=mx.float32
)
self._llt_ll[dropped_mx] = mx.zeros((dropped.size,), dtype=mx.float32)
self.good_idx = mx.array(good[keep])
self.numrej += 1
n_rejected = int(good.size - int(self.good_idx.size))
logger.info(
"Rejection %d at iter %d: dropped %d samples (%d good remaining).",
self.numrej,
self.iteration,
n_rejected,
int(self.good_idx.size),
)
# ------------------------------------------------------------------
# Component sharing (issue #263; AMICATorchNG's #60 port)
# ------------------------------------------------------------------
def _a_frozen(self) -> bool:
"""Whether the reference holds the A update (with its lrate ramp and
rho-rate reset) this iteration: once ``iter >= share_start``, every
iteration with ``mod(iter, share_iter) <= 5``, counted from 1
(amica15.f90:1803, :func:`pamica.schedule.share_freeze`;
``AMICATorchNG._a_frozen``).
The reference applies this whether or not ``share_comps`` is on and for
any number of models (issue #345), so this does too: with the defaults
``share_start = share_iter = 100``, every fit of 100 or more iterations
holds A on iterations 100-105, 200-205, and so on. The reference's own
scan never merges (see :meth:`_identify_shared_comps`), so this
unconditional schedule is the only freeze it ever shows. The
constructor requires ``share_iter >= 7``, since a shorter cycle would
hold A for good (:func:`pamica.schedule.validate_share_iter`).
"""
return schedule.share_freeze(self.iteration, self.share_start, self.share_iter)
def _component_sensor_maps(self) -> np.ndarray:
"""Every component's mixing vector in input-channel (sensor) space.
``pinv(sphere) @ A.T`` on the host in float64, shape ``(n_channels_in,
n_comps)``: column ``comp_list[i, h]`` is column ``i`` of
:meth:`get_sensor_mixing_matrix` for model ``h`` (issue #334). These
are the vectors the share metric compares.
"""
assert self.A is not None
return self._pinv_sphere() @ np.array(self.A, dtype=np.float64).T
def _identify_shared_comps(self) -> None:
"""Merge near-collinear components across models (Fortran
``identify_shared_comps``, amica15.f90:1916).
Two components (model ``h`` source ``i`` and model ``hh`` source ``ii``,
``h < hh``) are identified when the angle between their mixing vectors
(rows of ``A``, issue #334), measured in the original (de-sphered) data
space, is below the ``comp_thresh`` cutoff; on a match ``cj`` is folded
into ``ci``, so the two share one mixing vector and one density.
The decision itself is NOT reimplemented here: it runs
:func:`pamica.numpy_impl.utils.identify_shared_components` on host
float64 arrays, the same kernel the NumPy backend calls and the one whose
decisions are pinned equal to ``AMICATorchNG._identify_shared_comps``
(issue #258). A merge scan is a tiny greedy quadruple loop over
``n_models^2 * n_channels^2`` pairs, so nothing is gained by keeping it
on the GPU, and a third copy of the metric is exactly the cross-backend
divergence risk .rules/backend_parity.md forbids. ``A`` has just been
materialized by fit's per-iteration ``mx.eval``, so the host pull is
cheap.
No bit-exact oracle for the metric: the reference's ``Spinv2`` is
*declared* but never *allocated* in ``amica15.f90``, so the routine
there reads an unallocated array, every similarity comes out NaN and it
never merges (cf. the dead ``do_choose_pdfs`` switch, #26). The merged
state it produces is checked against the reference through
``load_comp_list`` on the float64 backends, and this backend is pinned
to AMICATorchNG.
"""
if self.n_models < 2:
return
assert self.A is not None and self.comp_list is not None
# _pinv_sphere raises on a non-finite sphere, so the metric below can
# only be garbage if A itself is (guarded per-pair inside the kernel).
atil = self._component_sensor_maps()
cl = np.array(self.comp_list)
new_cl, new_used = identify_shared_components(atil, cl, self.comp_thresh)
# Each fold removes exactly one component from the referenced set, so the
# drop in unique count IS the merge count (the kernel does not report it).
merged = int(np.unique(cl).size - np.unique(new_cl).size)
if merged:
self.comp_list = mx.array(new_cl)
self._comp_used_arr = mx.array(new_used)
logger.info(
"Component sharing (iter %d): %d merge(s), %d unique components.",
self.iteration,
merged,
int(np.unique(new_cl).size),
)
def _pinv_sphere(self) -> np.ndarray:
"""Cached ``pinv(sphere)``: the back-map from sphered to input-channel
space, in host float64.
This is the Fortran ``Spinv`` (amica15.f90:568-578), which the reference
also builds as a pseudo-inverse, ``Spinv(nx, numeigs)``, under rank/PCA
reduction. A pseudo-inverse rather than an inverse because reduction
leaves the sphere non-square (issue #223) and a square sphere fitted on
rank-deficient data is singular; for a full-rank square sphere the two
agree to ~1e-15. Computed from ``_sphere_np``, the float64 sphere
``_preprocess`` builds before the float32 GPU cast, so the merge metric
keeps the precision the PyTorch and NumPy backends use for it. Built on
first use and invalidated per fit in :meth:`_preprocess`, so it can never
describe a sphere other than the current one.
"""
if self._sphere_np is None:
raise RuntimeError(
"AMICAMLXNG._pinv_sphere() requires a preprocessed model; call "
"fit() first."
)
if self._sphere_pinv is None:
if not np.all(np.isfinite(self._sphere_np)):
# Only a degenerate fit (non-finite input data) gets here. Say
# so, rather than letting LAPACK report a confusing
# "ill-conditioned / repeated singular values" SVD failure.
# Mid-loop invariant raise (PR #318 review): _pinv_sphere is
# called from _identify_shared_comps, itself called from
# inside _fit_once's iteration body under share_comps -- so
# this can fire with the instance already holding this
# iterate's updated A/mu/etc, mid-fit, propagating uncaught
# through the single-restart fit() path. Set the degenerate
# marker before raising, same reasoning as the #274 guard
# above and _choose_pdfs's invariant.
self.stop_reason = restarts.ERROR_STOP_REASON
raise RuntimeError(
"The sphere holds non-finite values, so it has no "
"pseudo-inverse: the fit is degenerate. Check the input "
"data for NaN/inf."
)
self._sphere_pinv = np.linalg.pinv(self._sphere_np)
return self._sphere_pinv
@property
def comp_used(self) -> mx.array:
"""Boolean mask (n_comps,) of components still referenced by comp_list.
A component drops out of use when it is folded into another by
:meth:`_identify_shared_comps`; an unused component (its row of ``A``
and its density columns) receives no update and is never read by the
E-step.
CACHED (set all-True at init, rewritten by each merge) rather than
derived from ``comp_list`` on every read, which is how
``AMICATorchNG.comp_used`` does it. Same semantics and same source of
truth: the merge kernel already derives the mask host-side, from the
merged ``comp_list``, as part of the decision it returns -- so caching
that result costs nothing and keeps the mask fixed between merges, which
is exactly the lifetime the M-step needs.
"""
if self._comp_used_arr is None:
raise RuntimeError(
"AMICAMLXNG.comp_used requires a fitted model; call fit() first."
)
return self._comp_used_arr
def shared_components(self) -> list:
"""Components shared across models by ``share_comps`` (issue #263).
``share_comps`` folds near-collinear components of different models onto
one shared component (one row of ``A`` and one density), recorded as a
repeated index in ``comp_list``. Returns one group per shared
component: a list of ``(model_idx, source_idx)`` pairs that all
reference it, whose columns of :meth:`get_sensor_mixing_matrix` are
therefore identical. Empty when no component is shared across two or
more models (always for one model, and for a default multi-model fit
with ``share_comps`` off).
Note that a merge synchronizes only the parameters routed through
``comp_list`` (the mixing vector and ``mu``/``alpha``/``beta``/``rho``);
the
per-source density *family* code ``pdtype`` is a separate array and is
not synchronized (issue #265, matching ``AMICATorchNG.shared_components``),
so under the adaptive switcher (``pdftype=1``) a shared pair can still
report different :meth:`get_pdftype` codes.
Returns
-------
list of list of tuple(int, int)
"""
if self.comp_list is None:
raise RuntimeError(
"AMICAMLXNG.shared_components() requires a fitted model; call "
"fit() first."
)
self._check_usable("get the shared components")
cl = np.array(self.comp_list) # (n_channels, n_models)
groups = []
for col in np.unique(cl):
src, mdl = np.where(cl == col)
if np.unique(mdl).size >= 2:
groups.append([(int(h), int(i)) for i, h in zip(src, mdl)])
return groups
# ------------------------------------------------------------------
# Fit
# ------------------------------------------------------------------
# ------------------------------------------------------------------
# Best-of-N restarts (issue #198); mirrors AMICATorchNG's implementation
# ------------------------------------------------------------------
# Everything a fit writes, and therefore everything a restart snapshot must
# copy for the winning restart to be indistinguishable from a single fit
# from that seed. Together with the invariants below this must account for
# every ``self.x =`` in the fit path, which ``test_restart_policy.py``
# enforces by parsing this module -- so a field added to a fit-path method
# fails the suite until it is classified here.
_RESTART_STATE_ATTRS = (
# Fitted parameters and the derived per-iteration arrays ...
"A", "W", "c", "mu", "alpha", "beta", "rho", "gm", "comp_list", "pdtype",
"_comp_used_arr", "_lgamma_table", "_logdet_W", "_nd_arr",
# ... the schedule/counters a fit mutates ...
"iteration", "ll_history", "final_ll_", "stop_reason",
"n_newton_fallbacks", "n_kurt_done",
"lrate", "lrate_cap", "newtrate", "rholrate", "rholrate_cap",
# ... the LLt stash and its materialized arrays (issue #157, epic
# #278 Phase 3/#289) ...
"_llt_logv", "_llt_ll", "_llt_lht", "_llt_lt",
# ... outlier-rejection state (issue #123's AMICATorchNG mechanism,
# epic #278 Phase 3/#289) ...
"numrej", "good_idx",
# ... the MIR waypoint trajectory (issue #137, epic #278 Phase
# 3/#289) ...
"mir_history_",
# ... the tuned block size (do_opt_block re-times per restart) and the
# seed the winning restart ran from.
"block_size", "seed",
) # fmt: skip
# Written by the fit path but identical across the restarts of one fit()
# call, because they are functions of the data alone.
_RESTART_INVARIANT_ATTRS = (
"mean", "sphere", "_sphere_np", "sldet", "_sphere_pinv",
"n_channels", "n_comps",
) # fmt: skip
@staticmethod
def _copy_state_value(value):
""":func:`pamica.restarts.copy_state_value` plus the MLX case.
``mx.array(value)`` is a real copy (MLX arrays support item assignment,
so aliasing one into a snapshot would not be safe), and it preserves
dtype for the int/bool arrays here (``comp_list``, ``pdtype``,
``_comp_used_arr``).
"""
if isinstance(value, mx.array):
return mx.array(value)
return restarts.copy_state_value(value)
def _capture_restart_state(self) -> dict:
"""Independent copy of every attribute the fit path writes."""
return {
name: self._copy_state_value(getattr(self, name))
for name in self._RESTART_STATE_ATTRS
}
def _apply_restart_state(self, state: dict) -> None:
"""Restore the state captured by :meth:`_capture_restart_state`."""
for name, value in state.items():
setattr(self, name, value)
def fit(
self,
X: np.ndarray,
max_iter: int = 100,
verbose: bool = True,
mir_step: int = 0,
) -> "AMICAMLXNG":
"""Fit the model, running ``n_restarts`` fits and keeping the best.
``X`` is ``(n_channels, n_samples)``. With the default ``n_restarts=1``
this is exactly :meth:`_fit_once` -- the restart machinery draws
nothing, copies nothing and changes nothing, so the trajectory is
bit-identical to a pre-issue-#198 fit. With ``n_restarts > 1`` the model
is fit once per seed in ``restart_seeds`` (serially) and the returned
model holds the highest-``final_ll_`` non-degenerate restart's complete
state, exactly as a single fit from that seed would have left it.
Records (index-aligned, always populated): ``restart_seeds_``,
``restart_lls_`` (NaN where a restart ended degenerate) and
``restart_stop_reasons_``; the winner is named in one INFO log line. A
degenerate restart (``nan_ll``/``singular_ll``/``nan_direction``/
``nan_params``) is excluded from selection but recorded; if every restart is degenerate the
model is left holding the last one.
``mir_step``, as :meth:`_fit_once`, is passed through to every
restart unchanged.
"""
seeds = self._restart_seeds
if len(seeds) == 1:
# Single-restart path: seeds[0] IS self.seed unless the caller
# passed an explicit one-element restart_seeds, so nothing here
# perturbs the pre-#198 fit.
self.seed = seeds[0]
self._fit_once(X, max_iter=max_iter, verbose=verbose, mir_step=mir_step)
self.restart_seeds_ = list(seeds)
self.restart_lls_ = [
float("nan") if self.final_ll_ is None else float(self.final_ll_)
]
self.restart_stop_reasons_ = [self.stop_reason]
return self
lls: List[float] = []
degenerate: List[bool] = []
stop_reasons: List[Optional[str]] = []
states: dict = {}
for index, seed in enumerate(seeds):
self.seed = seed
try:
self._fit_once(X, max_iter=max_iter, verbose=verbose, mir_step=mir_step)
except RuntimeError as exc:
# An ill-conditioned A makes _update_unmixing_matrices raise
# (the issue #274 condition-number guard, which replaced MLX's
# process abort with a catchable RuntimeError). That guard keeps
# the process alive; this keeps the *search* alive, so one bad
# basin cannot discard the restarts that already succeeded.
# Mirrors AMICATorchNG._fit_restarts exactly, including catching
# only RuntimeError so a ValueError from _fit_once's argument
# checks still propagates.
self.stop_reason = restarts.ERROR_STOP_REASON
self.final_ll_ = float("nan")
logger.warning(
"%s", restarts.error_message(index, len(seeds), seed, exc)
)
ll = float("nan") if self.final_ll_ is None else float(self.final_ll_)
is_degenerate = self.stop_reason in self._DEGENERATE_STOP_REASONS
lls.append(ll)
degenerate.append(is_degenerate)
stop_reasons.append(self.stop_reason)
logger.info(
"%s",
restarts.progress_message(
index, len(seeds), seed, ll, self.stop_reason, is_degenerate
),
)
# Keep only the best state seen so far: one copy at a time.
if restarts.select_best(lls, degenerate) == index:
states = {index: self._capture_restart_state()}
winner = restarts.select_best(lls, degenerate)
if winner is None:
logger.warning(
"%s", restarts.all_degenerate_message(len(seeds), stop_reasons)
)
else:
logger.info(
"%s",
restarts.winner_message(winner, len(seeds), seeds[winner], lls[winner]),
)
if winner != len(seeds) - 1:
self._apply_restart_state(states[winner])
self.restart_seeds_ = list(seeds)
self.restart_lls_ = lls
self.restart_stop_reasons_ = stop_reasons
return self
def _fit_once(
self,
X: np.ndarray,
max_iter: int = 100,
verbose: bool = True,
mir_step: int = 0,
) -> "AMICAMLXNG":
"""Run one fit (one initialization, one EM loop) -- what :meth:`fit`
calls once per restart. ``X`` is ``(n_channels, n_samples)``.
Iteration order (issue #339), as in ``AMICATorchNG._fit_once``: each
iteration runs the E-step (its likelihood is appended to
``ll_history``, and the step for ``A`` and its norm are built), then
the likelihood-decrease response and the stopping checks, and only
then, unless a check fired, the parameter update with the rates just
set. ``iteration`` is the 0-based index of the last iteration whose
E-step ran. A convergence stop (``"min_dll"``, ``"grad_norm"``,
``"grad_norm_floor"``, ``"lrate_floor"``) takes no update on the
stopping iteration, so ``final_ll_`` (without a keep-best restore) is
exactly the log-likelihood of the returned parameters; a fit that runs
to ``max_iter`` takes the last iteration's update, so its
``final_ll_ == ll_history[-1]`` is the likelihood one update before the
returned parameters, as in the reference. A ``"nan_direction"`` stop
(a non-finite step or gradient norm, caught before any check reads it)
and a ``"nan_params"`` stop (non-finite parameters right after an
update) both record their iteration's (finite) likelihood; a
``"nan_ll"``/``"singular_ll"`` stop does not record the non-finite one.
All four are degenerate (``_DEGENERATE_STOP_REASONS``).
Under ``share_comps``, if a merge fires on the LAST iteration, the
returned ``A``/``W``/``comp_list`` are already post-merge but
``final_ll_`` still reports the pre-merge log-likelihood; see that
attribute's comment (issue #269).
LLt semantics (issue #157, epic #278 Phase 3/#289). The exported
``LLt`` (``_llt_lht``/``_llt_lt``, written by
:meth:`write_amica_output`) is the per-sample log-likelihood stashed
by the E-step that produced ``final_ll_``, never a separate post-fit
forward pass -- so after a fit that ran to ``max_iter`` it is one
M-step older than the returned ``W``/``A`` (Fortran's own convention;
see ``docs/guides/amica-differences.md``'s "one M-step older"
section). After a convergence stop, or a ``keep_best`` restore of an
earlier iterate, the stash and the returned parameters come from the
same point in the loop and there is no staleness at all.
``mir_step`` (issue #137, epic #278 Phase 3/#289), if > 0, computes
MIR from the current ``W``/``sphere`` every ``mir_step`` iterations
and appends it to ``mir_history_`` as ``(iteration, mir_nats,
variance)``. ``0`` (default) disables the waypoints and leaves fit
behavior byte-for-byte unchanged. ``mir_history_`` is a true
trajectory like ``ll_history``: a ``keep_best`` restore does not
rewrite it, so the fit-end MIR is ``self.mir(X)`` on the returned
parameters, not ``mir_history_[-1]``. Not index-aligned with
``ll_history``: entry ``i`` is computed after iteration ``i``'s
parameter update, while ``ll_history[i]`` is the likelihood of the
parameters before it, so the two are one update apart (issue #161);
an iteration that ends the fit on a stop takes no update and records
no waypoint.
Incompatible with PCA reduction, same as :meth:`mir` itself, and
gated exactly as ``AMICATorchNG._fit_once`` gates it (issue #323):
an explicit reduction request, ``pcakeep < n_channels`` or any
``pcadb`` while sphering (``pcakeep >= n_channels`` keeps every
dimension, and ``do_sphere=False`` never reduces, so neither is one),
raises ``ValueError`` up front, before :meth:`_preprocess`
runs, so a bad explicit config fails without paying for any
preprocessing. AUTOMATIC ``mineig``/``mineig_rel`` reduction is not a
request, and is caught the same way on both backends: downstream,
per-waypoint, inside :meth:`mir`'s own :meth:`_pca_reduced` guard,
whose ``ValueError`` is caught and logged here rather than propagated
(issue #283; issue #300 chose to keep that upfront-vs-downstream
split).
"""
if X.ndim != 2:
raise ValueError(f"X must be 2D (n_channels, n_samples), got {X.shape}")
if X.shape[0] != self._n_input_channels:
raise ValueError(
f"X has {X.shape[0]} channels, model expects {self._n_input_channels}"
)
if mir_step < 0:
raise ValueError(f"mir_step must be >= 0, got {mir_step}")
if max_iter < 1:
# PR #318 review: max_iter=0 used to run the loop zero times and
# complete "successfully" with stop_reason="max_iter" (not a
# _DEGENERATE_STOP_REASONS marker) and final_ll_=NaN -- an
# untrained model that every state_dict()/write_amica_output()
# degenerate-fit guard then accepted, since none of them check
# "did an E-step ever actually run", only "did stop_reason end
# up degenerate". Reject up front instead.
raise ValueError(f"max_iter must be >= 1, got {max_iter}")
# Same gate and message as AMICATorchNG._fit_once (issue #323). It
# sees only the explicit request; automatic reduction is caught per
# waypoint by mir()'s fitted-geometry guard below.
if mir_step > 0 and self._pca_reduction_requested(X.shape[0]):
raise ValueError(
"mir_step > 0 is incompatible with PCA reduction "
"(pcakeep/pcadb): the sphere is rank-deficient, so MIR's "
"log-Jacobian term is undefined. Rejected up front rather "
"than failing mid-fit at the first waypoint."
)
# Size every fit from the input geometry, as AMICATorchNG._fit_once
# does. _preprocess shrinks n_channels/n_comps to the kept rank, so
# without this a refit, or the second of n_restarts, would start from
# the previous fit's rank. A no-op for full-rank data.
self.n_channels = self._n_input_channels
self.n_comps = self.n_channels * self.n_models
X_t = self._preprocess(X)
n_total = X_t.shape[1]
self._initialize_parameters()
self.ll_history = []
self.mir_history_ = []
self.numrej = 0
self.stop_reason = "max_iter"
self.good_idx = mx.arange(n_total) if self.do_reject else None
# Block-size search (issue #232): after preprocessing and parameter
# initialization, before the first EM iteration, so it times the real
# data on the real device with the parameters the fit starts from. A
# no-op when off, and its probes leave no state behind, so a fit with
# the search off is byte-for-byte what it was before this existed.
if self.do_opt_block:
self._tune_block_size(X_t[:, self.good_idx] if self.do_reject else X_t)
# LLt buffers (issue #157), Fortran's permanently-allocated
# modloglik/loglik (amica15.f90:2617-2620). Zero-filled: a do_reject
# sample that is never scored again keeps the zero that Fortran's
# load_rej reads as the rejection sentinel. Re-allocated per fit so a
# refit on a different dataset cannot serve a stale array.
self._llt_logv = mx.zeros((n_total, self.n_models), dtype=mx.float32)
self._llt_ll = mx.zeros((n_total,), dtype=mx.float32)
self._llt_lht = None
self._llt_lt = None
numdecs = 0
# Consecutive-small-likelihood-gain counter for the min_dll stop (Fortran
# numincs, amica15.f90:1079-1089). Reset here so a refit starts clean.
numincs = 0
# MIR waypoint flood guard (PR #318 review): a ValueError from mir()
# (PCA reduction, or metrics.mir's own near-singular-unmixing check)
# is a per-fit-geometry condition, not a per-iteration one -- it does
# not spontaneously resolve, so leaving the schedule running would
# log the identical warning on every remaining waypoint of a long
# fit. A local (not self.<attr>): it only matters within this one
# _fit_once call, never needs to survive a restart snapshot or be
# inspected after fit() returns.
mir_waypoints_disabled = False
# Best-iterate safeguard (issue #51, epic #278 Phase 2/#288): track the
# highest-LL iterate so a late Newton-fallback overshoot cannot leave
# the returned model below a peak it already reached. Inactive under
# share_comps (a merge drops parameters, so pre- and post-merge LLs
# are not comparable AND the snapshot's comp_list would revert the
# merge -- the returned model would silently be unmerged; #60) and
# under do_reject (the good-sample set, and so the LL normalization,
# changes across iterations -- AMICATorchNG excludes it there for the
# same reason). Fit returns the last iterate when inactive, matching
# Fortran.
track_best = self.keep_best and not self.share_comps and not self.do_reject
best_ll = -math.inf
best_snapshot: Optional[dict] = None
if self.keep_best and (self.share_comps or self.do_reject):
# keep_best defaults on, so a user enabling sharing/rejection
# would otherwise silently lose the safeguard; surface it once.
# do_reject checked first, matching AMICATorchNG's precedence
# exactly (``AMICATorchNG._fit_once``) -- both can be true at once
# (see test_mlx_reject.py's genuine-merge-plus-reject test), and
# the two backends must report the same reason for the same
# configuration.
reason = "do_reject" if self.do_reject else "share_comps"
logger.warning(
"keep_best is inactive under %s: best-iterate selection by "
"LL is not well-defined (%s), so fit() returns the last "
"iterate.",
reason,
"the good-sample set / LL normalization changes across iterations"
if self.do_reject
else "a merge changes the parameter count and reverting to an "
"earlier snapshot would undo the merge",
)
rng = range(max_iter)
if verbose:
try:
from tqdm import tqdm
rng = tqdm(rng, desc="AMICA-MLX")
except ImportError:
pass
# One iteration follows the reference's main loop (amica15.f90:949-1142,
# issue #339; ``AMICATorchNG._fit_once``): the E-step with LL(iter), the
# step and its norm; the likelihood-decrease response and the stopping
# checks; an exit BEFORE any parameter moves if a check fired; otherwise
# the update with the rates the response just set, then the remaining
# hooks and rejection.
for it in rng:
self.iteration = it
X_use = X_t[:, self.good_idx] if self.do_reject else X_t
n_use = X_use.shape[1]
acc = self._accumulate_blocks(X_use, stash_llt=True)
# The step and its norm come from the same E-step and are read by
# the checks below (the reference builds both in
# accum_updates_and_likelihood, amica15.f90:1749-1761). Built lazily
# here so the sync below materializes them with the likelihood, one
# sync as before; a non-finite likelihood simply discards them.
step = self._update_direction(acc)
ll_arr = acc["ll"] / (n_use * self.n_channels)
# The stash scatter (_accumulate_blocks) is part of the same lazy
# graph as ll_arr, so materializing both here costs exactly the
# one sync this line already paid -- not a second one (issue #157).
mx.eval(ll_arr, self._llt_logv, self._llt_ll, step.dAk, self._nd_arr)
ll = float(ll_arr.item())
if not math.isfinite(ll):
self.stop_reason = "nan_ll" if math.isnan(ll) else "singular_ll"
logger.warning(
"Non-finite log-likelihood (%s) at iteration %d; stopping.", ll, it
)
break
# Best-iterate safeguard (issue #51): remember the parameters that
# produced this LL when it is the best seen, so a later overshoot
# does not leave the returned model below this peak. Nothing has
# moved them since the E-step, so the snapshot pairs them with ll.
if track_best and ll > best_ll:
best_ll = ll
best_snapshot = self._snapshot_params()
self.ll_history.append(ll)
# A non-finite step or norm would pass both gradient-norm checks
# below (NaN <= min_nd is False) and then be applied, so stop on it
# here, before any check reads it and before the update, with the
# parameters whose (finite) likelihood was just recorded (as
# AMICATorchNG does). Both were materialized with the likelihood.
nd_now = self._ndtmpsum
if not (
nd_now is not None
and math.isfinite(nd_now)
and bool(mx.all(mx.isfinite(step.dAk)).item())
):
self.stop_reason = "nan_direction"
logger.warning(
"Non-finite update direction (ndtmpsum %s) at iteration %d; "
"stopping before the update.",
nd_now,
it,
)
break
# Learning-rate control (Fortran amica15.f90:1051-1097): anneal the
# working lrate (and the working rho rate) on an LL decrease; ratchet
# the ceilings after maxdecs persistent decreases. All of it runs
# before this iteration's update, as in the reference, so the update
# below already uses the new rates.
#
# The working rho rate is scaled on every decrease, as the
# reference's is (:1063), but every update of A resets it to its
# ceiling before rho moves (_update_parameters, amica15.f90:1806/
# 1813), so the scaling only reaches rho on an iteration on which A
# is held. It is not a monotone decay (the issue #195 collapse):
# nothing but the maxdecs ratchet lowers the ceiling.
#
# have_prev mirrors Fortran's outer ``if (iter > 1)``
# (amica15.f90:1051), which wraps the decrease branch AND the two
# stops below, so none of the three can fire on the first iteration.
#
# PRECEDENCE NOTE (mirroring the same note in AMICATorchNG.fit): the
# three blocks are independent -- none is gated on ``leave`` already
# being True from an earlier block this same iteration, matching
# Fortran's own structure of independent ``leave = .true.``
# assignments with no declared precedence. Whichever block runs LAST
# and finds its condition true wins the reported stop_reason, so with
# this source order the standalone grad_norm block always has final
# say; under the shipped use_grad_norm=True default that makes the
# decrease branch's "grad_norm_floor" unreachable as the FINAL reason
# (its condition is strictly narrower). Deliberately not restructured
# into an explicit precedence, to keep this a direct port of
# amica15.f90:1051-1098.
have_prev = len(self.ll_history) > 1
leave = False
if have_prev and ll < self.ll_history[-2]:
if self.lrate <= self.minlrate:
logger.warning(
"lrate floor (%g) reached at iter %d; stopping.",
self.minlrate,
it,
)
self.stop_reason = "lrate_floor"
leave = True
elif self._ndtmpsum is not None and self._ndtmpsum <= self.min_nd:
# Fortran amica15.f90:1058's ``.or. (ndtmpsum .le. min_nd)``
# half of the decrease stop (issue #207 gap 3, #248 here):
# the same per-iteration value use_grad_norm reads below, so
# a run whose lrate oscillates instead of annealing still
# stops instead of burning the whole budget.
logger.warning(
"gradient-norm floor (%g) reached at iter %d on a "
"likelihood decrease; stopping.",
self.min_nd,
it,
)
self.stop_reason = "grad_norm_floor"
leave = True
else:
self.lrate *= self.lratefact
self.rholrate *= self.rholratefact
numdecs += 1
if numdecs >= self.maxdecs:
self.lrate_cap *= self.lratefact
if schedule.past_newton_start(it, self.newt_start):
self.rholrate_cap *= self.rholratefact
if self.do_newton and schedule.past_newton_start(
it, self.newt_start
):
# The Newton ceiling ratchets on the same maxdecs
# cadence as lrate_cap/rholrate_cap (Fortran
# amica15.f90:1056-1077), so a run that keeps
# overshooting at newtrate anneals instead of
# oscillating there.
self.newtrate *= self.lratefact
numdecs = 0
# Small-likelihood-increase stop (Fortran amica15.f90:1078-1090,
# use_min_dll/min_dll/maxincs). Independent of the decrease branch
# above: it runs every iteration once have_prev, including iterations
# where the LL just decreased (a decrease is always "less than" a
# positive min_dll, so it also increments numincs there, matching
# Fortran exactly). numincs resets to 0 on any gain >= min_dll; stops
# only after MORE than maxincs *consecutive* small gains.
if have_prev and self.use_min_dll:
if ll - self.ll_history[-2] < self.min_dll:
numincs += 1
if numincs > self.maxincs:
logger.warning(
"likelihood increasing by less than %g for more than "
"%d iterations; stopping at iter %d.",
self.min_dll,
self.maxincs,
it,
)
self.stop_reason = "min_dll"
leave = True
else:
numincs = 0
# Weight-gradient-norm stop (Fortran amica15.f90:1091-1097,
# use_grad_norm/min_nd). Also independent of the decrease branch:
# this is the unconditional every-iteration check, as opposed to the
# decrease-gated grad_norm_floor above.
if (
have_prev
and self.use_grad_norm
and self._ndtmpsum is not None
and self._ndtmpsum <= self.min_nd
):
logger.warning(
"norm of weight gradient <= %g at iter %d; stopping.",
self.min_nd,
it,
)
self.stop_reason = "grad_norm"
leave = True
# Switching Newton on changes the step direction, so the decrease
# counter accumulated during the natural-gradient phase no longer
# describes the schedule now running: Fortran clears it on the
# switch-on iteration (amica15.f90:1099-1102,
# ``AMICATorchNG._fit_once``).
if schedule.newton_switches_on(self.do_newton, it, self.newt_start):
numdecs = 0
# Stop before this iteration's update, as the reference does
# (amica15.f90:1111, ahead of update_params at :1122): the returned
# parameters are the ones whose LL was just recorded.
if leave:
break
# Whether rejection fires this iteration (Fortran schedule,
# amica15.f90:1136, with rejstart counted from 1 as the reference
# counts it; issue #335). Captured here, before _update_parameters, so
# the statistic is the PRE-update per-sample log-likelihood --
# matching Fortran's ordering (loglik is filled in
# get_updates_and_likelihood, before update_params runs).
#
# DESIGN DECISION (epic #278 Phase 3/#289): the statistic is read
# FROM the stash just written above (self._llt_ll indexed by
# good_idx), the NumPy backend's design (numpy_impl/core.py's
# _last_ll_samples), rather than AMICATorchNG's extra
# _sample_ll forward pass over the good set -- which open issue
# #298 records as the pass to eliminate there. The statistic is
# mathematically identical either way (both read the same
# per-sample logsumexp); this backend simply never pays for the
# second pass in the first place. Fancy-indexed straight out of
# the mx stash, so no host round-trip is needed to decide whether
# to reject.
will_reject = schedule.rejection_due(
self.do_reject, it, self.rejstart, self.rejint, self.numrej, self.maxrej
)
if will_reject:
assert self.good_idx is not None and self._llt_ll is not None
reject_ll = np.array(self._llt_ll[self.good_idx])
else:
reject_ll = None
self._update_parameters(acc, n_use, step)
# One eval per iteration bounds the lazy graph to a single iteration's
# worth of ops (the updated params feed the next accumulate). gm/c are
# included so their dependency chain is materialized each iteration too
# (c depends on the prior iteration's c), not left to grow unbounded.
# _nd_arr (the grad-norm stops' input) was materialized with the
# likelihood above, before the checks read it.
mx.eval(
self.A,
self.W,
self.mu,
self.alpha,
self.beta,
self.rho,
self.gm,
self.c,
self._logdet_W,
)
# Surface a corrupted M-step (component collapse / float32 overflow)
# at the iteration it happens. The ll check above only catches a
# corruption via the NEXT iteration's E-step, so a final-iteration
# blow-up would otherwise complete as max_iter with silently NaN
# parameters (the torch backend has state_dict as a backstop; the
# MLX backend does not, so guard in fit()). Params are already
# materialized by the mx.eval above, so this is a cheap read. The
# gradient norm is not among them: a non-finite one already stopped
# the fit before the update ("nan_direction", above). AMICATorchNG
# and the NumPy backend run the same check with the same message.
checked = {
"A": self.A,
"mu": self.mu,
"alpha": self.alpha,
"beta": self.beta,
"rho": self.rho,
"gm": self.gm,
"c": self.c,
# W and its log-determinant are DERIVED from A by
# mx.linalg.inv/slogdet, so a non-finite value can reach the
# caller while A itself is still finite -- and nothing else
# would catch it on the LAST iteration, where there is no next
# E-step to turn it into a nan_ll stop. The fit would then
# return stop_reason="max_iter" with a healthy-looking final_ll_
# (computed from the PREVIOUS iteration's W) and a silently
# non-finite unmixing matrix, which is precisely the outcome
# this guard exists to prevent. Verified by injection: with
# W/_logdet_W excluded, that state passes every other entry
# here.
#
# Defense in depth rather than a route known to be reachable:
# the obvious candidate, a near-singular A whose inverse
# overflows float32, is NOT reachable -- _update_unmixing_matrices
# now raises RuntimeError on such an A before calling inv at all
# (issue #274's condition-number guard), and before #274 it was
# unreachable for a different reason (MLX's LU aborted the whole
# process first, which this guard replaces with a catchable
# error). Cheap enough to keep regardless -- both are already
# materialized above.
"W": self.W,
"logdet_W": self._logdet_W,
}
params_finite = mx.array(True)
for value in checked.values():
params_finite = params_finite & mx.all(mx.isfinite(value))
if not bool(params_finite.item()):
# Name the offenders. Everything here is already materialized, so
# the per-tensor reads add no mid-graph sync (a check inside
# _update_parameters would sync the lazy graph mid-update).
bad = [
name
for name, value in checked.items()
if not bool(mx.all(mx.isfinite(value)).item())
]
logger.warning(
"Non-finite %s at iter %d (a mixture component likely "
"collapsed); stopping.",
", ".join(bad),
it,
)
self.stop_reason = "nan_params"
break
# Extended-Infomax adaptive PDF switch (Fortran do_choose_pdfs,
# ``AMICATorchNG._fit_once``). Runs on the
# kurt_start/num_kurt/kurt_int schedule using the just-updated W;
# the new per-source families take effect from the next E-step.
# kurt_start counts from 1, like every reference schedule. num_kurt=0
# disables switching (the family stays at its pdftype=1
# super-Gaussian init). Placed BEFORE the sharing hook below, matching
# AMICATorchNG's source order -- component sharing does not
# synchronize pdtype across merged columns (see
# shared_components()), so running the switch first means a
# just-merged pair still gets independently re-evaluated kurtosis
# this same iteration. This ordering is documentation, not a
# regression-tested contract: no test here pins the hooks' relative
# order (both are no-ops for most configurations, and share_comps
# x pdftype=1 has no bit-exact oracle either way to pin against), so
# a future accidental swap would not be caught by the suite.
if self.do_choose_pdfs and self.n_kurt_done < self.num_kurt:
if schedule.periodic_due(it, self.kurt_start, self.kurt_int):
self._choose_pdfs(X_use)
self.n_kurt_done += 1
# Component sharing (Fortran identify_shared_comps schedule,
# amica15.f90:1856): once per share_iter cycle from share_start,
# merge near-collinear components across models using the
# just-updated A. Fortran runs identify_shared_comps BEFORE
# get_unmixing_matrices (amica15.f90:1858,1863), so rebuild W from
# the merged comp_list -- otherwise the next E-step would read a
# stale W (pre-merge comp_list) while indexing the densities by the
# merged comp_list. No-op when share_comps is off or n_models == 1.
#
# This runs AFTER ``ll`` (this iteration's LL) was recorded above,
# so a merge on the final iteration lands in the returned
# A/W/comp_list but not in ``ll_history``/``final_ll_`` -- see
# final_ll_'s comment (issue #269).
if self.share_comps and schedule.periodic_due(
it, self.share_start, self.share_iter
):
self._identify_shared_comps()
self._update_unmixing_matrices()
# MIR waypoint (issue #137), following AMICATorchNG's idiom.
# Computed from the CURRENT W/sphere (just rebuilt above by
# _update_parameters / the share_comps block) against the raw,
# un-preprocessed X.
#
# A failed waypoint must never kill the fit. mir() raises on a
# near-singular unmixing or PCA reduction, and a near-singular W
# mid-fit is a transient the natural gradient can pass through.
# Warn and record NaN instead: the gap stays visible in
# mir_history_ rather than being silently absent.
#
# ValueError vs LinAlgError get different treatment (PR #318
# review): a ValueError (PCA reduction, or metrics.mir's own
# near-singular-unmixing check) reflects the fit's GEOMETRY --
# the sphere shape or the current unmixing's conditioning as a
# structural fact -- not a one-off numerical hiccup, so it will
# keep firing identically on every remaining scheduled waypoint
# of a long fit. Warn once, then stop scheduling waypoints for
# the rest of THIS fit (mir_history_ simply gets no more
# entries -- every one it would have gotten is the same NaN
# anyway, so nothing is lost). LinAlgError stays per-waypoint:
# it is the genuinely transient case the comment above already
# describes, which the natural gradient can pass through.
if mir_waypoints_disabled:
pass
elif mir_step > 0 and it % mir_step == 0:
try:
mir_nats, mir_var = self.mir(X)
except np.linalg.LinAlgError as exc:
logger.warning(
"MIR waypoint failed at iter %d (%s: %s); recording "
"NaN and continuing. The fit itself is unaffected.",
it,
type(exc).__name__,
exc,
)
mir_nats = mir_var = float("nan")
self.mir_history_.append((it, mir_nats, mir_var))
except ValueError as exc:
logger.warning(
"MIR waypoint failed at iter %d (%s: %s); this "
"condition will not resolve mid-fit, so MIR "
"waypoints are now disabled for the rest of this "
"fit (mir_history_ gets no further entries). The "
"fit itself is unaffected.",
it,
type(exc).__name__,
exc,
)
self.mir_history_.append((it, float("nan"), float("nan")))
mir_waypoints_disabled = True
else:
self.mir_history_.append((it, mir_nats, mir_var))
# Outlier rejection, after the parameter update (Fortran order,
# amica15.f90:1136-1140) but using the pre-update per-sample LL
# captured above.
if will_reject:
assert reject_ll is not None
self._reject_outliers(reject_ll)
# Log-likelihood of the parameters fit() returns. A degenerate stop
# leaves the model on the diverged parameters, whose LL is NOT the
# last finite ll_history value (the guard breaks before appending),
# so report NaN there rather than a stale healthy-looking number.
# Otherwise it is the last trajectory value, overwritten with the
# best iterate's LL below if the safeguard restores it.
if self.stop_reason in self._DEGENERATE_STOP_REASONS:
self.final_ll_ = float("nan")
else:
self.final_ll_ = self.ll_history[-1] if self.ll_history else float("nan")
# Restore the best iterate if the run ended materially below it
# (issue #51). Skipped for a degenerate stop -- state_dict() already
# refuses to persist any model whose stop_reason is degenerate, and
# salvaging a diverged run here would pre-empt that contract (issue
# #50). Also skipped when the final LL is within _KEEP_BEST_TOL of the
# best -- a monotone single-model run has final == best, so no
# restore fires and issue #24 parity stays byte-for-byte.
# ll_history is NEVER rewritten here: it stays the true per-iteration
# trajectory regardless of whether a restore fires below.
if (
track_best
and best_snapshot is not None
and self.stop_reason not in self._DEGENERATE_STOP_REASONS
and self.ll_history
and best_ll - self.ll_history[-1] > _KEEP_BEST_TOL
):
logger.info(
"Restoring best iterate (LL %.6f) over final LL %.6f "
"(issue #51 best-iterate safeguard).",
best_ll,
self.ll_history[-1],
)
self._restore_params(best_snapshot)
self.final_ll_ = best_ll
# LLt (Fortran's per-sample/per-model log-likelihood, issue #155):
# materialized from the stash the training E-step filled, with NO
# extra forward pass (issue #157). Runs strictly after the keep-best
# restore above, which rolls the stash back alongside the parameters
# (see _snapshot_params/_restore_params), so these arrays are always
# the E-step of the iterate fit() returns -- i.e. the very E-step
# whose total is ``final_ll_``:
# Lt.sum() / (n_good_samples * n_channels) == final_ll_
# where n_good_samples is the count that E-step ran over. This is
# Fortran's own convention (see the module docstring's staleness
# note) and holds for the reference binary's own output too.
# Converted to compact numpy here so the mx buffers can be freed; a
# refit reallocates them.
if self._llt_logv is not None and self._llt_ll is not None and self.ll_history:
self._llt_lht = np.array(self._llt_logv).T
self._llt_lt = np.array(self._llt_ll)
# Else: no iteration ever completed an E-step whose LL was recorded
# (max_iter=0, or a degenerate first iteration that broke before
# ll_history.append). The buffers hold nothing but zeros, which
# load_rej would misread as "every sample rejected", so leave
# _llt_lht/_llt_lt None and let write_amica_output omit the file
# with its existing warning rather than write a misleading one.
self._llt_logv = None
self._llt_ll = None
return self
def transform(self, X: np.ndarray, model_idx: int = 0) -> np.ndarray:
"""Apply the learned unmixing matrix to (new) data (issue #287, port of
``AMICATorchNG.transform``).
Sources are ``S = W[model_idx]^T @ (sphere @ (X - mean) - c[:,
model_idx])`` (issue #24 transpose convention, issue #27 per-model
center) -- the exact composition ``_forward`` uses to build its ``b``
activation, just laid out as ``(n_channels, n_samples)`` rather than
``_forward``'s ``(batch, n_channels)``: ``_forward`` computes ``b = (Xb
- c[:, h]).T @ W[h]``, and ``S = W[h].T @ (Xb - c[:, h])`` is exactly
``b.T`` by the transpose identity ``(W^T v)^T = v^T W``. CAUTION: MLX's
``W`` is ``(n_models, n, n)`` (``_update_unmixing_matrices`` stacks on
axis 0), NOT torch's ``(n, n, n_models)`` -- so the per-model slice
here is ``W[model_idx]``, not torch's ``W[:, :, model_idx]``.
Accepts any float ``np.ndarray``; computed in float32 (this backend's
only precision) and returned as a float32 ``np.ndarray``.
"""
if self.sphere is None or self.mean is None or self.W is None or self.c is None:
raise RuntimeError(
"AMICAMLXNG.transform() requires a fitted model; call fit() first."
)
self._check_model_idx(model_idx)
self._check_usable("transform")
self._check_input_shape(X)
X_arr = mx.array(np.ascontiguousarray(X).astype(np.float32))
X_t = self.sphere @ (X_arr - self.mean)
S = self.W[model_idx].T @ (X_t - self.c[:, model_idx : model_idx + 1])
return np.array(S)
# ------------------------------------------------------------------
# Fitted-parameter metadata (issue #265; AMICATorchNG's #142 port)
# ------------------------------------------------------------------
def _check_model_idx(self, model_idx: int) -> None:
"""Validate a model index against the fitted ``n_models`` (AMICATorchNG
``_check_model_idx``). Raises a clear ``ValueError``
(rejecting negatives, which MLX's negative indexing would otherwise turn
into a silent wrong-model result) instead of an opaque array error."""
if not isinstance(model_idx, (int, np.integer)):
raise TypeError(
f"model_idx must be an int, got {type(model_idx).__name__}."
)
if not (0 <= model_idx < self.n_models):
raise ValueError(
f"model_idx={model_idx} out of range for a {self.n_models}-model "
f"fit (valid: 0..{self.n_models - 1})."
)
def _nonfinite_params(self) -> list:
"""Names of :attr:`_PARAM_ARRAYS` currently holding a non-finite
value (matching the legacy NumPy backend's own
``_nonfinite_params``; issue #306 cross-backend review; port of
``AMICATorchNG._nonfinite_params``). Parameters not yet allocated
(``None``, e.g. a partially initialized instance) are skipped rather
than treated as bad.
Single-sync fast path: every parameter's ``isfinite().all()`` is
stacked into one small array and reduced with exactly one
``.item()`` MLX graph evaluation, instead of one evaluation per
parameter -- MLX is lazy, so each ``bool(mx.array)``/``.item()``
forces the whole pending graph to materialize on the accelerator,
making the naive per-parameter loop essentially the entire cost of a
small ``transform()`` call (measured ~2.2 ms across a 12-array
sweep, PR #329 review). The per-parameter breakdown -- one further,
small evaluation -- is only computed once that reduction is already
``False``.
"""
names = [name for name in self._PARAM_ARRAYS if getattr(self, name) is not None]
if not names:
return []
flags = mx.stack([mx.all(mx.isfinite(getattr(self, name))) for name in names])
if bool(mx.all(flags).item()):
return []
bad = (~flags).tolist()
return [name for name, is_bad in zip(names, bad) if is_bad]
def _check_usable(self, action: str) -> None:
"""Refuse to serve output from a degenerate fit (issue #306; port of
``AMICATorchNG._check_usable``).
Callers first check their own unfitted marker(s) and raise the
existing ``requires a fitted model`` ``RuntimeError`` (unchanged);
this assumes a fit has actually run and adds the two layers
:meth:`state_dict`/:meth:`write_amica_output` already use beyond
that: the ``stop_reason`` gate, then a defense-in-depth isfinite
sweep via :meth:`_nonfinite_params`. Mirrors the
:class:`~pamica.AMICA` wrapper's ``_check_usable`` (issue #50) for
callers using this backend directly.
"""
if self.stop_reason in self._DEGENERATE_STOP_REASONS:
raise RuntimeError(
f"Refusing to {action}: fit ended degenerate (stop_reason="
f"{self.stop_reason!r}), so the model holds non-finite "
f"parameters and would produce NaN output. Lower lrate, "
f"disable Newton, or check data conditioning, then refit."
)
nonfinite = self._nonfinite_params()
if nonfinite:
raise RuntimeError(
f"Refusing to {action}: parameters {nonfinite} hold "
f"non-finite values (stop_reason={self.stop_reason!r})."
)
def _check_input_shape(self, X: np.ndarray) -> None:
"""Validate a data array against the fitted input channel count,
mirroring :meth:`fit`'s own ``X`` validation (issue #306; port of
``AMICATorchNG._check_input_shape``)."""
if X.ndim != 2:
raise ValueError(f"X must be 2D (n_channels, n_samples), got {X.shape}")
if X.shape[0] != self.n_channels_in:
raise ValueError(
f"X has {X.shape[0]} channels, model expects {self.n_channels_in}"
)
def get_pdftype(self, model_idx: int = 0) -> np.ndarray:
"""Per-source density-family code for model ``model_idx`` (AMICATorchNG
``get_pdftype``).
One integer per source component (0-4; 0 generalized Gaussian, 1
super-Gaussian cosh, 2 Gaussian, 3 logistic, 4 sub-Gaussian cosh). All
sources share ``pdftype`` unless the adaptive switcher (``pdftype=1``)
moved them individually (issue #265). ``rho`` does not describe the
fitted density for codes 1-4 (it is frozen at ``rho0`` and only ever
meaningful for the generalized-Gaussian family, code 0).
Returns
-------
np.ndarray of int, shape (n_sources,)
"""
if self.pdtype is None:
raise RuntimeError(
"AMICAMLXNG.get_pdftype() requires a fitted model; call fit() first."
)
self._check_model_idx(model_idx)
self._check_usable("get the density family")
codes = np.array(self.pdtype[:, model_idx], dtype=np.int64)
# Silent-failure guard: an out-of-range stored code would otherwise
# fall through the _score/_log_pdf mx.where chain to a stale GG
# density (with rho frozen by self.dorho) with no diagnostic at all --
# surface it here instead, at the point a caller reads it.
bad = set(np.unique(codes).tolist()) - set(PDFTYPE_NAMES)
if bad:
raise RuntimeError(
f"AMICAMLXNG.get_pdftype(): stored pdtype has code(s) outside "
f"the valid set {sorted(PDFTYPE_NAMES)}: {sorted(bad)}."
)
return codes
def get_mixing_matrix(self, model_idx: int = 0) -> np.ndarray:
"""True mixing matrix of model ``model_idx``: the reference's
``A(:, comp_list(:, h))``, i.e. that model's component rows of the
stored ``A`` transposed (issue #24 convention; issue #334 layout;
issue #287 port of ``AMICATorchNG.get_mixing_matrix``)."""
if self.A is None or self.comp_list is None:
raise RuntimeError(
"AMICAMLXNG.get_mixing_matrix() requires a fitted model; call "
"fit() first."
)
self._check_model_idx(model_idx)
self._check_usable("get the mixing matrix")
return np.array(self.A[self.comp_list[:, model_idx], :].T)
@property
def n_channels_in(self) -> int:
"""Input channel count, i.e. the width of the sphere (issue #287 port
of ``AMICATorchNG.n_channels_in``).
Differs from ``n_channels`` only when rank reduction shrank the model
to the detected numerical rank (issue #223); equal to it for
full-rank data and before :meth:`fit`/:meth:`from_state_dict`.
Read off the sphere whenever one exists, so it cannot drift from the
sphere it describes, including on a reloaded rank-reduced model,
whose ``sphere`` width is exactly this value (see
:meth:`_load_params`'s shape guard). Before the first fit it is the
constructor's channel count, the width :meth:`fit` accepts.
"""
if self.sphere is None:
return self._n_input_channels
return int(self.sphere.shape[1])
def get_sensor_mixing_matrix(self, model_idx: int = 0) -> np.ndarray:
"""Mixing matrix mapped back to input-channel space (issue #287 port of
``AMICATorchNG.get_sensor_mixing_matrix``): ``pinv(sphere) @ A``, via
:meth:`_pinv_sphere` -- the only correct back-map when rank reduction has left
the sphere non-square (issue #223).
"""
if self.sphere is None:
raise RuntimeError(
"AMICAMLXNG.get_sensor_mixing_matrix() requires a fitted "
"model; call fit() first."
)
if self.A is None or self.comp_list is None:
raise RuntimeError(
"AMICAMLXNG.get_sensor_mixing_matrix() requires a fitted "
"model; call fit() first."
)
self._check_model_idx(model_idx)
self._check_usable("get the sensor mixing matrix")
A = np.array(self.A[self.comp_list[:, model_idx], :].T, dtype=np.float64)
return self._pinv_sphere() @ A
def get_unmixing_matrix(self, model_idx: int = 0) -> np.ndarray:
"""True unmixing matrix ``W_fort`` = (stored W)^T (issue #24
convention; issue #287 port of ``AMICATorchNG.get_unmixing_matrix``). MLX's
``W`` is model-major (``(n_models, n, n)``), so the per-model slice is
``W[model_idx]`` rather than torch's ``W[:, :, model_idx]``."""
if self.W is None:
raise RuntimeError(
"AMICAMLXNG.get_unmixing_matrix() requires a fitted model; "
"call fit() first."
)
self._check_model_idx(model_idx)
self._check_usable("get the unmixing matrix")
return np.array(self.W[model_idx].T)
# ------------------------------------------------------------------
# Preprocessing accessors (issue #313). Same names, shapes and float64
# return type as AMICATorchNG's, so a consumer that composes the transform
# itself (the MNE export, pamica.mne_compat) reads every backend alike.
# ------------------------------------------------------------------
def get_sphere(self) -> np.ndarray:
"""Fitted sphering matrix, shape ``(n_channels, n_channels_in)``
(port of ``AMICATorchNG.get_sphere``).
Square for a full-rank fit and ``(n_kept, n_channels_in)`` after rank
reduction (issue #223). Read from ``_sphere_np``, the float64 host
copy: after :meth:`fit` it is the float64 sphere :meth:`_preprocess`
computed (the GPU computes with its float32 cast, which agrees to
float32 rounding), and after :meth:`load` it is the persisted float32
sphere upcast (see :meth:`_load_params`). Returned as an independent
float64 copy.
"""
if self.sphere is None or self._sphere_np is None:
raise RuntimeError(
"AMICAMLXNG.get_sphere() requires a fitted model; call fit() first."
)
self._check_usable("get the sphere")
return np.array(self._sphere_np, dtype=np.float64)
def get_mean(self) -> np.ndarray:
"""Per-channel mean removed before sphering, shape ``(n_channels_in,)``
(port of ``AMICATorchNG.get_mean``).
All zeros for a ``do_mean=False`` fit. The stored mean is float32
(this backend's only precision); it is returned as an independent
float64 copy of those float32 values.
"""
if self.mean is None:
raise RuntimeError(
"AMICAMLXNG.get_mean() requires a fitted model; call fit() first."
)
self._check_usable("get the mean")
return np.array(self.mean, dtype=np.float64).ravel()
def get_model_center(self, model_idx: int = 0) -> np.ndarray:
"""Model ``model_idx``'s center ``c`` in sphered space, shape
``(n_channels,)`` (port of ``AMICATorchNG.get_model_center``).
The per-model offset :meth:`transform` subtracts after sphering
(issue #27). Identically zero for a single-model fit, since the ``c`` update
is gated to ``n_models > 1``. Returned as an independent float64 copy
of the stored float32 values.
"""
if self.c is None:
raise RuntimeError(
"AMICAMLXNG.get_model_center() requires a fitted model; call "
"fit() first."
)
self._check_model_idx(model_idx)
self._check_usable("get the model center")
return np.array(self.c[:, int(model_idx)], dtype=np.float64)
def get_rho(self, model_idx: int = 0) -> np.ndarray:
"""Generalized-Gaussian shape parameter ``rho`` for model
``model_idx`` (issue #287 port of ``AMICATorchNG.get_rho``; issue #142).
One value per (mixture component, source): ``rho == 2`` is Gaussian-
shaped, ``rho == 1`` Laplacian, ``rho < 1`` heavier-tailed. Only the
generalized-Gaussian family (``pdftype=0``) updates ``rho``; for every
non-zero code (1-4) it stays frozen at ``rho0`` and does not describe
the fitted density (see :meth:`get_pdftype`).
Returns
-------
np.ndarray of float, shape (n_mix, n_sources)
"""
if self.rho is None or self.comp_list is None:
raise RuntimeError(
"AMICAMLXNG.get_rho() requires a fitted model; call fit() first."
)
self._check_model_idx(model_idx)
# Folded into the shared guard (issue #306): a degenerate multi-model
# fit can leave one model's rho non-finite without the aggregate LL
# tripping a _DEGENERATE_STOP_REASONS marker, which _check_usable's
# defense-in-depth isfinite sweep over _PARAM_ARRAYS (rho included)
# still catches. Refuse rather than return a silent NaN.
self._check_usable("get rho")
idx = self.comp_list[:, model_idx]
return np.array(self.rho[:, idx])
# ------------------------------------------------------------------
# EEGLAB drop-in output (issue #92; epic #278 polish port of
# AMICATorchNG.variance_order)
# ------------------------------------------------------------------
def variance_order(
self, model_idx: int = 0, return_svar: bool = False
) -> np.ndarray | tuple:
"""EEGLAB back-projected-variance component order (IC1 = highest
variance) (port of ``AMICATorchNG.variance_order``).
Returns the source indices sorted by descending back-projected
variance, the ordering EEGLAB's ``loadmodout15.m`` applies on load
(so ``order[0]`` is IC1). The de-sphered sensor-space mixing column
``a_i = pinv(W S)[:, i]`` contributes ``||a_i||^2 * sum_k alpha_ki
(mu_ki^2 + r_ki / sbeta_ki^2)`` with ``r_ki = gamma(3/rho_ki)/
gamma(1/rho_ki)`` (the source's mixture variance), matching
``loadmodout15`` exactly. Non-mutating: the stored parameters keep
their fit order; this only reports the display order.
The parameters are pulled off the GPU with ``np.array(...)`` and the
ordering arithmetic (the gamma ratio, the sphere pseudo-inverse) runs
host-side in float64 via NumPy/SciPy -- matching what
``AMICATorchNG.variance_order`` computes in at its default
``dtype=torch.float64`` -- so the only float32 step is
the fitted parameters themselves, not how the order is computed from
them.
Parameters
----------
model_idx : int, default=0
Which model's components to order.
return_svar : bool, default=False
If True, also return the per-source variance sorted to ``order``.
Returns
-------
order : np.ndarray of int, shape (n_sources,)
Source indices, highest back-projected variance first.
svar : np.ndarray, optional
Present only when ``return_svar``; the sorted variances.
"""
if (
self.comp_list is None
or self.alpha is None
or self.mu is None
or self.beta is None
or self.rho is None
or self.W is None
or self.sphere is None
):
raise RuntimeError(
"AMICAMLXNG.variance_order() requires a fitted model; call fit() first."
)
self._check_model_idx(model_idx)
self._check_usable("compute the variance order")
cl = self.comp_list[:, model_idx]
alpha = np.array(self.alpha[:, cl], dtype=np.float64)
mu = np.array(self.mu[:, cl], dtype=np.float64)
sbeta = np.array(self.beta[:, cl], dtype=np.float64)
rho = np.array(self.rho[:, cl], dtype=np.float64)
# source mixture variance (sum over the mixture components); unused
# mixtures carry alpha == 0 and drop out, matching loadmodout15.
ratio = gamma(3.0 / rho) / gamma(1.0 / rho)
mix_var = (alpha * (mu**2 + ratio / sbeta**2)).sum(axis=0)
# de-sphered sensor-space mixing: A = pinv(W_fort @ S), columns = maps.
# MLX's W is model-major ((n_models, n, n)), so the per-model slice is
# W[model_idx] rather than torch's W[:, :, model_idx].
w_fort = np.array(self.W[model_idx].T, dtype=np.float64)
sphere = np.array(self.sphere, dtype=np.float64)
a_sensor = np.linalg.pinv(w_fort @ sphere)
svar = mix_var * (a_sensor**2).sum(axis=0)
order = np.argsort(-svar)
if return_svar:
return order, svar[order]
return order
# ------------------------------------------------------------------
# MIR/PMI diagnostics (issue #137; epic #278 Phase 3/#289 port of
# AMICATorchNG.mir/pmi)
# ------------------------------------------------------------------
def _pca_reduction_requested(self, n_channels: int) -> bool:
"""Whether the explicit ``pcakeep``/``pcadb`` asks to fit fewer than
``n_channels`` dimensions (port of
``AMICATorchNG._pca_reduction_requested``;
both delegate to :func:`pamica.rank.pca_reduction_requested`,
issue #323).
``n_channels`` is the channel count of the data being fitted, not
``self.n_channels``, which :meth:`_preprocess` shrinks to the kept
rank. Config-only, not geometry: used solely by :meth:`_fit_once`'s
upfront ``mir_step`` gate, which runs before this fit's sphere exists,
so AUTOMATIC ``mineig``/``mineig_rel`` reduction is not knowable here.
Use :meth:`_pca_reduced` wherever a fitted sphere already exists.
"""
return pca_reduction_requested(
self.pcakeep, self.pcadb, n_channels, self.do_sphere
)
def _pca_reduced(self) -> bool:
"""Whether the fitted sphere is rank-reduced (non-square) -- the #300
fitted-geometry guard (port of ``AMICATorchNG._pca_reduced``).
Derived from the fitted geometry (``sphere.shape[0] !=
sphere.shape[1]``), so it catches rank reduction from an explicit
``pcakeep``/``pcadb`` and from AUTOMATIC numerical-rank detection
(``mineig``/``mineig_rel``) alike. The config-only
:meth:`_pca_reduction_requested` complements it for the upfront
``mir_step`` gate, which runs before this fit's sphere exists.
``False`` before :meth:`fit` (``sphere`` is ``None``) and for a
full-rank fit.
"""
return self.sphere is not None and self.sphere.shape[0] != self.sphere.shape[1]
def mir(
self, X: np.ndarray, *, model_idx: int = 0, nbins: Optional[int] = None
) -> Tuple[float, float]:
"""Mutual Information Reduction (issue #137) of this model's unmixing
on ``X``.
Composes the linear part of the raw-data-to-sources transform, ``W_fort @ sphere``
-- i.e. ``get_unmixing_matrix(model_idx) @ sphere`` -- and delegates
to :func:`pamica.metrics.mir`. MIR is shift-invariant, so the
data-space mean/``c`` centering :meth:`transform` applies is
irrelevant here. Computed through this backend's float32 parameters,
so treat the result as ~7-significant-digit, not float64-parity --
fine for a diagnostic (see the module docstring's precision note).
Parameters
----------
X : np.ndarray of shape (n_channels, n_samples)
Raw (unpreprocessed) data.
model_idx : int, default=0
Which model's unmixing to use.
nbins : int, optional
Histogram bin count; see :func:`pamica.metrics.mir`.
Returns
-------
mir_nats : float
variance : float
Raises
------
RuntimeError
If the model is unfitted, or the fit ended degenerate
(issue #306).
ValueError
If ``X`` is not a 2D array of the fitted input channel count, or
if the fitted sphere is rank-reduced (non-square): whether from
explicit ``pcakeep``/``pcadb`` or from automatic ``mineig``/
``mineig_rel`` numerical-rank detection, the sphere is
rank-deficient, so MIR's log-Jacobian term is undefined
(issue #283/#300).
"""
if self.A is None or self.W is None or self.sphere is None:
raise RuntimeError(
"AMICAMLXNG.mir() requires a fitted model; call fit() first."
)
self._check_model_idx(model_idx)
self._check_usable("compute MIR")
self._check_input_shape(X)
if self._pca_reduced():
raise ValueError(
"mir() is incompatible with PCA reduction: the fitted "
f"sphere is rank-deficient ({self.n_channels} of "
f"{self.n_channels_in} channels kept), whether from explicit "
"pcakeep/pcadb or automatic mineig/mineig_rel numerical-rank "
"detection, so MIR's log-Jacobian term is undefined for the "
"resulting non-square/non-invertible unmixing."
)
unmixing = np.array(self.W[model_idx].T @ self.sphere)
return mir_metric(unmixing, X, nbins)
def pmi(
self, X: np.ndarray, *, model_idx: int = 0, nbins: Optional[int] = None
) -> np.ndarray:
"""Pairwise Mutual Information (issue #137) between this model's
sources on ``X``.
Delegates to :func:`pamica.metrics.pairwise_mi` on
``transform(X, model_idx)``.
Parameters
----------
X : np.ndarray of shape (n_channels, n_samples)
Raw (unpreprocessed) data.
model_idx : int, default=0
Which model's sources to use.
nbins : int, optional
Histogram bin count; see :func:`pamica.metrics.pairwise_mi`.
Returns
-------
mi_matrix : np.ndarray of shape (n_sources, n_sources)
Raises
------
RuntimeError
If the model is unfitted, or the fit ended degenerate
(issue #306), both via :meth:`transform`.
ValueError
If ``X`` is not a 2D array of the fitted input channel count
(via :meth:`transform`).
"""
return pairwise_mi(self.transform(X, model_idx=model_idx), nbins)
# ------------------------------------------------------------------
# Multi-model posterior (issue #141; epic #278 Phase 3/#289 port of
# AMICATorchNG.model_loglik/model_probability)
# ------------------------------------------------------------------
def model_loglik(self, X: np.ndarray) -> np.ndarray:
"""Per-model, per-sample log-likelihood ``Lht`` on (new) data.
For each model ``h`` and sample ``t`` this is the joint log-likelihood
``log(gm[h]) + log|det W_h| + sldet + sum_i log p_h(s_i)`` (Fortran's
``Lht``/``modloglik``), evaluated on arbitrary raw data via the
STORED sphere/mean -- never re-preprocessing, which would overwrite
them. The per-sample posterior over models (model dominance) is
``softmax(Lht, axis=0)``; see :meth:`model_probability`.
This does not replicate a training-time ``do_reject`` mask: it
scores every sample of ``X``. On a ``do_reject`` fit's own training
data it therefore returns real values where the stored ``_llt_lht``
carries Fortran's sentinel zeros for rejected samples (issue #155),
so the two agree bit-for-bit only when the fit did not use
``do_reject``. Like :meth:`transform`, it assumes a usable
(non-degenerate) fit; the :class:`~pamica.AMICA` wrapper enforces
that via ``_check_usable``.
Parameters
----------
X : np.ndarray of shape (n_channels, n_samples)
Raw (unpreprocessed) data.
Returns
-------
Lht : np.ndarray of shape (n_models, n_samples)
Raises
------
RuntimeError
If the model is unfitted, or the fit ended degenerate
(issue #306).
ValueError
If ``X`` is not a 2D array of the fitted input channel count, or
contains non-finite (NaN/Inf) values.
"""
if self.sphere is None or self.mean is None or self.W is None:
raise RuntimeError(
"AMICAMLXNG.model_loglik() requires a fitted model; call fit() first."
)
self._check_usable("compute the model log-likelihood")
self._check_input_shape(X)
return self._model_loglik_unchecked(X)
def _model_loglik_unchecked(self, X: np.ndarray) -> np.ndarray:
"""Core ``Lht`` computation for :meth:`model_loglik`, with no
degenerate-fit guard or shape validation of its own (issue #306
PR #329 review; port of ``AMICATorchNG._model_loglik_unchecked``):
:meth:`model_loglik` and :meth:`model_probability` each do their own
single guard + shape check, with their own action wording, then
both call this -- so the guard no longer runs twice on a
:meth:`model_probability` call, which used to run its own
``_check_usable`` and then :meth:`model_loglik`'s (measured ~2x the
cost of a single guard evaluation on this backend)."""
X = np.ascontiguousarray(X)
if not np.isfinite(X).all():
bad = np.flatnonzero(~np.isfinite(X).all(axis=1))
raise ValueError(
"AMICAMLXNG.model_loglik(): input contains non-finite "
f"(NaN/Inf) values in {bad.size} channel(s) {bad.tolist()}; "
"clean bad segments before scoring."
)
X_arr = mx.array(X.astype(np.float32))
X_t = self.sphere @ (X_arr - self.mean)
n_samples = X_t.shape[1]
Lht = np.zeros((self.n_models, n_samples), dtype=np.float32)
for start in range(0, n_samples, self.block_size):
end = min(start + self.block_size, n_samples)
logV, *_ = self._forward(X_t[:, start:end])
Lht[:, start:end] = np.array(logV).T
return Lht
def model_probability(self, X: np.ndarray) -> np.ndarray:
"""Per-sample posterior probability of each model (model dominance).
The column-wise ``softmax`` over models of :meth:`model_loglik`,
i.e. ``P(model h | x_t)``; each column sums to 1. For a single model
this is all ones.
Parameters
----------
X : np.ndarray of shape (n_channels, n_samples)
Raw (unpreprocessed) data.
Returns
-------
prob : np.ndarray of shape (n_models, n_samples)
Raises
------
RuntimeError
If the model is unfitted, or the fit ended degenerate
(issue #306).
ValueError
If ``X`` is not a 2D array of the fitted input channel count, if
``X`` is non-finite, if every model underflows to ``-inf``
log-likelihood at some sample (the posterior is undefined
there), or if a log-likelihood is NaN (numerical corruption,
distinct from the ``-inf`` underflow case above).
"""
if self.sphere is None or self.mean is None or self.W is None:
raise RuntimeError(
"AMICAMLXNG.model_probability() requires a fitted model; "
"call fit() first."
)
self._check_usable("compute the model probability")
self._check_input_shape(X)
Lht = self._model_loglik_unchecked(X)
# NaN and -inf are different failure modes and must not share a
# message: -inf is every model underflowing at a real sample (an
# extreme outlier), while NaN is numerical corruption. isfinite alone
# conflates them (PR #311 review scope extension, issue #306). The
# diagnosis + normalization is shared with AMICATorchNG (PR #329
# review) rather than duplicated per backend.
return model_probability_from_loglik(
Lht, caller="AMICAMLXNG.model_probability()"
)
# ------------------------------------------------------------------
# EEGLAB export (issue #92; epic #278 Phase 3/#289 port of
# AMICATorchNG.write_amica_output)
# ------------------------------------------------------------------
def write_amica_output(self, outdir) -> None:
"""Write this fitted model as the Fortran/EEGLAB AMICA output
directory.
Produces the raw binary files that EEGLAB's ``loadmodout15.m`` (and
the Python port :func:`pamica.numpy_impl.load.loadmodout`) read:
``gm``, ``W``, ``S``, ``mean``, ``c``, ``alpha``, ``mu``, ``sbeta``,
``rho``, ``comp_list``, ``LL``, so an MLX fit drops directly into an
EEGLAB workflow, exactly like ``AMICATorchNG.write_amica_output``.
``loadmodout15`` performs the variance-ordering and unit-norm
normalization on load, so the on-disk parameters are written in fit
order. Single-model output is byte-compatible with the Fortran
reference.
Also writes ``LLt`` (the per-sample/per-model log-likelihood,
issue #155) for a model that was just :meth:`fit` in this process, from the
stash the training E-step filled (issue #157) -- so, exactly as in
the reference, ``LLt`` is the E-step of the returned iterate: one
M-step older than the ``W``/``A`` written beside it after a fit that ran
to ``max_iter``, their own after a convergence stop (see
:meth:`_fit_once`'s docstring). A model restored via
:meth:`from_state_dict`/:meth:`load` carries no stash, so ``LLt`` is
omitted for it (a warning is logged) -- the rest of the output is
unaffected. Under ``do_reject``, a rejected sample's ``LLt`` entries
are written as exactly 0.0 (the load-bearing sentinel ``load_rej``
reconstructs from, amica15.f90:2231-2234): this is automatic,
because :meth:`_reject_outliers` already zeroes the stash for
dropped samples as it drops them.
Raises if the model is unfitted or degenerate (a fit that ended on a
non-finite log-likelihood): a NaN model must not be written silently.
The scikit-learn-style :class:`~pamica.AMICA` wrapper
(``backend="mlx"``, issue #313) already refuses this via its own
usability gate, but a caller using :class:`AMICAMLXNG` directly has
no such gate in front of this method, so the guard lives here too --
mirrors :meth:`state_dict`'s two-layer guard (stop_reason
refusal, then a defense-in-depth isfinite sweep over the parameter
arrays) so the same protection applies to a direct
``write_amica_output`` call (PR #311 review).
Parameters
----------
outdir : str or path-like
Destination directory (created if absent).
"""
if self.A is None:
raise RuntimeError(
"write_amica_output requires a fitted model; call fit() first."
)
if self.stop_reason in self._DEGENERATE_STOP_REASONS:
raise RuntimeError(
f"Refusing to write output for a degenerate model (stop_reason="
f"{self.stop_reason!r}): fit() hit a non-finite value "
f"at iteration {self.iteration}. Fix the instability (lower "
f"lrate, disable Newton, or check data conditioning) before "
f"writing."
)
# Defense-in-depth, mirroring state_dict(): catch a non-finite
# parameter even if stop_reason bookkeeping ever misses it. Also
# neutralizes a stale LLt stash: a failed final iteration's
# _llt_lht/_llt_lt (from before a degenerate break) can no
# longer reach disk once this guard refuses the write outright.
nonfinite = self._nonfinite_params()
if nonfinite:
raise RuntimeError(
f"Refusing to write output for a model with non-finite "
f"parameters {nonfinite} (stop_reason={self.stop_reason!r})."
)
from ..numpy_impl.load import write_amicaout
# The exported parameters are the fit()-kept iterate (LL ==
# final_ll_). Under the keep_best safeguard (#51) that can be an
# earlier iterate than the last, so end the written LL trajectory at
# that iterate rather than at a later, discarded overshoot --
# otherwise LL[-1] would not match the model just written. Monotone
# runs keep the full trajectory unchanged.
ll = np.asarray(self.ll_history, dtype=np.float64)
if (
self.final_ll_ is not None
and np.isfinite(self.final_ll_)
and ll.size
and not np.isclose(ll[-1], self.final_ll_)
):
ll = ll[: int(np.argmax(ll)) + 1]
# LLt (Fortran's per-sample/per-model log-likelihood, issue #155):
# computed once at the end of fit() (after any keep-best restore)
# and stored compactly on self. A model restored via
# from_state_dict()/load() never ran fit() in this process, so it
# has neither -- warn rather than silently omitting the file
# (silent-failure review).
if self._llt_lht is not None and self._llt_lt is not None:
Lht, Lt = self._llt_lht, self._llt_lt
else:
logger.warning(
"No LLt data available (model was restored via "
"from_state_dict()/load(), not freshly fit()); writing "
"output without the LLt file."
)
Lht = Lt = None
write_amicaout(
outdir,
gm=np.array(self.gm),
# write_amicaout's contract is W(nw, nw, num_models) -- the
# SAME layout AMICATorchNG's W already is. MLX's W is
# model-major, (n_models, n, n) (see the module docstring and
# transform()'s CAUTION note), so move the model axis from
# front to back rather than transposing torch's tensor layout.
W=np.array(self.W).transpose(1, 2, 0),
sphere=np.array(self.sphere),
mean=np.array(self.mean),
c=np.array(self.c),
alpha=np.array(self.alpha),
mu=np.array(self.mu),
sbeta=np.array(self.beta), # Fortran's 'sbeta' is pamica's beta
rho=np.array(self.rho),
comp_list=np.array(self.comp_list),
ll=ll,
# The reference layout, (nw, num_comps) with component k in column
# k: the component-row A transposed (issue #334).
A=np.array(self.A).T,
Lht=Lht,
Lt=Lt,
)
# ------------------------------------------------------------------
# Persistence (issue #287)
# ------------------------------------------------------------------
# Full fitted-parameter snapshot -- the same 12-name set as AMICATorchNG's
# _PARAM_TENSORS: A/W/c/comp_list/mean/
# sphere are what transform()/get_*matrix() read back; mu/alpha/beta/rho/
# gm are the mixture-PDF EM state; pdtype is the per-source density-family
# code (issue #265) -- a non-default pdftype model, or the adaptive
# switcher's chosen 1/4 assignments, would otherwise silently revert to GG
# on reload. comp_list and pdtype are FORCE-CAST to their integer dtype on
# load (via _safe_int_cast, not a bare .astype -- see _load_params), not
# merely "preserved": a restored array that is already integer keeps its
# dtype unchanged, but one that arrives as float (e.g. a hand-edited or
# foreign-tool payload) is only accepted if every value is finite and
# whole, and rejected with a named error otherwise. The rest are float32.
_PARAM_ARRAYS = (
"A", "W", "c", "mu", "alpha", "beta", "rho", "gm",
"comp_list", "mean", "sphere", "pdtype",
) # fmt: skip
# Integer arrays in _PARAM_ARRAYS: their dtype is restored explicitly
# (via _safe_int_cast) rather than following the float32 default the rest
# take.
_INT_PARAM_DTYPES = {"comp_list": np.int64, "pdtype": np.int32}
# extra's full key set as of format_version 1's original release.
# from_state_dict()/load() require every one of these to be present
# (_load_params raises a named error naming what's missing, mirroring the
# params check above) rather than defaulting a subset via .get() --
# phase-1-era payloads genuinely have all of them, so silently
# defaulting one would hide real corruption. A field added LATER, once
# payloads without it already exist, is instead loaded additively via
# extra.get() with a documented fallback and deliberately left OUT of
# this tuple (the pattern AMICATorchNG's #198 restart_seeds_/
# restart_lls_/restart_stop_reasons_ and #207 convergence-stop config
# keys established): epic #278 Phase 3/#289's ``numrej``/``good_idx``
# (outlier rejection, issue #123's mechanism) are the first such case --
# see their extra.get() reads in _load_params.
_EXTRA_KEYS = (
"sldet", "iteration", "ll_history", "final_ll", "stop_reason",
"n_kurt_done", "n_newton_fallbacks", "lrate", "lrate_cap", "newtrate",
"rholrate", "restart_seeds_", "restart_lls_", "restart_stop_reasons_",
) # fmt: skip
# This backend owns its own format_version, independent of AMICATorchNG's
# (currently 4): the two payloads are never interchangeable (different
# param layouts, no dtype/device fields here), so there is no reason for
# the version numbers to track each other. Version 2 (issue #334) stores A
# with one component per row, shape (n_comps, n_channels); a version 1
# payload stored it as (n_channels, n_comps) and is converted on load when
# unmerged, refused when share_comps had merged components
# (pamica.component_layout, ADR 0007).
_SAVE_FORMAT_VERSION = 2
_COLUMN_LAYOUT_FORMAT_VERSION = 1
def state_dict(self) -> dict:
"""Serialize the fitted model to a plain, framework-agnostic dict.
The returned dict has three parts: ``config`` (the constructor
arguments needed to rebuild the object), ``params`` (the fitted
arrays, as numpy), and ``extra`` (scalar/schedule state). Every value
is a numpy array or a plain Python primitive, so the dict is JSON/
``.npz``-safe (see :meth:`save`). Rebuild with :meth:`from_state_dict`.
Raises if the model is unfitted or degenerate (a fit that ended on a
non-finite log-likelihood): a NaN model must not be persisted
silently.
"""
if self.A is None:
raise RuntimeError(
"AMICAMLXNG.state_dict() requires a fitted model; call fit() first."
)
if self.stop_reason in self._DEGENERATE_STOP_REASONS:
raise RuntimeError(
f"Refusing to serialize a degenerate model (stop_reason="
f"{self.stop_reason!r}): fit() hit a non-finite value "
f"at iteration {self.iteration}. Fix the instability (lower "
f"lrate, disable Newton, or check data conditioning) before "
f"saving."
)
# Defense-in-depth: catch a non-finite parameter even if stop_reason
# bookkeeping ever misses it (the codebase has known NaN-suppression
# risks). isfinite on the integer comp_list/pdtype is trivially
# all-True.
nonfinite = self._nonfinite_params()
if nonfinite:
raise RuntimeError(
f"Refusing to serialize a model with non-finite parameters "
f"{nonfinite} (stop_reason={self.stop_reason!r})."
)
config = {
"n_channels": self.n_channels,
"n_models": self.n_models,
"n_mix": self.n_mix,
# block_size is the value the fit actually ran at -- which, under
# do_opt_block, is the size the search chose rather than the one
# the constructor was given (issue #232), so a reloaded model
# reproduces the run it came from; the sweep bounds ride along so
# a re-fit can search again if asked.
"block_size": self.block_size,
"do_opt_block": self.do_opt_block,
"blk_min": self.blk_min,
"blk_max": self.blk_max,
"blk_step": self.blk_step,
# lrate/newtrate/rholrate are annealed during fit; persist the
# original constructor values (lrate0/newtrate0/rholrate0) and
# restore the mutated ones from ``extra`` below.
"lrate": self.lrate0,
"minlrate": self.minlrate,
"lratefact": self.lratefact,
"maxdecs": self.maxdecs,
"use_min_dll": self.use_min_dll,
"min_dll": self.min_dll,
"maxincs": self.maxincs,
"use_grad_norm": self.use_grad_norm,
"min_nd": self.min_nd,
"newt_ramp": self.newt_ramp,
"newt_start": self.newt_start,
"newtrate": self.newtrate0,
"do_newton": self.do_newton,
# Outlier rejection (issue #123's AMICATorchNG mechanism, epic
# #278 Phase 3/#289): a phase-1/2-era payload's config dict lacks
# these keys, and cls(**config) then falls back to the
# constructor's do_reject=False default -- no format_version bump
# needed (same precedent as keep_best above). The rejection state
# a fit actually reached (numrej/good_idx) is in ``extra`` below.
"do_reject": self.do_reject,
"rejsig": self.rejsig,
"rejstart": self.rejstart,
"rejint": self.rejint,
"maxrej": self.maxrej,
"rho0": self.rho0,
"minrho": self.minrho,
"maxrho": self.maxrho,
"rholrate": self.rholrate0,
"rholratefact": self.rholratefact,
# Density-family selection (issue #265): needed so a reloaded
# model rebuilds with the right pdftype/dorho/do_choose_pdfs and
# switch schedule instead of the GG default.
"pdftype": self.pdftype,
"kurt_start": self.kurt_start,
"num_kurt": self.num_kurt,
"kurt_int": self.kurt_int,
"invsigmin": self.invsigmin,
"invsigmax": self.invsigmax,
"doscaling": self.doscaling,
"scalestep": self.scalestep,
# Component sharing (issue #263): persisted so a reloaded
# multi-model run keeps its schedule; the merged comp_list itself
# is in params.
"share_comps": self.share_comps,
"share_start": self.share_start,
"share_iter": self.share_iter,
"comp_thresh": self.comp_thresh,
"do_mean": self.do_mean,
"do_sphere": self.do_sphere,
"do_approx_sphere": self.do_approx_sphere,
# Explicit PCA reduction (issue #323). Additive, like keep_best
# below: a payload written before #323 lacks both keys, and
# cls(**config) then falls back to the constructor defaults (None),
# which is what that fit ran with, so no format_version bump. Cast
# to plain int/float because save() JSON-encodes config and the
# validator accepts numpy scalars (np.int64), which json cannot.
"pcakeep": None if self.pcakeep is None else int(self.pcakeep),
"pcadb": None if self.pcadb is None else float(self.pcadb),
"mineig": self.mineig,
"mineig_rel": self.mineig_rel,
"seed": self.seed,
# Best-of-N restarts (issue #198). Persisted so a reloaded model
# reconstructs its exact configuration; the restart the fit
# actually kept is in ``extra`` below.
"n_restarts": self.n_restarts,
"restart_seeds": self.restart_seeds,
# Best-iterate safeguard flag (issue #51, epic #278 Phase 2/#288);
# only affects a re-fit, but persisted so a reloaded model
# reconstructs its exact configuration. Additive: a phase-1-era
# payload's config dict lacks this key, and ``cls(**config)`` then
# falls back to the constructor default (True) -- no
# format_version bump needed (same precedent as torch's #207).
"keep_best": self.keep_best,
}
params = {name: np.array(getattr(self, name)) for name in self._PARAM_ARRAYS}
extra = {
"sldet": float(self.sldet),
"iteration": int(self.iteration),
"ll_history": [float(v) for v in self.ll_history],
"final_ll": None if self.final_ll_ is None else float(self.final_ll_),
"stop_reason": self.stop_reason,
"n_kurt_done": int(self.n_kurt_done),
"n_newton_fallbacks": int(self.n_newton_fallbacks),
"lrate": float(self.lrate),
"lrate_cap": float(self.lrate_cap),
"newtrate": float(self.newtrate),
"rholrate": float(self.rholrate),
"rholrate_cap": float(self.rholrate_cap),
# Per-restart records (issue #198): which seeds ran, what each
# returned, and why each stopped.
"restart_seeds_": list(self.restart_seeds_),
"restart_lls_": [float(v) for v in self.restart_lls_],
"restart_stop_reasons_": list(self.restart_stop_reasons_),
# Outlier rejection (issue #123's AMICATorchNG mechanism, epic
# #278 Phase 3/#289). Deliberately NOT added to _EXTRA_KEYS
# (which _load_params checks strictly): these are the first
# extra fields added after format_version 1 shipped, so a
# phase-1/2-era payload genuinely lacks them, and _load_params
# falls back with extra.get() -- the additive pattern
# AMICATorchNG's #198 restart_seeds_/restart_lls_/
# restart_stop_reasons_ established there. ``good_idx`` is
# written as a plain list (not a numpy array): ``save()`` JSON-
# encodes ``extra`` via ``json.dumps``, which cannot serialize
# ndarrays.
"numrej": int(self.numrej),
"good_idx": None
if self.good_idx is None
else np.array(self.good_idx).astype(np.int64).tolist(),
}
return {
"format_version": self._SAVE_FORMAT_VERSION,
"config": config,
"params": params,
"extra": extra,
}
@classmethod
def from_state_dict(cls, state: dict) -> "AMICAMLXNG":
"""Rebuild a fitted :class:`AMICAMLXNG` from :meth:`state_dict` output.
Unlike ``AMICATorchNG.from_state_dict`` there is no ``device``
argument: this backend always runs on ``mx.default_device()``.
A ``format_version`` 1 state (components as columns of ``A``, before
issue #334) loads unchanged in every other respect: its ``A`` is
converted to component rows without loss, unless ``share_comps`` had
merged components, which raises ``ValueError`` asking for a refit
(:func:`pamica.component_layout.rows_from_legacy_columns`).
"""
version = state.get("format_version")
if version not in (cls._SAVE_FORMAT_VERSION, cls._COLUMN_LAYOUT_FORMAT_VERSION):
raise ValueError(
f"unsupported AMICAMLXNG state format_version: {version!r} "
f"(expected {cls._SAVE_FORMAT_VERSION}, or "
f"{cls._COLUMN_LAYOUT_FORMAT_VERSION} from before the "
"component-row layout)"
)
for section in ("config", "params", "extra"):
if section not in state:
raise ValueError(
f"malformed AMICAMLXNG state: missing {section!r} section "
f"(format_version={version}); the payload may be truncated."
)
config = dict(state["config"])
# A missing/unexpected key in a malformed or foreign-version payload
# surfaces as a bare TypeError from the constructor call; every other
# validation step in this method already names the payload as the
# culprit with a ValueError, so wrap this one the same way instead of
# letting a mismatched-keyword TypeError propagate unexplained
# (issue #306; :meth:`load`'s .npz path delegates to this method, so
# it is covered too).
try:
obj = cls(**config)
except TypeError as exc:
raise ValueError(
f"malformed AMICAMLXNG state: config does not match the "
f"AMICAMLXNG constructor ({exc}); the payload may be "
"truncated or from an incompatible version."
) from exc
if version == cls._COLUMN_LAYOUT_FORMAT_VERSION:
params = state["params"]
missing = [name for name in ("A", "comp_list") if name not in params]
if missing:
raise ValueError(
f"malformed AMICAMLXNG state: missing params {missing}"
)
A_rows = rows_from_legacy_columns(
np.asarray(params["A"]),
np.asarray(params["comp_list"]),
owner="AMICAMLXNG",
)
state = {**state, "params": {**params, "A": A_rows}}
obj._load_params(state)
return obj
def _load_params(self, state: dict) -> None:
"""Restore fitted arrays/scalars from :meth:`state_dict` output onto
this instance."""
params = state["params"]
missing = [name for name in self._PARAM_ARRAYS if name not in params]
if missing:
raise ValueError(f"malformed AMICAMLXNG state: missing params {missing}")
arrays = {name: np.asarray(params[name]) for name in self._PARAM_ARRAYS}
# Guard against config/params drift: every param must match the
# dimensions _initialize_parameters actually allocates, or
# transform()/the E-step would fail later with a confusing matmul
# error far from load() (or, worse, silently broadcast wrong).
# Shapes are read off _initialize_parameters/_update_unmixing_matrices
# (core.py, this module), not guessed: A/mu/alpha/beta/rho/gm/c/
# comp_list/pdtype/W are all sized from n_channels/n_models/n_mix/
# n_comps, every one of which the constructor (cls(**config) in
# from_state_dict) has already derived. ``sphere``'s SECOND
# dimension is the exception: it is the ORIGINAL input-channel count,
# which is not recoverable from config alone once rank reduction has
# happened (issue #223) -- self.n_channels is already the
# post-reduction value by the time state_dict() ran (see that
# method's config comment) -- so there is nothing independent to
# check sphere's width against. It is instead taken from the
# restored sphere itself, and mean (the only other array on that
# axis) is cross-checked against it.
n, m, ncomp, nmix = self.n_channels, self.n_models, self.n_comps, self.n_mix
sphere_shape = arrays["sphere"].shape
if len(sphere_shape) != 2 or sphere_shape[0] != n:
raise ValueError(
f"restored sphere has shape {sphere_shape}, expected "
f"(n_channels, n_channels_in) with n_channels={n}"
)
n_channels_in = sphere_shape[1]
# config's n_channels is the fitted (possibly reduced) rank, so the
# constructor set the input count to it; the restored sphere's width
# is the true input count, which a refit validates X against.
self._n_input_channels = int(n_channels_in)
expected_shapes = {
"A": (ncomp, n),
"W": (m, n, n),
"c": (n, m),
"mu": (nmix, ncomp),
"alpha": (nmix, ncomp),
"beta": (nmix, ncomp),
"rho": (nmix, ncomp),
"gm": (m,),
"comp_list": (n, m),
"mean": (n_channels_in, 1),
"pdtype": (n, m),
}
for name, expected in expected_shapes.items():
actual = arrays[name].shape
if actual != expected:
raise ValueError(
f"restored {name!r} has shape {actual}, expected "
f"{expected} for n_channels={n}, n_models={m}, "
f"n_mix={nmix}, n_comps={ncomp}, "
f"n_channels_in={n_channels_in}"
)
for name in self._PARAM_ARRAYS:
value = arrays[name]
if name in self._INT_PARAM_DTYPES:
value = _safe_int_cast(name, value, self._INT_PARAM_DTYPES[name])
else:
value = value.astype(np.float32)
# Named-error isfinite validation (PR #318 review): the 10
# float _PARAM_ARRAYS had NO finiteness check at all --
# _safe_int_cast above only covers comp_list/pdtype -- so a
# NaN/inf-poisoned payload (corrupted file, hand-edited
# .npz, truncated write) would load "successfully" and only
# surface later as a confusing downstream NaN with no
# diagnostic tying it back to load(), the exact failure mode
# the W-singularity check just below this loop already
# guards against for W specifically. Extends that same
# promise to every float param, matching load()'s docstring.
if not np.all(np.isfinite(value)):
raise ValueError(
f"malformed AMICAMLXNG state: restored {name!r} has "
f"non-finite values; the payload may be corrupted."
)
setattr(self, name, mx.array(value))
# Rebuild the MLX-only per-iteration caches (module docstring):
# unlike AMICATorchNG, which recomputes log|det W| and
# lgamma(1+1/rho) inline on every call, this backend hoists them to
# once-per-iteration cached arrays. A loaded model has no fit history
# to hoist them from, so they are rebuilt here from the restored
# params -- before this returns, transform()/_forward()/comp_used and
# every get_* accessor must work exactly as they would mid-fit.
self._comp_used_arr = mx.array(
np.isin(np.arange(self.n_comps), np.unique(np.array(self.comp_list)))
)
self._refresh_lgamma_table()
assert self.W is not None # just set by the loop above
logdets = [
mx.linalg.slogdet(self.W[h], stream=_CPU)[1] for h in range(self.n_models)
]
self._logdet_W = mx.stack(logdets)
# slogdet returns -inf (not a raised error) for a finite-but-singular
# matrix, so a corrupted or hand-edited W that is finite yet singular
# would otherwise load "successfully" and only fail later, deep
# inside _forward, with no diagnostic tying it back to load().
if not bool(mx.all(mx.isfinite(self._logdet_W)).item()):
raise ValueError(
"malformed AMICAMLXNG state: restored W is singular for at "
"least one model (log|det W| is non-finite); the payload may "
"be corrupted."
)
# sphere was just replaced, so any cached back-map describes the old
# one. _sphere_np backs _pinv_sphere (get_sensor_mixing_matrix,
# _identify_shared_comps) at float64 precision during a live fit, but
# only the float32 ``sphere`` is a persisted param (the fixed
# 12-name set above) -- so a reloaded model's _sphere_np is the
# float32 sphere upcast to float64, not the higher-precision value
# _preprocess originally computed. This cannot affect transform()
# (which reads self.sphere directly, so it stays bit-identical
# pre/post round trip): only get_sensor_mixing_matrix() on a
# reloaded model carries this small extra rounding.
self._sphere_np = np.array(self.sphere, dtype=np.float64)
self._sphere_pinv = None
extra = state["extra"]
missing_extra = [key for key in self._EXTRA_KEYS if key not in extra]
if missing_extra:
raise ValueError(
f"malformed AMICAMLXNG state: missing extra fields {missing_extra}"
)
self.sldet = extra["sldet"]
self.iteration = extra["iteration"]
self.ll_history = list(extra["ll_history"])
self.final_ll_ = extra["final_ll"]
self.stop_reason = extra["stop_reason"]
self.n_kurt_done = extra["n_kurt_done"]
self.n_newton_fallbacks = extra["n_newton_fallbacks"]
# The rates, validated (a non-finite one is a named ValueError);
# rholrate_cap is additive, NOT in _EXTRA_KEYS, see _saved_rates.
for name, value in _saved_rates(extra, "AMICAMLXNG").items():
setattr(self, name, value)
self.restart_seeds_ = list(extra["restart_seeds_"])
self.restart_lls_ = list(extra["restart_lls_"])
self.restart_stop_reasons_ = list(extra["restart_stop_reasons_"])
# Additive-only (issue #123's AMICATorchNG mechanism, epic #278
# Phase 3/#289), NOT in _EXTRA_KEYS: a payload written before Phase 3
# simply has no rejection state, and a good_idx of None / numrej of 0
# is the honest description of a model that never rejected anything
# (the same fallback AMICATorchNG's _load_params uses for its own
# #198-era restart_seeds_/restart_lls_/restart_stop_reasons_ keys).
# Wrapped in the same named-malformed-state pattern every other
# field in this method uses (PR #318 review): a corrupted or
# hand-edited payload's numrej/good_idx previously reached raw
# int()/np.asarray() calls with no try/except, so a malformed value
# there raised an opaque bare TypeError/ValueError instead of the
# "malformed AMICAMLXNG state: ..." message this method promises
# for everything else.
try:
self.numrej = int(extra.get("numrej", 0))
except (TypeError, ValueError) as exc:
raise ValueError(
f"malformed AMICAMLXNG state: extra['numrej'] is not a "
f"valid integer ({exc})."
) from exc
good_idx = extra.get("good_idx")
if good_idx is None:
self.good_idx = None
else:
try:
good_idx_arr = np.asarray(good_idx, dtype=np.int64)
except (TypeError, ValueError) as exc:
raise ValueError(
f"malformed AMICAMLXNG state: extra['good_idx'] could "
f"not be converted to an integer index array ({exc})."
) from exc
self.good_idx = mx.array(good_idx_arr)
def save(self, filepath: str) -> None:
"""Persist the fitted model to ``filepath`` as a single ``.npz``
(issue #287).
Device- and framework-agnostic by construction: ``config``/``extra``
are embedded as JSON-encoded 0-d string arrays and ``params`` as
native numpy arrays, all written by ``np.savez_compressed`` -- no
torch coupling, no pickle. Reload with :meth:`load`. Raises the same
refusal guards as :meth:`state_dict` (unfitted, degenerate, or
non-finite parameters).
"""
state = self.state_dict()
np.savez_compressed(
filepath,
format_version=state["format_version"],
config=json.dumps(state["config"]),
extra=json.dumps(state["extra"]),
**state["params"],
)
@classmethod
def load(cls, filepath: str) -> "AMICAMLXNG":
"""Rebuild a fitted :class:`AMICAMLXNG` from a file written by
:meth:`save`.
Wrong ``format_version``, a missing ``config``/``extra``/``params``
section, a truncated archive missing one of the 12 param arrays, or a
genuinely corrupt/byte-truncated ``.npz`` (not a valid zip, or a zip
whose central directory or a member's compressed bytes were cut off)
each raise a named ``ValueError`` naming what is wrong, rather than
raising an opaque ``zipfile``/``numpy`` error or loading a silently
partial model. A ``format_version`` 1 file (before issue #334) is
converted or refused exactly as :meth:`from_state_dict` describes.
"""
# Eagerly materialize every array the archive actually contains
# INSIDE this try, so any corruption -- an unreadable zip (raised by
# np.load itself), or one truncated member's compressed bytes
# (raised lazily, on that member's own read) -- surfaces here as one
# of the well-known exception types below, not scattered across the
# validation logic beneath. Everything from here down reads only the
# already-materialized ``raw`` dict, so those checks keep raising
# their own specific ValueErrors unchanged.
try:
with np.load(filepath, allow_pickle=False) as data:
raw = {name: data[name] for name in data.files}
except (zipfile.BadZipFile, EOFError, OSError, ValueError) as exc:
raise ValueError(
f"malformed AMICAMLXNG save file {filepath!r}: could not read "
f"the archive ({exc}); the file may be truncated or corrupted."
) from exc
for section in ("format_version", "config", "extra"):
if section not in raw:
raise ValueError(
f"malformed AMICAMLXNG save file {filepath!r}: missing "
f"{section!r} (the file may be truncated or corrupted)."
)
version = int(raw["format_version"])
if version not in (cls._SAVE_FORMAT_VERSION, cls._COLUMN_LAYOUT_FORMAT_VERSION):
raise ValueError(
f"unsupported AMICAMLXNG save format_version: {version!r} "
f"(expected {cls._SAVE_FORMAT_VERSION}, or "
f"{cls._COLUMN_LAYOUT_FORMAT_VERSION} from before the "
"component-row layout)"
)
config = json.loads(raw["config"].item())
extra = json.loads(raw["extra"].item())
missing_params = [name for name in cls._PARAM_ARRAYS if name not in raw]
if missing_params:
raise ValueError(
f"malformed AMICAMLXNG save file {filepath!r}: missing "
f"params {missing_params} (the file may be truncated or "
f"corrupted)."
)
params = {name: np.array(raw[name]) for name in cls._PARAM_ARRAYS}
state = {
"format_version": version,
"config": config,
"params": params,
"extra": extra,
}
return cls.from_state_dict(state)