Skip to content

Phase D: Multi-model support for AMICAMLXNG (Apple GPU) #81

Description

@neuromechanist

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 > 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).
  • Single-model (Phase C: MLX backend port for AMICATorchNG (Apple-native GPU) #76) stays byte-for-byte unchanged (guard the n_models=1 path).
  • The Phase B benchmark's multi-model configs then run MLX too (drop the single-model gate).

Dependencies

Builds on #76 (single-model MLX). Part of epic #74 (Phase D).

Activity

  1. neuromechanist commented on Jul 8, 2026

    @neuromechanist
    MemberAuthor

    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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    featureNew feature or enhancement

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions