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
140 changes: 140 additions & 0 deletions py_hdWGCNA/_kernels.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
"""
Optimized compute kernels for py-hdWGCNA.

Strategy:
- Use numpy BLAS for matrix multiply (adj @ adj) - already optimal
- Use Numba for element-wise operations that are slow in pure Python
- Falls back gracefully to pure numpy if numba is not available
"""

from __future__ import annotations

import numpy as np

try:
from numba import njit, prange

HAS_NUMBA = True
except ImportError:
HAS_NUMBA = False


def _rank_rows(x: np.ndarray) -> np.ndarray:
"""Rank each row of a 2D array (for Spearman correlation)."""
from scipy.stats import rankdata

return rankdata(x, axis=1).astype(np.float64)


def compute_tom_numba(
adj_matrix: np.ndarray,
tom_type: str = "signed",
tom_denom: str = "min",
parallel: bool = True,
) -> np.ndarray:
"""Compute TOM using optimized kernel.

Uses numpy BLAS for the critical matrix multiply (adj @ adj),
which is already multi-threaded and SIMD-optimized.

Parameters
----------
adj_matrix : np.ndarray
Adjacency matrix
tom_type : str
'signed' or 'unsigned'
tom_denom : str
'min' or 'max'
parallel : bool
Unused (BLAS handles threading internally)

Returns
-------
np.ndarray
TOM dissimilarity matrix
"""
adj_work = adj_matrix.astype(np.float64, copy=True)
np.fill_diagonal(adj_work, 0.0)

k_i = adj_work.sum(axis=1)

if tom_denom == "min":
denominator = np.minimum(k_i[:, np.newaxis], k_i[np.newaxis, :])
else:
denominator = np.maximum(k_i[:, np.newaxis], k_i[np.newaxis, :])

denominator += 1.0
denominator -= adj_work
denominator[denominator < 1e-6] = 1e-6

# Critical path: matrix multiply (numpy BLAS, already optimal)
num = adj_work @ adj_work + adj_work
np.fill_diagonal(num, 0.0)

TOM = num / denominator
np.fill_diagonal(TOM, 1.0)
np.clip(TOM, 0.0, 1.0, out=TOM)

dissTOM = 1.0 - TOM
np.fill_diagonal(dissTOM, 0.0)
np.clip(dissTOM, 0.0, 1.0, out=dissTOM)

return dissTOM


def compute_correlation_numba(
expr_mat: np.ndarray,
method: str = "pearson",
parallel: bool = True,
) -> np.ndarray:
"""Compute correlation matrix using optimized kernel.

For Pearson: uses numpy BLAS (already optimal).
For Spearman: uses vectorized scipy.stats.rankdata.

Parameters
----------
expr_mat : np.ndarray
Genes x Samples matrix
method : str
'pearson', 'spearman', or 'bicor'
parallel : bool
Unused (BLAS handles threading internally)

Returns
-------
np.ndarray
Correlation matrix
"""
x = np.asarray(expr_mat, dtype=np.float64)

if method == "spearman":
x = _rank_rows(x)

if method == "bicor":
from .utils import _bicor_vectorized

return _bicor_vectorized(x)

# Handle NaN
nan_mask = np.isnan(x)
if nan_mask.any():
x = x.copy()
col_means = np.nanmean(x, axis=1, keepdims=True)
for i in range(x.shape[0]):
x[i, nan_mask[i]] = col_means[i, 0]

# Pearson: use numpy BLAS (already optimal)
n_genes, n_samples = x.shape
means = np.mean(x, axis=1, keepdims=True)
stds = np.std(x, axis=1, ddof=1, keepdims=True)
stds = np.where(stds < 1e-15, 1.0, stds)
centered = x - means
normalized = centered / stds
cor_matrix = (normalized @ normalized.T) / (n_samples - 1)
return np.clip(cor_matrix, -1.0, 1.0)


