Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/cnmf/cnmf.py
Original file line number Diff line number Diff line change
Expand Up @@ -1267,7 +1267,7 @@ def main():

args = parser.parse_args()
try:
engine_commands = ('factorize',)
engine_commands = ('factorize', 'consensus')
validate_engine_args_for_command(args, engine_commands)
except ValueError as e:
parser.error(str(e))
Expand Down
109 changes: 88 additions & 21 deletions src/cnmf/nmf_gpu.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,8 @@

Standalone (no cnmf dependency); exposes a core NMF factorization and a thin `cNMF._nmf` adapter.

nmf_kwargs (from cnmf): reads n_components, max_iter, tol, random_state, init.
nmf_kwargs (from cnmf): reads n_components, max_iter, tol, random_state, init, and
consensus refit args H/update_H.
Ignores beta_loss/solver/alpha_W/alpha_H/l1_ratio by design — always Frobenius MU, no regularization
(matches cnmf defaults; would diverge if a run enables regularization). Defaults are centralized in
`DEFAULT_NMF` and `DEFAULT_GPU`. Init runs through sklearn's _initialize_nmf for parity
Expand Down Expand Up @@ -72,14 +73,14 @@
def parse_gpu_args(parser):
"""Register cNMF CLI flags for the optional PyTorch GPU NMF engine."""
group = parser.add_argument_group("NMF engine options")
group.add_argument("--engine", type=str.lower, choices=["cpu", "gpu"], help="[factorize] NMF engine to use (default cpu)")
group.add_argument("--gpu-device", type=str, help="[factorize,gpu] Device for GPU NMF: auto, cpu, cuda, cuda:N, or mps")
group.add_argument("--gpu-dtype", type=str.lower, choices=["auto", "fp32", "fp64", "bf16"], help="[factorize,gpu] Storage and matmul dtype for GPU NMF (default auto)")
group.add_argument("--gpu-allow-tf32", action="store_const", const=True, help="[factorize,gpu] Allow TF32 for CUDA fp32 matrix multiplication")
group.add_argument("--gpu-compile", action="store_const", const=True, help="[factorize,gpu] Enable torch.compile for the GPU NMF update step")
group.add_argument("--gpu-eps", type=float, help="[factorize,gpu] Multiplicative-update denominator guard")
group.add_argument("--gpu-check-every", type=int, help="[factorize,gpu] Eager-mode convergence check interval")
group.add_argument("--gpu-compile-block", type=int, help="[factorize,gpu] Number of MU iterations per compiled block")
group.add_argument("--engine", type=str.lower, choices=["cpu", "gpu"], help="[factorize,consensus] NMF engine to use (default cpu)")
group.add_argument("--gpu-device", type=str, help="[factorize,consensus,gpu] Device for GPU NMF: auto, cpu, cuda, cuda:N, or mps")
group.add_argument("--gpu-dtype", type=str.lower, choices=["auto", "fp32", "fp64", "bf16"], help="[factorize,consensus,gpu] Storage and matmul dtype for GPU NMF (default auto)")
group.add_argument("--gpu-allow-tf32", action="store_const", const=True, help="[factorize,consensus,gpu] Allow TF32 for CUDA fp32 matrix multiplication")
group.add_argument("--gpu-compile", action="store_const", const=True, help="[factorize,consensus,gpu] Enable torch.compile for the GPU NMF update step")
group.add_argument("--gpu-eps", type=float, help="[factorize,consensus,gpu] Multiplicative-update denominator guard")
group.add_argument("--gpu-check-every", type=int, help="[factorize,consensus,gpu] Eager-mode convergence check interval")
group.add_argument("--gpu-compile-block", type=int, help="[factorize,consensus,gpu] Number of MU iterations per compiled block")
return parser


Expand Down Expand Up @@ -326,7 +327,12 @@ def _mu_step(W, H, Xg, eps):
return W, H


def _execution_plan(torch, opt, device):
def _mu_step_fixed_h(W, H, Xg, eps):
"""One Frobenius MU update for consensus refit: keep H fixed and update W only."""
return W * ((Xg @ H.T) / (W @ (H @ H.T) + eps))


def _execution_plan(torch, opt, device, step_fn=_mu_step):
"""Resolve compile policy and convergence-check cadence in one place.

MPS does not support torch.compile reliably for this path, so compile requests are treated as
Expand All @@ -336,8 +342,8 @@ def _execution_plan(torch, opt, device):
"""
use_compile = opt["compile"] and not device.startswith("mps")
if use_compile:
return torch.compile(_mu_step), opt["compile_block"]
return _mu_step, opt["check_every"]
return torch.compile(step_fn), opt["compile_block"]
return step_fn, opt["check_every"]


def _to_device_factors(torch, W0, H0, dtype, device):
Expand All @@ -359,6 +365,22 @@ def _check_runtime_tensors(Xg, W, H, eps):
raise RuntimeError(f"NMF runtime tensors must share dtype; got {sorted(map(str, dtypes))}")


