Paper2Agent_decoupleRpy / src /workflows /manifest_data_validation.py
Annie Voigt
style: apply ruff lint --fix + ruff format across the tree
c3b49d6
Raw
History Blame Contribute Delete
21.4 kB
"""
Validates a loaded dataset against its manifest declaration.
DESIGN
------
All functions accept plain numpy arrays and pandas DataFrames — no scanpy or
AnnData imports. The MCP tool (dataset_validate_manifest_against_data in
dataset_tools.py) handles h5ad loading and calls validate_manifest_against_data.
Five check categories
---------------------
data_level Detected expression type vs. manifest declaration.
feature_id_type Spot-check var.index format (gene symbols, Ensembl, etc.).
metadata_columns Declared columns present in obs; values within allowed set.
group_columns Each group column exists and has ≥2 usable groups.
default_contrasts Each contrast has ≥3 samples per group after subset_query.
Return structure
----------------
All public functions return a dict with at minimum:
status "pass" | "warning" | "error"
message short human-readable summary
validate_manifest_against_data returns:
overall_valid bool
n_errors int
n_warnings int
checks dict of category → result dict
errors list[str] — blocking issues
warnings list[str] — non-blocking issues
recommendations list[str] — suggested next steps
"""
from __future__ import annotations
import re
from typing import Any
import numpy as np
import pandas as pd
# ---------------------------------------------------------------------------
# Data-level compatibility map
# (detected_type → set of manifest data_level values that are compatible)
# ---------------------------------------------------------------------------
_COMPAT: dict[str, set[str]] = {
"raw_counts": {"raw_counts"},
"log_expression": {"log_expression", "normalized", "tpm", "fpkm", "protein_abundance"},
"log_ratio_microarray": {"log_ratio", "log_expression", "normalized"},
"unknown": {
"raw_counts",
"log_expression",
"log_ratio",
"normalized",
"tpm",
"fpkm",
"protein_abundance",
},
}
# Feature-ID patterns (applied to a sample of var.index values)
_GENE_SYMBOL_RE = re.compile(r"^[A-Z][A-Z0-9\-\.]{1,19}$")
_ENSEMBL_HUMAN = re.compile(r"^ENSG\d{11}$")
_ENSEMBL_MOUSE = re.compile(r"^ENSMUSG\d{11}$")
_ENTREZ_RE = re.compile(r"^\d+$")
_MANIFEST_TO_DETECTED: dict[str, str] = {
"gene_symbol": "gene_symbol",
"ensembl_gene_id": "ensembl",
"entrez_id": "entrez",
"probe_id": "probe_id",
"protein_id": "protein_id",
}
# ---------------------------------------------------------------------------
# Individual checks
# ---------------------------------------------------------------------------
def check_data_level(
X_flat: np.ndarray,
declared_data_level: str,
) -> dict[str, Any]:
"""
Compare manifest-declared data_level against detected expression type.
Parameters
----------
X_flat:
1-D numpy array sampled from adata.X (up to ~50 000 values is enough).
declared_data_level:
Value of manifest.data_level (e.g. 'log_expression', 'raw_counts').
Returns
-------
dict with: status, declared, detected, compatible, message, details.
"""
from src.workflows.microarray import classify_expression_data_type
info = classify_expression_data_type(X_flat)
detected = info["data_type"]
compatible = declared_data_level in _COMPAT.get(detected, set())
if detected == "unknown":
status = "warning"
msg = (
f"Could not confidently classify expression data "
f"(min={info['value_min']}, max={info['value_max']}, "
f"is_integer={info['is_integer']}, has_negatives={info['has_negatives']}). "
f"Manifest declares '{declared_data_level}' — verify manually."
)
elif compatible:
status = "pass"
msg = (
f"Detected '{detected}' is compatible with declared data_level='{declared_data_level}'."
)
else:
status = "error"
msg = (
f"MISMATCH: manifest declares data_level='{declared_data_level}' "
f"but data appears to be '{detected}'. "
f"Check whether the correct data level is declared in the manifest, "
f"or whether a normalisation step was already applied."
)
return {
"status": status,
"declared": declared_data_level,
"detected": detected,
"compatible": compatible,
"message": msg,
"details": {
k: info[k]
for k in (
"is_integer",
"has_negatives",
"value_min",
"value_max",
"value_mean",
"value_median",
)
},
}
def check_feature_id_type(
var_index_sample: list[str],
declared_type: str,
) -> dict[str, Any]:
"""
Spot-check var.index values to verify they match the declared feature_id_type.
Samples up to 200 values and checks the fraction matching each pattern.
Returns a warning (not error) when confidence is low, since probe ID
formats vary too much across platforms for a definitive check.
Parameters
----------
var_index_sample:
List of feature names from adata.var.index (sample or full list).
declared_type:
manifest.feature_id_type value.
"""
sample = var_index_sample[:200]
n = len(sample)
if n == 0:
return {
"status": "error",
"declared": declared_type,
"detected_pattern": "unknown",
"message": "var.index is empty — cannot check feature ID type.",
}
n_gene = sum(1 for v in sample if _GENE_SYMBOL_RE.match(str(v)))
n_ensh = sum(1 for v in sample if _ENSEMBL_HUMAN.match(str(v)))
n_ensm = sum(1 for v in sample if _ENSEMBL_MOUSE.match(str(v)))
n_entrez = sum(1 for v in sample if _ENTREZ_RE.match(str(v)))
fracs = {
"gene_symbol": n_gene / n,
"ensembl": (n_ensh + n_ensm) / n,
"entrez": n_entrez / n,
}
# Ensembl IDs satisfy the gene-symbol regex too, so check Ensembl first.
ensembl_frac = (n_ensh + n_ensm) / n
if ensembl_frac >= 0.5:
detected_pattern = "ensembl"
best_frac = ensembl_frac
best_pattern = "ensembl"
else:
best_pattern, best_frac = max(fracs.items(), key=lambda x: x[1])
detected_pattern = best_pattern if best_frac >= 0.5 else "probe_id_or_unknown"
expected_pattern = _MANIFEST_TO_DETECTED.get(declared_type, declared_type)
# probe_id and protein_id can't be detected reliably — treat as pass if
# the data doesn't strongly look like something else
if declared_type in ("probe_id", "protein_id"):
if best_frac > 0.7 and best_pattern in ("gene_symbol", "ensembl", "entrez"):
status = "warning"
msg = (
f"Manifest declares feature_id_type='{declared_type}' but "
f"{best_frac:.0%} of var.index values look like '{best_pattern}'. "
f"If probes were already collapsed to gene symbols, "
f"update feature_id_type to 'gene_symbol'."
)
else:
status = "pass"
msg = (
f"feature_id_type='{declared_type}' — pattern not automatically "
f"verifiable (platform-specific). Sample: {sample[:5]}"
)
elif detected_pattern == expected_pattern:
status = "pass"
msg = (
f"var.index looks like '{detected_pattern}' ({best_frac:.0%} match), "
f"consistent with declared feature_id_type='{declared_type}'."
)
elif detected_pattern == "probe_id_or_unknown":
status = "warning"
msg = (
f"Could not confidently classify var.index format "
f"(best match '{best_pattern}' at only {best_frac:.0%}). "
f"Manifest declares '{declared_type}'. Sample: {sample[:5]}"
)
else:
status = "warning"
msg = (
f"var.index looks like '{detected_pattern}' ({best_frac:.0%} match) "
f"but manifest declares feature_id_type='{declared_type}'. "
f"Sample: {sample[:5]}"
)
return {
"status": status,
"declared": declared_type,
"detected_pattern": detected_pattern,
"pattern_fractions": {k: round(v, 3) for k, v in fracs.items()},
"n_sampled": n,
"sample_features": sample[:10],
"message": msg,
}
def check_metadata_columns(
obs_df: pd.DataFrame,
manifest: Any,
) -> dict[str, Any]:
"""
Verify declared metadata_columns exist in obs and values match allowed set.
Parameters
----------
obs_df:
adata.obs as a pandas DataFrame.
manifest:
DatasetManifest instance.
"""
if not manifest.metadata_columns:
return {
"status": "pass",
"message": "No metadata_columns declared in manifest — skipped.",
"n_checked": 0,
"columns": {},
}
columns: dict[str, Any] = {}
errors: list[str] = []
warnings: list[str] = []
for col_name, col_def in manifest.metadata_columns.items():
if hasattr(col_def, "role"):
decoded_col = col_def.decoded_column or col_name
source_col = col_def.source_column
allowed = set(col_def.allowed_values)
missing_vals = set(col_def.missing_values) | {"", "nan", "None"}
else:
decoded_col = col_def.get("decoded_column") or col_name
source_col = col_def.get("source_column")
allowed = set(col_def.get("allowed_values") or [])
missing_vals = set(col_def.get("missing_values") or []) | {"", "nan", "None"}
check_col = decoded_col
result: dict[str, Any] = {"check_column": check_col, "source_column": source_col}
if check_col not in obs_df.columns:
if source_col and source_col in obs_df.columns:
result["status"] = "warning"
result["message"] = (
f"Source column '{source_col}' present but decoded column "
f"'{check_col}' missing — numeric decoder may not have run."
)
warnings.append(result["message"])
else:
result["status"] = "error"
result["message"] = f"Column '{check_col}' not found in obs."
errors.append(result["message"])
columns[col_name] = result
continue
actual_vals = obs_df[check_col].astype(str)
non_missing = actual_vals[~actual_vals.isin(missing_vals)]
value_counts = actual_vals.value_counts().to_dict()
result["value_counts"] = {str(k): int(v) for k, v in value_counts.items()}
result["n_missing"] = int(actual_vals.isin(missing_vals).sum())
result["n_annotated"] = int(len(non_missing))
if allowed:
unexpected = sorted(set(non_missing.unique()) - allowed)
result["unexpected_values"] = unexpected
if unexpected:
result["status"] = "warning"
result["message"] = (
f"Unexpected values in '{check_col}': {unexpected}. "
f"Declared allowed: {sorted(allowed)}"
)
warnings.append(result["message"])
else:
result["status"] = "pass"
result["message"] = (
f"All non-missing values in '{check_col}' match declared "
f"allowed_values ({len(non_missing)} annotated, "
f"{result['n_missing']} missing)."
)
else:
result["unexpected_values"] = []
result["status"] = "pass"
result["message"] = (
f"Column '{check_col}' present ({len(actual_vals)} values; "
f"no allowed_values constraint declared)."
)
columns[col_name] = result
overall = "error" if errors else ("warning" if warnings else "pass")
return {
"status": overall,
"n_checked": len(columns),
"columns": columns,
"message": (
f"{len(errors)} error(s), {len(warnings)} warning(s) "
f"across {len(columns)} declared metadata columns."
),
}
def check_group_columns(
obs_df: pd.DataFrame,
manifest: Any,
min_group_size: int = 3,
) -> dict[str, Any]:
"""
Verify each group_column exists in obs and has ≥2 groups with enough samples.
Parameters
----------
obs_df:
adata.obs as a pandas DataFrame.
manifest:
DatasetManifest instance.
min_group_size:
Minimum samples per group (default 3).
"""
columns: dict[str, Any] = {}
errors: list[str] = []
warnings: list[str] = []
for col in manifest.group_columns:
if col not in obs_df.columns:
columns[col] = {
"status": "error",
"present": False,
"message": f"group_column '{col}' not found in obs.",
}
errors.append(f"group_column '{col}' missing from obs")
continue
vc = obs_df[col].astype(str).value_counts()
qualifying = {str(k): int(v) for k, v in vc.items() if v >= min_group_size}
small = {str(k): int(v) for k, v in vc.items() if v < min_group_size}
n_usable = len(qualifying)
if n_usable < 2:
status = "warning"
msg = (
f"Column '{col}' has fewer than 2 groups with ≥{min_group_size} "
f"samples (qualifying: {list(qualifying.keys())})."
)
warnings.append(msg)
else:
status = "pass"
msg = f"Column '{col}' has {n_usable} usable groups (≥{min_group_size} samples each)."
columns[col] = {
"status": status,
"present": True,
"n_unique": int(vc.nunique()),
"qualifying_groups": qualifying,
"small_groups": small,
"message": msg,
}
overall = "error" if errors else ("warning" if warnings else "pass")
return {
"status": overall,
"n_checked": len(manifest.group_columns),
"columns": columns,
"message": (
f"{len(errors)} error(s), {len(warnings)} warning(s) "
f"across {len(manifest.group_columns)} group_column(s)."
),
}
def check_default_contrasts(
obs_df: pd.DataFrame,
manifest: Any,
min_group_size: int = 3,
) -> dict[str, Any]:
"""
Verify each default_contrast has enough samples in both groups.
Applies subset_query if declared, then counts test and control samples.
Parameters
----------
obs_df:
adata.obs as a pandas DataFrame.
manifest:
DatasetManifest instance.
min_group_size:
Minimum samples per group for a contrast to be feasible (default 3).
"""
if not manifest.default_contrasts:
return {
"status": "pass",
"message": "No default_contrasts declared — skipped.",
"contrasts": [],
}
results: list[dict] = []
errors: list[str] = []
warnings: list[str] = []
for i, contrast in enumerate(manifest.default_contrasts):
factor = contrast.get("design_factor", "")
test_grp = contrast.get("test_group", "")
ctrl_grp = contrast.get("control_group", "")
subset_q = contrast.get("subset_query")
method = contrast.get("method", "ttest")
label = f"{test_grp} vs {ctrl_grp} (contrast {i})"
entry: dict[str, Any] = {
"design_factor": factor,
"test_group": test_grp,
"control_group": ctrl_grp,
"subset_query": subset_q,
"method": method,
}
# Apply subset_query
working_df = obs_df
if subset_q:
try:
working_df = obs_df.query(subset_q)
entry["n_after_subset"] = len(working_df)
except Exception as exc:
entry["status"] = "error"
entry["message"] = f"subset_query failed: {exc}"
errors.append(entry["message"])
results.append(entry)
continue
if factor not in working_df.columns:
entry["status"] = "error"
entry["message"] = (
f"design_factor '{factor}' not found in obs (after applying subset_query)."
)
errors.append(entry["message"])
results.append(entry)
continue
col = working_df[factor].astype(str)
n_test = int((col == test_grp).sum())
n_ctrl = int((col == ctrl_grp).sum())
entry["n_test"] = n_test
entry["n_control"] = n_ctrl
if n_test == 0 or n_ctrl == 0:
entry["status"] = "error"
entry["message"] = (
f"{label}: one or both groups have 0 samples "
f"(test='{test_grp}': {n_test}, control='{ctrl_grp}': {n_ctrl}). "
f"Check group label spelling and subset_query."
)
errors.append(entry["message"])
elif n_test < min_group_size or n_ctrl < min_group_size:
entry["status"] = "warning"
entry["message"] = (
f"{label}: groups are small "
f"(test={n_test}, control={n_ctrl}, min={min_group_size}). "
f"Results may be underpowered."
)
warnings.append(entry["message"])
else:
entry["status"] = "pass"
entry["message"] = f"{label}: feasible — test={n_test}, control={n_ctrl}."
results.append(entry)
overall = "error" if errors else ("warning" if warnings else "pass")
return {
"status": overall,
"n_checked": len(manifest.default_contrasts),
"contrasts": results,
"message": (
f"{len(errors)} error(s), {len(warnings)} warning(s) "
f"across {len(manifest.default_contrasts)} contrast(s)."
),
}
# ---------------------------------------------------------------------------
# Top-level validator
# ---------------------------------------------------------------------------
def validate_manifest_against_data(
X_flat: np.ndarray,
var_index: list[str],
obs_df: pd.DataFrame,
manifest: Any,
) -> dict[str, Any]:
"""
Run all five checks and return a consolidated validation report.
Parameters
----------
X_flat:
1-D numpy array of expression values (sample up to ~50 000 values).
var_index:
Feature names from adata.var.index.
obs_df:
adata.obs as a pandas DataFrame.
manifest:
DatasetManifest instance.
Returns
-------
dict with: overall_valid, n_errors, n_warnings, checks, errors,
warnings, recommendations.
"""
checks = {
"data_level": check_data_level(X_flat, manifest.data_level),
"feature_id_type": check_feature_id_type(var_index, manifest.feature_id_type),
"metadata_columns": check_metadata_columns(obs_df, manifest),
"group_columns": check_group_columns(obs_df, manifest),
"default_contrasts": check_default_contrasts(obs_df, manifest),
}
all_errors: list[str] = []
all_warnings: list[str] = []
for name, result in checks.items():
status = result.get("status", "pass")
if status == "error":
all_errors.append(f"[{name}] {result.get('message', '')}")
elif status == "warning":
all_warnings.append(f"[{name}] {result.get('message', '')}")
overall_valid = len(all_errors) == 0
recommendations: list[str] = []
if overall_valid and not all_warnings:
recommendations.append(
"Manifest is consistent with loaded data. "
"Proceed with dataset_plan_analysis to start analysis."
)
if all_errors:
recommendations.append(
"Fix errors before running analysis — they indicate manifest "
"declarations that contradict the actual data."
)
if any("[feature_id_type]" in w for w in all_warnings):
recommendations.append(
"Verify feature_id_type by inspecting adata.var.index directly "
"and updating the manifest if the format has changed."
)
if any("[metadata_columns]" in w for w in all_warnings):
recommendations.append(
"Review unexpected metadata values — they may indicate new "
"categories in the data or a stale manifest."
)
if any("decoder" in w for w in all_warnings):
recommendations.append(
"Re-run decoupler_load_geo_series_matrix with decode_numeric=True "
"to apply numeric-to-label mapping."
)
return {
"dataset_id": manifest.dataset_id,
"n_samples": len(obs_df),
"n_features": len(var_index),
"overall_valid": overall_valid,
"n_errors": len(all_errors),
"n_warnings": len(all_warnings),
"checks": checks,
"errors": all_errors,
"warnings": all_warnings,
"recommendations": recommendations,
}