def is_numba_available() -> bool:
"""Check if Numba JIT is available."""
return HAS_NUMBA
24 changes: 22 additions & 2 deletions py_hdWGCNA/hdWGCNA.py
Original file line number Diff line number Diff line change
Expand Up @@ -292,6 +292,7 @@ def test_soft_powers(
power_range: list = None,
network_type: str = "signed",
cor_method: str = "bicor",
n_threads: int | None = None,
wgcna_name: str = None,
):
"""Test different soft-thresholding powers for scale-free topology fit.
Expand All @@ -304,6 +305,8 @@ def test_soft_powers(
'signed', 'unsigned', or 'signed hybrid'
cor_method : str
Correlation method: 'bicor', 'pearson', or 'spearman'
n_threads : int or None
Number of threads for parallel computation. None = auto.
wgcna_name : str
Name of hdWGCNA experiment

Expand All @@ -317,6 +320,7 @@ def test_soft_powers(
power_range=power_range,
network_type=network_type,
cor_method=cor_method,
n_threads=n_threads,
wgcna_name=wgcna_name,
)
return self
Expand All @@ -333,6 +337,7 @@ def construct_network(
pamRespectsDendro: bool = True,
pamStage: bool = False,
mergeCutHeight: float = 0.2,
n_threads: int | None = None,
wgcna_name: str = None,
**kwargs,
):
Expand Down Expand Up @@ -360,6 +365,8 @@ def construct_network(
Whether to perform PAM stage
mergeCutHeight : float
Cut height for merging
n_threads : int or None
Number of threads for parallel computation. None = auto.
wgcna_name : str
Name of hdWGCNA experiment
**kwargs
Expand All @@ -382,6 +389,7 @@ def construct_network(
pamRespectsDendro=pamRespectsDendro,
pamStage=pamStage,
mergeCutHeight=mergeCutHeight,
n_threads=n_threads,
wgcna_name=wgcna_name,
**kwargs,
)
Expand Down Expand Up @@ -722,7 +730,9 @@ def generate_motif_data(self, n_tfs=100, density=0.05, seed=42, wgcna_name=None)
)
return self

def construct_tf_network(self, model_params=None, nfold=5, wgcna_name=None):
def construct_tf_network(
self, model_params=None, nfold=5, n_threads=None, wgcna_name=None
):
"""Construct directed TF-gene network using XGBoost.

Parameters
Expand All @@ -731,6 +741,8 @@ def construct_tf_network(self, model_params=None, nfold=5, wgcna_name=None):
XGBoost parameters
nfold : int
CV folds
n_threads : int or None
Number of threads for parallel gene processing. None = auto.
wgcna_name : str
Experiment name

Expand All @@ -742,7 +754,11 @@ def construct_tf_network(self, model_params=None, nfold=5, wgcna_name=None):
from .tf_network import construct_tf_network as _ctf

self.adata = _ctf(
self.adata, model_params=model_params, nfold=nfold, wgcna_name=wgcna_name
self.adata,
model_params=model_params,
nfold=nfold,
n_threads=n_threads,
wgcna_name=wgcna_name,
)
return self