def _to_checked_fixed_h(H, k, n_features):
"""Materialize and validate fixed H for update_H=False consensus refits."""
if H is None:
raise ValueError("update_H=False requires a fixed H matrix")
Hnp = H.toarray() if hasattr(H, "toarray") else np.asarray(H)
if Hnp.ndim != 2:
raise ValueError("fixed H must be a 2D matrix")
if Hnp.shape != (k, n_features):
raise ValueError(f"fixed H shape must be ({k}, {n_features}); got {Hnp.shape}")
if not np.isfinite(Hnp).all():
raise ValueError("fixed H contains NaN/inf")
if Hnp.size and Hnp.min() < 0:
raise ValueError(f"NMF requires non-negative fixed H; found min(H) = {float(Hnp.min()):.4g}")
return Hnp


def _fit_mu(torch, Xg, W, H, eps, max_iter, tol, step, block, tf32, device):
"""Run Frobenius MU updates until convergence or max_iter."""
_check_runtime_tensors(Xg, W, H, eps)
Expand All @@ -382,6 +404,28 @@ def _fit_mu(torch, Xg, W, H, eps, max_iter, tol, step, block, tf32, device):
return W, H


def _fit_mu_fixed_h(torch, Xg, W, H, eps, max_iter, tol, step, block, tf32, device):
"""Run Frobenius MU updates for W with H fixed."""
_check_runtime_tensors(Xg, W, H, eps)
err_init = prev_err = None
with torch.no_grad(), _cuda_tf32(torch, tf32, device):
it = 0
while it < max_iter:
n = min(block, max_iter - it)
for _ in range(n):
W = step(W, H, Xg, eps)
it += n
err = float(torch.linalg.norm(Xg - W @ H))
if err_init is None:
err_init = err
if err_init == 0.0:
break
elif prev_err is not None and (prev_err - err) / err_init < tol:
break
prev_err = err
return W


# ---------------------------------------------------------------------
# Public API and adapters
# ---------------------------------------------------------------------
Expand All @@ -390,6 +434,33 @@ def _to_nmf_output(H, W):
return H.cpu().double().numpy(), W.cpu().double().numpy()


def _factorize_nmf_gpu_fixed_h(torch, Xnp, Xg, nmf_kwargs, opt, device, dtype, k, max_iter, tol, eps, seed):
"""Consensus refit path: use fixed H and update W only."""
Hnp = _to_checked_fixed_h(nmf_kwargs.get("H"), k, Xnp.shape[1])

W0, _ = _init_wh(Xnp, k, seed, nmf_kwargs.get("init"))
W, H = _to_device_factors(torch, W0, Hnp, dtype, device)

step, block = _execution_plan(torch, opt, device, _mu_step_fixed_h)
tf32 = device.startswith("cuda") and opt["allow_tf32"] and dtype is torch.float32

W = _fit_mu_fixed_h(torch, Xg, W, H, eps, max_iter, tol, step, block, tf32, device)
return _to_nmf_output(H, W)


def _factorize_nmf_gpu_full(torch, Xnp, Xg, nmf_kwargs, opt, device, dtype, k, max_iter, tol, eps, seed):
"""Normal NMF path: initialize and update both W and H."""
# sklearn-parity init on CPU (_initialize_nmf), then move to device (cnmf uses init='random')
W0, H0 = _init_wh(Xnp, k, seed, nmf_kwargs.get("init"))
W, H = _to_device_factors(torch, W0, H0, dtype, device)

step, block = _execution_plan(torch, opt, device)
tf32 = device.startswith("cuda") and opt["allow_tf32"] and dtype is torch.float32

W, H = _fit_mu(torch, Xg, W, H, eps, max_iter, tol, step, block, tf32, device)
return _to_nmf_output(H, W)


def factorize_nmf_gpu(X, nmf_kwargs, gpu_kwargs=None):
"""Standalone Frobenius MU NMF: X -> (spectra, usages) = (H, W).

Expand Down Expand Up @@ -418,15 +489,11 @@ def factorize_nmf_gpu(X, nmf_kwargs, gpu_kwargs=None):
# should stream X blocks to GPU/MPS instead of requiring full dense X in VRAM.
Xg = torch.as_tensor(Xnp, dtype=dtype, device=device) # cells x genes

# sklearn-parity init on CPU (_initialize_nmf), then move to device (cnmf uses init='random')
W0, H0 = _init_wh(Xnp, k, seed, nmf_kwargs.get("init"))
W, H = _to_device_factors(torch, W0, H0, dtype, device)

step, block = _execution_plan(torch, opt, device)
tf32 = device.startswith("cuda") and opt["allow_tf32"] and dtype is torch.float32

W, H = _fit_mu(torch, Xg, W, H, eps, max_iter, tol, step, block, tf32, device)
return _to_nmf_output(H, W)
if nmf_kwargs.get("update_H", True) is False:
return _factorize_nmf_gpu_fixed_h(torch, Xnp, Xg, nmf_kwargs, opt, device, dtype,
k, max_iter, tol, eps, seed)
return _factorize_nmf_gpu_full(torch, Xnp, Xg, nmf_kwargs, opt, device, dtype,
k, max_iter, tol, eps, seed)


def _nmf_gpu(self, X, nmf_kwargs):
Expand Down
Loading