Skip to content
Merged
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
103 changes: 60 additions & 43 deletions pertpy/tools/_augur.py
Original file line number Diff line number Diff line change
Expand Up @@ -518,7 +518,13 @@ def run_cross_validation(
categorical=True, random_state=42)
>>> results = ag_rfc.run_cross_validation(subsample=subsample, folds=3, subsample_idx=0, random_state=42, zero_division=0)
"""
x = subsample.to_df()
# Pass the dense matrix instead of subsample.to_df(); its arrow-backed string columns make scikit-learn re-validate dtypes on every fold and scorer, while the values (and results) stay identical.
x = subsample.X
if sparse.issparse(x):
x = x.toarray()
x = np.asarray(x)
genes = subsample.var_names.tolist()
n_genes = len(genes)
y = subsample.obs["y_"]
scorer = self.set_scorer(multiclass=len(y.unique()) > 2, zero_division=zero_division)
folds = StratifiedKFold(n_splits=folds, random_state=random_state, shuffle=True)
Expand All @@ -540,21 +546,21 @@ def run_cross_validation(
feature_importances = defaultdict(list)
if isinstance(self.estimator, RandomForestClassifier | RandomForestRegressor):
for fold, estimator in list(zip(range(len(results["estimator"])), results["estimator"], strict=False)):
feature_importances["genes"].extend(x.columns.tolist())
feature_importances["genes"].extend(genes)
feature_importances["feature_importances"].extend(estimator.feature_importances_.tolist())
feature_importances["subsample_idx"].extend(len(x.columns) * [subsample_idx])
feature_importances["fold"].extend(len(x.columns) * [fold])
feature_importances["subsample_idx"].extend(n_genes * [subsample_idx])
feature_importances["fold"].extend(n_genes * [fold])

# standardized coefficients with Agresti method
# cf. https://think-lab.github.io/d/205/#3
if isinstance(self.estimator, LogisticRegression):
for fold, estimator in list(zip(range(len(results["estimator"])), results["estimator"], strict=False)):
feature_importances["genes"].extend(x.columns.tolist())
feature_importances["genes"].extend(genes)
feature_importances["feature_importances"].extend(
(estimator.coef_ * estimator.coef_.std()).flatten().tolist()
)
feature_importances["subsample_idx"].extend(len(x.columns) * [subsample_idx])
feature_importances["fold"].extend(len(x.columns) * [fold])
feature_importances["subsample_idx"].extend(n_genes * [subsample_idx])
feature_importances["fold"].extend(n_genes * [fold])

results["feature_importances"] = feature_importances

Expand Down Expand Up @@ -795,7 +801,9 @@ def predict(
logger.warning("Set smaller span value in the case of a `segmentation fault` error.")
logger.warning("Set larger span in case of svddc or other near singularities error.")
adata.obs["augur_score"] = nan
for cell_type in track(adata.obs["cell_type"].unique(), description="Processing data..."):

eligible_cell_types = []
for cell_type in adata.obs["cell_type"].unique():
cell_type_subsample = adata[adata.obs["cell_type"] == cell_type].copy()

if augur_mode in ("default", "permute") and len(cell_type_subsample) >= min_cells:
Expand Down Expand Up @@ -825,42 +833,51 @@ def predict(
f"subsample size {subsample_size}."
)
else:
results[cell_type] = Parallel(n_jobs=n_threads)(
delayed(self.cross_validate_subsample)(
adata=cell_type_subsample,
augur_mode=augur_mode,
subsample_size=subsample_size,
folds=folds,
feature_perc=feature_perc,
subsample_idx=i,
random_state=random_state,
zero_division=zero_division,
)
for i in range(n_subsamples)
)
# summarize scores for cell type
results["summary_metrics"][cell_type] = self.average_metrics(results[cell_type])

# add scores as observation to anndata
mask = adata.obs["cell_type"].str.startswith(cell_type)
adata.obs.loc[mask, "augur_score"] = results["summary_metrics"][cell_type]["mean_augur_score"]

# concatenate feature importances for each subsample cv
subsample_feature_importances_dicts = [cv["feature_importances"] for cv in results[cell_type]]

for dictionary in subsample_feature_importances_dicts:
for key, value in dictionary.items():
results["feature_importances"][key].extend(value)
results["feature_importances"]["cell_type"].extend(
[cell_type]
* (len(results["feature_importances"]["genes"]) - len(results["feature_importances"]["cell_type"]))
)
eligible_cell_types.append((cell_type, cell_type_subsample))

# Fan subsamples out across all eligible cell types in one pool so workers don't idle at each cell-type boundary; submission order is preserved for regrouping.
flat_tasks = [(cell_type, sub, i) for cell_type, sub in eligible_cell_types for i in range(n_subsamples)]
cv_results = Parallel(n_jobs=n_threads, return_as="generator")(
delayed(self.cross_validate_subsample)(
adata=sub,
augur_mode=augur_mode,
subsample_size=subsample_size,
folds=folds,
feature_perc=feature_perc,
subsample_idx=i,
random_state=random_state,
zero_division=zero_division,
)
for _, sub, i in flat_tasks
)
cv_results = list(track(cv_results, total=len(flat_tasks), description="Processing data..."))

for offset, (cell_type, _) in enumerate(eligible_cell_types):
results[cell_type] = cv_results[offset * n_subsamples : (offset + 1) * n_subsamples]

# summarize scores for cell type
results["summary_metrics"][cell_type] = self.average_metrics(results[cell_type])

# add scores as observation to anndata
mask = adata.obs["cell_type"].str.startswith(cell_type)
adata.obs.loc[mask, "augur_score"] = results["summary_metrics"][cell_type]["mean_augur_score"]

# concatenate feature importances for each subsample cv
subsample_feature_importances_dicts = [cv["feature_importances"] for cv in results[cell_type]]

for dictionary in subsample_feature_importances_dicts:
for key, value in dictionary.items():
results["feature_importances"][key].extend(value)
results["feature_importances"]["cell_type"].extend(
[cell_type]
* (len(results["feature_importances"]["genes"]) - len(results["feature_importances"]["cell_type"]))
)

for idx, cv in zip(range(n_subsamples), results[cell_type], strict=False):
results["full_results"]["idx"].extend([idx] * folds)
results["full_results"]["augur_score"].extend(cv["test_augur_score"])
results["full_results"]["folds"].extend(range(folds))
results["full_results"]["cell_type"].extend([cell_type] * folds * n_subsamples)
for idx, cv in zip(range(n_subsamples), results[cell_type], strict=False):
results["full_results"]["idx"].extend([idx] * folds)
results["full_results"]["augur_score"].extend(cv["test_augur_score"])
results["full_results"]["folds"].extend(range(folds))
results["full_results"]["cell_type"].extend([cell_type] * folds * n_subsamples)
# make sure one cell type worked
if len(results) <= 2:
logger.warning("No cells types had more than min_cells needed. Please adjust data or min_cells parameter.")
Expand Down
Loading