Expand Down Expand Up @@ -786,6 +802,7 @@ def regulon_scores(
target_type="positive",
cor_thresh=0.05,
exclude_grey_genes=True,
n_threads=None,
wgcna_name=None,
):
"""Compute regulon activity scores.
Expand All @@ -798,6 +815,8 @@ def regulon_scores(
Correlation threshold
exclude_grey_genes : bool
Exclude grey module genes
n_threads : int or None
Number of threads for parallel computation. None = auto.
wgcna_name : str
Experiment name

Expand All @@ -813,6 +832,7 @@ def regulon_scores(
target_type=target_type,
cor_thresh=cor_thresh,
exclude_grey_genes=exclude_grey_genes,
n_threads=n_threads,
wgcna_name=wgcna_name,
)
return self
Expand Down
62 changes: 45 additions & 17 deletions py_hdWGCNA/network.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ def test_soft_powers(
power_range: list = None,
network_type: str = "signed",
cor_method: str = "pearson",
n_threads: int | None = None,
wgcna_name: str = None,
**kwargs,
):
Expand All @@ -33,6 +34,8 @@ def test_soft_powers(
Network type: 'signed', 'unsigned', or 'signed hybrid'
cor_method : str
Correlation method: 'bicor', 'pearson', or 'spearman'
n_threads : int or None
Number of threads for parallel power testing. None = auto.
wgcna_name : str
Name of hdWGCNA experiment
**kwargs
Expand Down Expand Up @@ -102,26 +105,45 @@ def test_soft_powers(
results_list = []

chunk_size = 100
for power in power_range:

def _compute_single_power(power):
k = np.zeros(n_genes, dtype=np.float64)
for ci in range(0, n_genes, chunk_size):
ci_end = min(ci + chunk_size, n_genes)
chunk = log_cor[ci:ci_end, :]
k[ci:ci_end] = np.exp(power * chunk).sum(axis=1)

sft = scale_free_fit_index_full(k, nBreaks=10)
return {
"Power": power,
"SFT.R.sq": sft["R2"],
"slope": sft["slope"],
"truncated.R.sq": sft["truncated_R2"],
"mean.k.": float(np.mean(k)),
"median.k.": float(np.median(k)),
"max.k.": float(np.max(k)),
}

results_list.append(
{
"Power": power,
"SFT.R.sq": sft["R2"],
"slope": sft["slope"],
"truncated.R.sq": sft["truncated_R2"],
"mean.k.": float(np.mean(k)),
"median.k.": float(np.median(k)),
"max.k.": float(np.max(k)),
}
)
# Parallelize across power values using threads
# (numpy releases GIL during matrix operations)
from .parallel import _get_max_workers
from concurrent.futures import ThreadPoolExecutor

n_workers = _get_max_workers(n_threads)
if n_workers > 1 and len(power_range) > 1:
with ThreadPoolExecutor(max_workers=n_workers) as executor:
results_list = list(executor.map(_compute_single_power, power_range))
else:
for power in power_range:
results_list.append(_compute_single_power(power))

for r in results_list:
r.setdefault("Power", 0)
r.setdefault("SFT.R.sq", 0.0)
r.setdefault("slope", 0.0)
r.setdefault("truncated.R.sq", 0.0)
r.setdefault("mean.k.", 0.0)
r.setdefault("median.k.", 0.0)
r.setdefault("max.k.", 0.0)

power_table = pd.DataFrame(results_list)

Expand Down Expand Up @@ -167,7 +189,7 @@ def construct_network(
detectCutHeight: float = 0.995,
minKMEtoStay: float = 0,
mergeCutHeight: float = 0.2,
n_threads: int = 1,
n_threads: int | None = None,
verbose: int = 3,
saveTOMs: bool = False,
loadTOMs: bool = False,
Expand Down Expand Up @@ -270,7 +292,9 @@ def construct_network(

cor_method = wgcna_data.get("cor_method", "pearson")
print(f"Computing correlation matrix ({cor_method})...")
cor_matrix = compute_correlation_matrix(dat_expr, method=cor_method)
cor_matrix = compute_correlation_matrix(
dat_expr, method=cor_method, n_threads=n_threads
)
np.clip(cor_matrix, -1, 1, out=cor_matrix)
print(
f" Cor matrix shape: {cor_matrix.shape}, range=[{cor_matrix.min():.4f}, {cor_matrix.max():.4f}]"
Expand All @@ -285,7 +309,9 @@ def construct_network(
del cor_matrix

print("Computing Topological Overlap Matrix (TOM)...")
tom_dissim = compute_tom(adj_matrix, tom_type=tom_type, tom_denom=tom_denom)
tom_dissim = compute_tom(
adj_matrix, tom_type=tom_type, tom_denom=tom_denom, n_threads=n_threads
)
print(
f" TOM dissim shape: {tom_dissim.shape}, range=[{tom_dissim.min():.4f}, {tom_dissim.max():.4f}]"
)
Expand Down Expand Up @@ -361,7 +387,9 @@ def construct_network(
}
)

kME_all = compute_kme(dat_expr, MEs_merged, method=cor_method)
kME_all = compute_kme(
dat_expr, MEs_merged, method=cor_method, n_threads=n_threads
)

kME_cols = {}

Expand Down
Loading
Loading