You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
Phase D: Multi-model support for AMICAMLXNG (Apple GPU) #81
Phase B (#77) measured that the MLX backend is a ~7x Apple-GPU win for single-model
AMICA (and beats an RTX 4090 at EEG scale), while multi-model has no GPU path: AMICAMLXNG
(v1 MVP, #76) is single-model only, and PyTorch-MPS loses to CPU. Extending the MLX backend to
multi-model carries the ~7x win to real multi-model AMICA -- the highest-value follow-up, and the
last piece before epic #74 closes.
Scope
Add n_models > 1 support to AMICAMLXNG, porting the multi-model machinery from the PyTorch
backend (AMICATorchNG, torch_impl/core.py):
comp_list indirection (shapes go from (n_mix, n_channels) to (n_mix, n_comps)), per-model W = inv(A[:, comp_list[:, h]]) and slogdet on the CPU stream (hoisted per iteration).
E-step per-model loop producing logV (batch, n_models), cross-model responsibilities v = softmax(logV, over models), and u = v*z.
M-step: gm = sum_t v_h / n, the per-model exact-EM bias c update (Fortran update_c),
and the gm-weighted A-update scattered through comp_list (byte-identical to single-model when n_models=1).
Component sharing (share_comps) stays deferred (a further follow-up; it builds on this).
Acceptance
AMICAMLXNG(n_models=2) fits the real sample EEG on the Apple GPU (float32) and its converged
LL / per-block sufficient stats match AMICATorchNG(dtype=float32, n_models=2) within the
established multi-model tolerance (multi-model is not partition-identifiable, so compare the
same-init one-iteration stats tightly and the converged LL within distributional spread).
Done via PR #82 (squash-merged into the epic branch, epic #74 Phase D). AMICAMLXNG now supports n_models>1: ported comp_list indirection, per-model W/slogdet (CPU stream), cross-model responsibilities, the per-model exact-EM bias c update, and the gm-weighted comp_list-scattered A-update. Single-model (#76) stays byte-identical. Validated vs AMICATorchNG float32: one-iteration sufficient stats match to float32 precision, converged LL ~1e-5. Benchmark: MLX wins multi-model too (~38-45 ms/it, ~5x over torch-CPU). Sharing remains a fast-follow.
Context
Phase B (#77) measured that the MLX backend is a ~7x Apple-GPU win for single-model
AMICA (and beats an RTX 4090 at EEG scale), while multi-model has no GPU path:
AMICAMLXNG(v1 MVP, #76) is single-model only, and PyTorch-MPS loses to CPU. Extending the MLX backend to
multi-model carries the ~7x win to real multi-model AMICA -- the highest-value follow-up, and the
last piece before epic #74 closes.
Scope
Add
n_models > 1support toAMICAMLXNG, porting the multi-model machinery from the PyTorchbackend (
AMICATorchNG,torch_impl/core.py):comp_listindirection (shapes go from(n_mix, n_channels)to(n_mix, n_comps)), per-modelW = inv(A[:, comp_list[:, h]])andslogdeton the CPU stream (hoisted per iteration).logV (batch, n_models), cross-model responsibilitiesv = softmax(logV, over models), andu = v*z.gm = sum_t v_h / n, the per-model exact-EM biascupdate (Fortranupdate_c),and the gm-weighted A-update scattered through
comp_list(byte-identical to single-model whenn_models=1).Component sharing (
share_comps) stays deferred (a further follow-up; it builds on this).Acceptance
AMICAMLXNG(n_models=2)fits the real sample EEG on the Apple GPU (float32) and its convergedLL / per-block sufficient stats match
AMICATorchNG(dtype=float32, n_models=2)within theestablished multi-model tolerance (multi-model is not partition-identifiable, so compare the
same-init one-iteration stats tightly and the converged LL within distributional spread).
n_models=1path).Dependencies
Builds on #76 (single-model MLX). Part of epic #74 (Phase D).