Skip to content
Closed
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
280 changes: 279 additions & 1 deletion pf2rnaseq/factorization.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,29 @@ def correct_conditions(X: anndata.AnnData):
return X.uns["Pf2_A"] / counts_correct


def filter_components_by_cytokine_quantile(
threshold: float,
csv_path: str = "/home/nicoleb/Pf2-scRNAseq-1/pf2rnaseq/Data/donor_vs_cytokine_variability_sd.csv",
) -> np.ndarray:
"""Return 0-indexed components whose cytokine_quantile is at/above a threshold.

Parameters
----------
threshold : float
Minimum cytokine_quantile value a component must have to be kept.
csv_path : str
Path to the CSV with columns component, donor_sd, cytokine_sd,
donor_quantile, cytokine_quantile.

Returns
-------
np.ndarray
0-indexed component numbers with cytokine_quantile >= threshold.
"""
Comment on lines +46 to +64
df = pd.read_csv(csv_path)
return df.loc[df["cytokine_quantile"] >= threshold, "component"].to_numpy()


def pf2(
X: anndata.AnnData,
rank: int,
Expand Down Expand Up @@ -288,7 +311,11 @@ def gradient(x):

# ===== Gradient w.r.t. W =====
# 1. Reconstruction term: ∂/∂W [||A - WH||²] = 2(error @ H^T), L1 penalty: ∂/∂W [α||W||₁] = α * sign(W)
grad_W = 2 * ((W @ H - A) @ H.T) + alpha * np.sign(W) - np.diag(alpha * np.sign(np.diag(W)))
grad_W = (
2 * ((W @ H - A) @ H.T)
+ alpha * np.sign(W)
- np.diag(alpha * np.sign(np.diag(W)))
)

# ===== Gradient w.r.t. H =====
# 1. Reconstruction term: ∂/∂H [||A - WH||²] = 2(W^T @ error), L1 penalty: ∂/∂H [α||H||₁] = α * sign(H)
Expand Down Expand Up @@ -346,3 +373,254 @@ def gradient(x):
print(f" Non-zeros: {np.sum(np.abs(H) > 1e-3)}/{H.size}")

return W, H


def deconvolution_cytokine_admm(
A: np.ndarray,
alpha_h: float = 0.1,
alpha_w: float = 0.01,
rho_w_init: float | None = None,
rho_h_init: float | None = None,
max_iter: int = 10000,
tol_abs: float = 1e-4,
tol_rel: float = 1e-3,
random_state: int = 1,
adaptive_rho: bool = True,
rho_bounds: tuple[float, float] = (1e-4, 1e4),
non_negative_w: bool = True,
non_negative_h: bool = True,
) -> tuple[np.ndarray, np.ndarray, dict]:
"""
Decompose cytokine factor matrix using ADMM: A ≈ W @ H

Parameters
----------
A : np.ndarray
Input matrix (n_cytokines, n_components)
alpha_h : float
L1 regularization for H
alpha_w : float
L1 regularization for W (off-diagonal only)
rho_w_init, rho_h_init : float or None
Initial ADMM penalty for the W- and H-subproblems. If None,
computed via the spectral rule rho* = sqrt(lambda_min * lambda_max)
of the relevant Gram matrix (H H^T for W-step, W^T W for H-step).
max_iter : int
Maximum iterations
tol_abs, tol_rel : float
Absolute and relative tolerance terms for the combined stopping
criterion (Boyd et al. 2011, §3.3.1), applied per-block.
random_state : int
Random seed
adaptive_rho : bool
Whether to adaptively adjust rho_w and rho_h independently
rho_bounds : tuple[float, float]
(min, max) clip range for adaptive rho updates
non_negative_w : bool
If True, enforce W ≥ 0 (cytokines only activate, not inhibit)
non_negative_h : bool
If True, enforce H ≥ 0

Returns
-------
Z_W : np.ndarray
Cytokine interaction matrix (n_cytokines, n_cytokines)
Z_H : np.ndarray
Effect basis matrix (n_cytokines, n_components)
history : dict
Optimization history
"""
n_cytokines, n_components = A.shape
np.random.seed(random_state)

# Initialize
W = np.random.rand(n_cytokines, n_cytokines) * 0.1 + np.eye(n_cytokines)
H = np.random.rand(n_cytokines, n_components) * np.mean(np.abs(A))
Z_W = W.copy()
Z_H = H.copy()
U_W = np.zeros_like(W)
U_H = np.zeros_like(H)

off_diag_mask = ~np.eye(n_cytokines, dtype=bool)
rho_min, rho_max = rho_bounds

def soft_threshold(X, threshold):
return np.sign(X) * np.maximum(np.abs(X) - threshold, 0)

def spectral_rho(M):
"""Geometric mean of the extreme eigenvalues of M @ M.T."""
eigs = np.linalg.eigvalsh(M @ M.T)
eigs = eigs[eigs > 1e-12]
if eigs.size == 0:
return 1.0
return float(np.sqrt(eigs.min() * eigs.max()))

# --- Initialize rho_w, rho_h independently ---
rho_w = rho_w_init if rho_w_init is not None else spectral_rho(H)
rho_h = rho_h_init if rho_h_init is not None else spectral_rho(W)
rho_w = np.clip(rho_w, rho_min, rho_max)
rho_h = np.clip(rho_h, rho_min, rho_max)

def update_W(H, Z_W, U_W, rho_w):
lhs = H @ H.T + rho_w * np.eye(n_cytokines)
rhs = A @ H.T + rho_w * (Z_W - U_W)
return np.linalg.solve(lhs, rhs.T).T

def update_H(W, Z_H, U_H, rho_h):
W_TW = W.T @ W
W_TA = W.T @ A
lhs = W_TW + rho_h * np.eye(n_cytokines)
rhs = W_TA + rho_h * (Z_H - U_H)
return np.linalg.solve(lhs, rhs)

def update_Z_W(W, U_W, alpha, rho_w):
Z_W_new = soft_threshold(W + U_W, alpha / rho_w)
if non_negative_w:
np.maximum(Z_W_new, 0, out=Z_W_new)
np.fill_diagonal(Z_W_new, 1.0)
return Z_W_new

def update_Z_H(H, U_H, alpha, rho_h):
Z_H_new = soft_threshold(H + U_H, alpha / rho_h)
if non_negative_h:
np.maximum(Z_H_new, 0, out=Z_H_new)
return Z_H_new

history = {
"objective": [],
"primal_residual_w": [],
"dual_residual_w": [],
"primal_residual_h": [],
"dual_residual_h": [],
"rho_w": [],
"rho_h": [],
"w_sparsity": [],
"h_sparsity": [],
}

print("Cytokine deconvolution with ADMM:")
print(f" A shape: {A.shape}")
print(f" Alpha_W: {alpha_w}, Alpha_H: {alpha_h}")
print(f" Initial rho_w: {rho_w:.4g}, rho_h: {rho_h:.4g}")
print(f" Tolerance (abs, rel): ({tol_abs:.1e}, {tol_rel:.1e})")
print(f" Non-negative W: {non_negative_w}")
print(f" Non-negative H: {non_negative_h}")
print("\nStarting ADMM iterations...")

for iteration in range(max_iter):
Z_W_old = Z_W.copy()
Z_H_old = Z_H.copy()

# ADMM updates
W = update_W(H, Z_W, U_W, rho_w)
H = update_H(W, Z_H, U_H, rho_h)
Z_W = update_Z_W(W, U_W, alpha_w, rho_w)
Z_H = update_Z_H(H, U_H, alpha_h, rho_h)
U_W = U_W + (W - Z_W)
U_H = U_H + (H - Z_H)

# Per-block residuals, normalized by entry count
r_w = np.linalg.norm(W - Z_W) / np.sqrt(W.size)
s_w = np.linalg.norm(rho_w * (Z_W - Z_W_old)) / np.sqrt(W.size)
r_h = np.linalg.norm(H - Z_H) / np.sqrt(H.size)
s_h = np.linalg.norm(rho_h * (Z_H - Z_H_old)) / np.sqrt(H.size)

# Combined absolute + relative stopping criteria (Boyd et al. 2011, §3.3.1)
eps_pri_w = tol_abs + tol_rel * max(
np.linalg.norm(W) / np.sqrt(W.size), np.linalg.norm(Z_W) / np.sqrt(W.size)
)
eps_dual_w = tol_abs + tol_rel * np.linalg.norm(rho_w * U_W) / np.sqrt(W.size)
eps_pri_h = tol_abs + tol_rel * max(
np.linalg.norm(H) / np.sqrt(H.size), np.linalg.norm(Z_H) / np.sqrt(H.size)
)
eps_dual_h = tol_abs + tol_rel * np.linalg.norm(rho_h * U_H) / np.sqrt(H.size)

# Compute objective
recon_error = np.sum((A - W @ H) ** 2)
l1_W = alpha_w * np.sum(np.abs(Z_W[off_diag_mask]))
l1_H = alpha_h * np.sum(np.abs(Z_H))
objective = recon_error + l1_W + l1_H

# Track sparsity
w_sparsity = np.sum(np.abs(Z_W[off_diag_mask]) < 1e-3) / np.sum(off_diag_mask)
h_sparsity = np.sum(np.abs(Z_H) < 1e-3) / Z_H.size

# Store history
history["objective"].append(objective)
history["primal_residual_w"].append(r_w)
history["dual_residual_w"].append(s_w)
history["primal_residual_h"].append(r_h)
history["dual_residual_h"].append(s_h)
history["rho_w"].append(rho_w)
history["rho_h"].append(rho_h)
history["w_sparsity"].append(w_sparsity)
history["h_sparsity"].append(h_sparsity)

if iteration % 100 == 0 or iteration < 10:
print(
f" Iter {iteration:4d}: Obj={objective:.4e}, "
f"r_w={r_w:.3e}, s_w={s_w:.3e}, ρ_w={rho_w:.3g} | "
f"r_h={r_h:.3e}, s_h={s_h:.3e}, ρ_h={rho_h:.3g}"
)

# Adaptive rho update — W and H blocks handled independently
if adaptive_rho and iteration > 0:
if r_w > 10 * s_w:
rho_w = np.clip(rho_w * 2, rho_min, rho_max)
U_W = U_W / 2
elif s_w > 10 * r_w:
rho_w = np.clip(rho_w / 2, rho_min, rho_max)
U_W = U_W * 2

if r_h > 10 * s_h:
rho_h = np.clip(rho_h * 2, rho_min, rho_max)
U_H = U_H / 2
elif s_h > 10 * r_h:
rho_h = np.clip(rho_h / 2, rho_min, rho_max)
U_H = U_H * 2

# Convergence check — both blocks must satisfy their combined criteria
if (
r_w < eps_pri_w
and s_w < eps_dual_w
and r_h < eps_pri_h
and s_h < eps_dual_h
):
print(f"\n✓ Converged at iteration {iteration}")
print(
f" W block: r={r_w:.4e} < {eps_pri_w:.4e}, s={s_w:.4e} < {eps_dual_w:.4e}"
)
print(
f" H block: r={r_h:.4e} < {eps_pri_h:.4e}, s={s_h:.4e} < {eps_dual_h:.4e}"
)
break

# Final statistics
A_recon = W @ H
rel_error = np.linalg.norm(A - A_recon, "fro") / np.linalg.norm(A, "fro")

w_sparsity = np.sum(np.abs(Z_W[off_diag_mask]) < 1e-3) / np.sum(off_diag_mask)
h_sparsity = np.sum(np.abs(Z_H) < 1e-3) / Z_H.size

print("\nOptimization complete:")
print(f" Iterations: {iteration + 1}/{max_iter}")
print(f" Relative reconstruction error: {rel_error:.4%}")
print(f" Final rho_w: {rho_w:.4g}, rho_h: {rho_h:.4g}")

print("\n W (cytokine interactions):")
print(f" Off-diagonal sparsity: {w_sparsity:.2%}")
print(f" Off-diagonal non-zeros: {np.sum(np.abs(Z_W[off_diag_mask]) > 1e-3)}")
print(f" Mean |W_offdiag|: {np.abs(Z_W[off_diag_mask]).mean():.4f}")
print(f" Min value: {W.min():.4f}")
print(f" Max value: {W.max():.4f}")
print(f" Diagonal: all 1.0 (constrained)")
Comment on lines +614 to +616

print("\n H (effect patterns):")
print(f" Sparsity: {h_sparsity:.2%}")
print(f" Non-zeros: {np.sum(np.abs(Z_H) > 1e-3)}/{Z_H.size}")
print(f" Mean |H|: {np.abs(Z_H).mean():.4f}")
print(f" Min value: {H.min():.4f}")
print(f" Max value: {H.max():.4f}")
print(f" Negative values: {np.sum(H < 0)} ({100 * np.sum(H < 0) / H.size:.1f}%)")
Comment on lines +622 to +624

return Z_W, Z_H, history
93 changes: 93 additions & 0 deletions pf2rnaseq/figures/figureParseADMM.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
"""
Parse data: Plotting factors
"""

import numpy as np
import pandas as pd
import seaborn as sns
from anndata import read_h5ad
from matplotlib import pyplot as plt

from ..factorization import correct_conditions, deconvolution_cytokine_admm
from .common import getSetup, subplotLabel
from .commonFuncs.plotFactors import (
plot_condition_factors,
)


def samples_only(X) -> pd.DataFrame:
"""Obtain samples once only with corresponding observations"""
samples = X.obs
df_samples = samples.drop_duplicates(subset="condition_unique_idxs")
df_samples = df_samples.sort_values("condition_unique_idxs")
return df_samples


def makeFigure():
"""Get a list of the axis objects and create a figure."""
# Get list of axis objects
ax, f = getSetup((25, 15), (1, 3))

# Add subplot labels
subplotLabel(ax)

# Load data
X = read_h5ad("/home/nicoleb/ParsePf2_100_D11_filt.h5ad")
X.uns["Pf2_A"] = correct_conditions(X)
A = X.uns["Pf2_A"]

# Center A by cytokine medians
cytokine_medians = np.median(A, axis=1, keepdims=True)
A_centered = A - cytokine_medians
X.uns["Pf2_A"] = A_centered

W, H, _ = deconvolution_cytokine_admm(A_centered, alpha_h=0.05, alpha_w=0.05, rho=2)

# Get cytokine names in correct order
samples_df = samples_only(X)

# Create deconvolved version for plotting
X_deconv = X.copy()
X_deconv.uns["Pf2_A"] = H # Use primary effects only

plot_condition_factors(
X_deconv,
ax[0],
samples_df["cytokine"],
groupConditions=True,
cond="cytokine",
log_scale=False,
)
ax[0].set_title("Deconvolved matrix (H)", fontsize=12, fontweight="bold")

#Plot original median subtracted factor matrix for reference
plot_condition_factors(
X,
ax[1],
samples_df["cytokine"],
groupConditions=True,
cond="cytokine",
log_scale=False,
)
ax[1].set_title("Original Effects (A)", fontsize=12, fontweight="bold")

cytokine_names = samples_df["cytokine"].values

# Plot 2: W heatmap (primary effects)
sns.heatmap(
W,
ax=ax[2],
cmap="YlOrRd",
robust=False,
square=True,
cbar_kws={"label": "Signaling Strength"},
xticklabels=cytokine_names,
yticklabels=cytokine_names,
)
ax[2].set_title("Cytokine Signaling (W)", fontsize=12, fontweight="bold")
ax[2].set_xlabel("Inducing Cytokine →", fontsize=10)
ax[2].set_ylabel("← Induced Cytokine", fontsize=10)
plt.setp(ax[2].get_xticklabels(), rotation=90, ha="center", fontsize=6)
plt.setp(ax[2].get_yticklabels(), rotation=0, fontsize=6)

return f
Loading