deepscenic 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- deepscenic/__init__.py +90 -0
- deepscenic/_data.py +131 -0
- deepscenic/_datasets.py +486 -0
- deepscenic/_genome.py +501 -0
- deepscenic/_io.py +364 -0
- deepscenic/_types.py +304 -0
- deepscenic/models/__init__.py +21 -0
- deepscenic/models/_decoder.py +60 -0
- deepscenic/models/_encoder.py +79 -0
- deepscenic/models/_layers.py +135 -0
- deepscenic/models/_motifnet.py +94 -0
- deepscenic/models/_vae.py +226 -0
- deepscenic/pl/__init__.py +59 -0
- deepscenic/pl/_colors.py +118 -0
- deepscenic/pl/_embedding.py +155 -0
- deepscenic/pl/_utils.py +149 -0
- deepscenic/pl/genomics/__init__.py +9 -0
- deepscenic/pl/genomics/_arc.py +183 -0
- deepscenic/pl/genomics/_browser.py +304 -0
- deepscenic/pl/grn/__init__.py +12 -0
- deepscenic/pl/grn/_heatmap.py +306 -0
- deepscenic/pl/grn/_network.py +388 -0
- deepscenic/pl/perturbation/__init__.py +12 -0
- deepscenic/pl/perturbation/_heatmap.py +254 -0
- deepscenic/pl/perturbation/_pca.py +292 -0
- deepscenic/pl/perturbation/_volcano.py +137 -0
- deepscenic/pl/sequence/__init__.py +10 -0
- deepscenic/pl/sequence/_ism.py +109 -0
- deepscenic/pl/sequence/_logo.py +201 -0
- deepscenic/pl/training/__init__.py +13 -0
- deepscenic/pl/training/_diagnostics.py +389 -0
- deepscenic/pp/__init__.py +391 -0
- deepscenic/pp/basic.py +416 -0
- deepscenic/pp/search_space.py +387 -0
- deepscenic/tl/__init__.py +75 -0
- deepscenic/tl/_dataloaders.py +334 -0
- deepscenic/tl/_grn.py +1181 -0
- deepscenic/tl/_inference.py +136 -0
- deepscenic/tl/_logging.py +172 -0
- deepscenic/tl/_loss.py +343 -0
- deepscenic/tl/_model.py +920 -0
- deepscenic/tl/_perturbation.py +313 -0
- deepscenic/tl/_sequence.py +242 -0
- deepscenic/tl/_train.py +1378 -0
- deepscenic/tl/_training_state.py +182 -0
- deepscenic-0.1.0.dist-info/METADATA +203 -0
- deepscenic-0.1.0.dist-info/RECORD +49 -0
- deepscenic-0.1.0.dist-info/WHEEL +4 -0
- deepscenic-0.1.0.dist-info/licenses/LICENSE +60 -0
deepscenic/__init__.py
ADDED
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
"""deepSCENIC: Deep learning for single-cell Gene Regulatory Networks."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import logging
|
|
6
|
+
from typing import TYPE_CHECKING
|
|
7
|
+
|
|
8
|
+
try:
|
|
9
|
+
from importlib.metadata import version
|
|
10
|
+
|
|
11
|
+
__version__ = version("deepscenic")
|
|
12
|
+
except Exception:
|
|
13
|
+
__version__ = "0.1.0" # Fallback for development
|
|
14
|
+
|
|
15
|
+
# Eager imports (lightweight, frequently used)
|
|
16
|
+
from ._data import (
|
|
17
|
+
SchemaError,
|
|
18
|
+
SchemaWarning,
|
|
19
|
+
is_valid_schema,
|
|
20
|
+
validate_schema,
|
|
21
|
+
)
|
|
22
|
+
from ._datasets import (
|
|
23
|
+
clear_cache,
|
|
24
|
+
fetch_gene_annotation,
|
|
25
|
+
fetch_tf_collection,
|
|
26
|
+
get_cache_info,
|
|
27
|
+
)
|
|
28
|
+
from ._genome import (
|
|
29
|
+
Genome,
|
|
30
|
+
GenomeIntervalDataset,
|
|
31
|
+
clear_genome,
|
|
32
|
+
get_genome,
|
|
33
|
+
register_genome,
|
|
34
|
+
)
|
|
35
|
+
from ._io import read, read_bed, write
|
|
36
|
+
|
|
37
|
+
if TYPE_CHECKING:
|
|
38
|
+
from . import models, pl, pp, tl
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def __getattr__(name: str):
|
|
42
|
+
"""Lazy load heavy submodules."""
|
|
43
|
+
import importlib
|
|
44
|
+
|
|
45
|
+
if name in {"pp", "tl", "pl", "models"}:
|
|
46
|
+
return importlib.import_module(f".{name}", __name__)
|
|
47
|
+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def __dir__() -> list[str]:
|
|
51
|
+
return __all__
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
__all__ = [
|
|
55
|
+
"__version__",
|
|
56
|
+
# I/O
|
|
57
|
+
"read",
|
|
58
|
+
"read_bed",
|
|
59
|
+
"write",
|
|
60
|
+
# Datasets
|
|
61
|
+
"fetch_tf_collection",
|
|
62
|
+
"fetch_gene_annotation",
|
|
63
|
+
"clear_cache",
|
|
64
|
+
"get_cache_info",
|
|
65
|
+
# Genome
|
|
66
|
+
"Genome",
|
|
67
|
+
"GenomeIntervalDataset",
|
|
68
|
+
"register_genome",
|
|
69
|
+
"get_genome",
|
|
70
|
+
"clear_genome",
|
|
71
|
+
# Data utilities
|
|
72
|
+
"validate_schema",
|
|
73
|
+
"is_valid_schema",
|
|
74
|
+
"SchemaError",
|
|
75
|
+
"SchemaWarning",
|
|
76
|
+
# Submodules (lazy loaded)
|
|
77
|
+
"pp",
|
|
78
|
+
"tl",
|
|
79
|
+
"pl",
|
|
80
|
+
"models",
|
|
81
|
+
]
|
|
82
|
+
|
|
83
|
+
# Configure package-level logging to show INFO messages by default
|
|
84
|
+
_logger = logging.getLogger("deepscenic")
|
|
85
|
+
_logger.setLevel(logging.INFO)
|
|
86
|
+
|
|
87
|
+
if not _logger.handlers:
|
|
88
|
+
_handler = logging.StreamHandler()
|
|
89
|
+
_handler.setFormatter(logging.Formatter("%(name)s - %(levelname)s - %(message)s"))
|
|
90
|
+
_logger.addHandler(_handler)
|
deepscenic/_data.py
ADDED
|
@@ -0,0 +1,131 @@
|
|
|
1
|
+
"""MuData schema validation for deepSCENIC."""
|
|
2
|
+
|
|
3
|
+
import warnings
|
|
4
|
+
|
|
5
|
+
import mudata as md
|
|
6
|
+
|
|
7
|
+
__all__ = [
|
|
8
|
+
"validate_schema",
|
|
9
|
+
"is_valid_schema",
|
|
10
|
+
"SchemaError",
|
|
11
|
+
"SchemaWarning",
|
|
12
|
+
]
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class SchemaError(Exception):
|
|
16
|
+
"""Raised when MuData doesn't conform to deepSCENIC schema."""
|
|
17
|
+
|
|
18
|
+
pass
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class SchemaWarning(UserWarning):
|
|
22
|
+
"""Warning for non-critical schema issues."""
|
|
23
|
+
|
|
24
|
+
pass
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def validate_schema(
|
|
28
|
+
mdata: md.MuData,
|
|
29
|
+
mode: str = "training",
|
|
30
|
+
strict: bool = False,
|
|
31
|
+
) -> list[str]:
|
|
32
|
+
"""
|
|
33
|
+
Validate MuData against deepSCENIC schema.
|
|
34
|
+
|
|
35
|
+
Parameters
|
|
36
|
+
----------
|
|
37
|
+
mdata
|
|
38
|
+
Data to validate
|
|
39
|
+
mode
|
|
40
|
+
Validation mode:
|
|
41
|
+
- 'training': require rna + atac + r2g (strict)
|
|
42
|
+
- 'inference': require rna only (flexible)
|
|
43
|
+
strict
|
|
44
|
+
If True, raise errors. If False, collect warnings.
|
|
45
|
+
|
|
46
|
+
Returns
|
|
47
|
+
-------
|
|
48
|
+
List of validation issues (empty if valid)
|
|
49
|
+
|
|
50
|
+
Raises
|
|
51
|
+
------
|
|
52
|
+
SchemaError
|
|
53
|
+
If strict=True and validation fails
|
|
54
|
+
|
|
55
|
+
Examples
|
|
56
|
+
--------
|
|
57
|
+
>>> import deepscenic as ds
|
|
58
|
+
>>> # Validate for training (strict)
|
|
59
|
+
>>> ds.validate_schema(mdata, mode="training", strict=True)
|
|
60
|
+
>>> # Validate for inference (RNA-only OK)
|
|
61
|
+
>>> ds.validate_schema(mdata, mode="inference")
|
|
62
|
+
"""
|
|
63
|
+
issues = []
|
|
64
|
+
|
|
65
|
+
if mode not in ("training", "inference"):
|
|
66
|
+
raise ValueError(f"mode must be 'training' or 'inference', got: {mode}")
|
|
67
|
+
|
|
68
|
+
# Check RNA modality (always required)
|
|
69
|
+
if "rna" not in mdata.mod:
|
|
70
|
+
issues.append("Missing 'rna' modality")
|
|
71
|
+
|
|
72
|
+
# Check ATAC modality (required for training, optional for inference)
|
|
73
|
+
if "atac" not in mdata.mod:
|
|
74
|
+
if mode == "training":
|
|
75
|
+
issues.append("Missing 'atac' modality (required for training)")
|
|
76
|
+
|
|
77
|
+
if issues and strict:
|
|
78
|
+
raise SchemaError(f"Schema validation failed: {issues}")
|
|
79
|
+
|
|
80
|
+
# Check RNA modality
|
|
81
|
+
if "rna" in mdata.mod:
|
|
82
|
+
rna = mdata.mod["rna"]
|
|
83
|
+
|
|
84
|
+
if "is_tf" not in rna.var.columns:
|
|
85
|
+
issues.append("rna.var missing 'is_tf' column")
|
|
86
|
+
elif rna.var["is_tf"].dtype != bool:
|
|
87
|
+
issues.append("rna.var['is_tf'] should be bool")
|
|
88
|
+
|
|
89
|
+
if mode == "training" and "split" not in rna.var.columns:
|
|
90
|
+
issues.append("rna.var missing 'split' column")
|
|
91
|
+
|
|
92
|
+
# Check ATAC modality
|
|
93
|
+
if "atac" in mdata.mod:
|
|
94
|
+
atac = mdata.mod["atac"]
|
|
95
|
+
|
|
96
|
+
for col in ["chromosome", "start", "end"]:
|
|
97
|
+
if col not in atac.var.columns:
|
|
98
|
+
issues.append(f"atac.var missing '{col}' column")
|
|
99
|
+
|
|
100
|
+
if "split" not in atac.var.columns:
|
|
101
|
+
issues.append("atac.var missing 'split' column")
|
|
102
|
+
|
|
103
|
+
# Check shared obs
|
|
104
|
+
if "split" not in mdata.obs.columns:
|
|
105
|
+
issues.append("mdata.obs missing 'split' column")
|
|
106
|
+
|
|
107
|
+
# Check r2g (required for training, not for inference)
|
|
108
|
+
if mode == "training":
|
|
109
|
+
if "r2g" not in mdata.uns:
|
|
110
|
+
issues.append("mdata.uns missing 'r2g' (required for training)")
|
|
111
|
+
else:
|
|
112
|
+
r2g = mdata.uns["r2g"]
|
|
113
|
+
# Check required keys
|
|
114
|
+
for key in ["config"]:
|
|
115
|
+
if key not in r2g:
|
|
116
|
+
issues.append(f"mdata.uns['r2g'] missing '{key}'")
|
|
117
|
+
|
|
118
|
+
if issues:
|
|
119
|
+
if strict:
|
|
120
|
+
raise SchemaError(f"Schema validation failed: {issues}")
|
|
121
|
+
else:
|
|
122
|
+
for issue in issues:
|
|
123
|
+
warnings.warn(issue, SchemaWarning, stacklevel=2)
|
|
124
|
+
|
|
125
|
+
return issues
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def is_valid_schema(mdata: md.MuData, mode: str = "training") -> bool:
|
|
129
|
+
"""Check if MuData conforms to schema without raising."""
|
|
130
|
+
issues = validate_schema(mdata, mode=mode, strict=False)
|
|
131
|
+
return len(issues) == 0
|
deepscenic/_datasets.py
ADDED
|
@@ -0,0 +1,486 @@
|
|
|
1
|
+
"""Dataset fetching utilities for deepSCENIC."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import hashlib
|
|
6
|
+
import logging
|
|
7
|
+
import re
|
|
8
|
+
import time
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from typing import Literal
|
|
11
|
+
|
|
12
|
+
import pandas as pd
|
|
13
|
+
import pooch
|
|
14
|
+
import requests
|
|
15
|
+
|
|
16
|
+
__all__ = ["fetch_tf_collection", "fetch_gene_annotation", "clear_cache", "get_cache_info"]
|
|
17
|
+
|
|
18
|
+
log = logging.getLogger("deepscenic.datasets")
|
|
19
|
+
_NCBI_MAX_RETRIES = 3
|
|
20
|
+
|
|
21
|
+
# TF collection URLs (SCENIC+ resources)
|
|
22
|
+
TF_COLLECTION_URLS = {
|
|
23
|
+
"human": "https://resources.aertslab.org/cistarget/tf_lists/allTFs_hg38.txt",
|
|
24
|
+
"mouse": "https://resources.aertslab.org/cistarget/tf_lists/allTFs_mm.txt",
|
|
25
|
+
"fly": "https://resources.aertslab.org/cistarget/tf_lists/allTFs_dmel.txt",
|
|
26
|
+
}
|
|
27
|
+
|
|
28
|
+
# Known checksums (None allows download without verification)
|
|
29
|
+
TF_COLLECTION_CHECKSUMS = {
|
|
30
|
+
"human": "3953034f84112c60d3d8ef15b0e0c8ac5fce0b40d2c7c0824c2945c70cee2523",
|
|
31
|
+
"mouse": "17a95e142147fb7dc063d7b9e84262746b0b64f622793b3cc5df0eddf2f1194c",
|
|
32
|
+
"fly": "20d7e11540b595dda3ed133f86af559c9fc708810dc39d4bef106961fedf348d",
|
|
33
|
+
}
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class NCBIError(Exception):
|
|
37
|
+
"""Error fetching data from NCBI."""
|
|
38
|
+
|
|
39
|
+
pass
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _get_cache_dir() -> Path:
|
|
43
|
+
"""Get the deepscenic cache directory."""
|
|
44
|
+
cache_dir: Path = pooch.os_cache("deepscenic")
|
|
45
|
+
cache_dir.mkdir(parents=True, exist_ok=True)
|
|
46
|
+
return cache_dir
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _get_cache_key(
|
|
50
|
+
species: str,
|
|
51
|
+
biomart_host: str,
|
|
52
|
+
use_ucsc_chromosome_style: bool,
|
|
53
|
+
transcript_type: str | None,
|
|
54
|
+
) -> str:
|
|
55
|
+
"""Generate a unique cache key for the query parameters."""
|
|
56
|
+
# Hash the host to keep filename reasonable
|
|
57
|
+
host_hash = hashlib.md5(biomart_host.encode()).hexdigest()[:8]
|
|
58
|
+
ucsc = "ucsc" if use_ucsc_chromosome_style else "ensembl"
|
|
59
|
+
ttype = transcript_type or "all"
|
|
60
|
+
return f"gene_annot_{species}_{host_hash}_{ucsc}_{ttype}"
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def fetch_tf_collection(
|
|
64
|
+
species: Literal["human", "mouse", "fly"],
|
|
65
|
+
) -> list[str]:
|
|
66
|
+
"""
|
|
67
|
+
Fetch SCENIC+ transcription factor collection.
|
|
68
|
+
|
|
69
|
+
Downloads and caches the TF list for the specified species from
|
|
70
|
+
the SCENIC+ resources at aertslab.org.
|
|
71
|
+
|
|
72
|
+
Parameters
|
|
73
|
+
----------
|
|
74
|
+
species
|
|
75
|
+
Species to fetch TF list for.
|
|
76
|
+
|
|
77
|
+
Returns
|
|
78
|
+
-------
|
|
79
|
+
List of transcription factor gene names.
|
|
80
|
+
|
|
81
|
+
Examples
|
|
82
|
+
--------
|
|
83
|
+
>>> import deepscenic as ds
|
|
84
|
+
>>> tfs = ds.fetch_tf_collection(species="mouse")
|
|
85
|
+
>>> len(tfs)
|
|
86
|
+
1390
|
|
87
|
+
>>> tfs[:3]
|
|
88
|
+
['Adnp', 'Aebp1', 'Aebp2']
|
|
89
|
+
"""
|
|
90
|
+
if species not in TF_COLLECTION_URLS:
|
|
91
|
+
raise ValueError(f"Unknown species: {species}. Available: {list(TF_COLLECTION_URLS.keys())}")
|
|
92
|
+
|
|
93
|
+
url = TF_COLLECTION_URLS[species]
|
|
94
|
+
known_hash = TF_COLLECTION_CHECKSUMS[species]
|
|
95
|
+
|
|
96
|
+
cache_dir = pooch.os_cache("deepscenic")
|
|
97
|
+
cache_dir.mkdir(parents=True, exist_ok=True)
|
|
98
|
+
|
|
99
|
+
local_path = pooch.retrieve(
|
|
100
|
+
url=url,
|
|
101
|
+
known_hash=known_hash,
|
|
102
|
+
path=cache_dir,
|
|
103
|
+
fname=f"allTFs_{species}.txt",
|
|
104
|
+
progressbar=True,
|
|
105
|
+
)
|
|
106
|
+
|
|
107
|
+
# Read TF list
|
|
108
|
+
with open(local_path) as f:
|
|
109
|
+
tf_names = [line.strip() for line in f if line.strip()]
|
|
110
|
+
|
|
111
|
+
return tf_names
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def fetch_gene_annotation(
|
|
115
|
+
species: str = "hsapiens",
|
|
116
|
+
biomart_host: str = "http://www.ensembl.org",
|
|
117
|
+
use_ucsc_chromosome_style: bool = True,
|
|
118
|
+
transcript_type: str = "protein_coding",
|
|
119
|
+
force_download: bool = False,
|
|
120
|
+
) -> tuple[pd.DataFrame, pd.DataFrame | None]:
|
|
121
|
+
"""
|
|
122
|
+
Download gene annotation from Ensembl Biomart.
|
|
123
|
+
|
|
124
|
+
Results are cached locally as parquet files for fast subsequent access.
|
|
125
|
+
After the first download, this function loads from cache without
|
|
126
|
+
making network requests.
|
|
127
|
+
|
|
128
|
+
Parameters
|
|
129
|
+
----------
|
|
130
|
+
species
|
|
131
|
+
Species name for Ensembl (e.g., "hsapiens", "mmusculus", "dmelanogaster").
|
|
132
|
+
biomart_host
|
|
133
|
+
Biomart host URL. Use archived hosts for reproducibility:
|
|
134
|
+
- "http://nov2020.archive.ensembl.org/" for GRCm38
|
|
135
|
+
- "http://www.ensembl.org" for latest
|
|
136
|
+
use_ucsc_chromosome_style
|
|
137
|
+
Convert chromosome names to UCSC style (chr1, chr2, etc.).
|
|
138
|
+
transcript_type
|
|
139
|
+
Filter for transcript type. Set to None to include all types.
|
|
140
|
+
force_download
|
|
141
|
+
Force re-download even if cached data exists.
|
|
142
|
+
|
|
143
|
+
Returns
|
|
144
|
+
-------
|
|
145
|
+
gene_annotation
|
|
146
|
+
Gene annotation with columns:
|
|
147
|
+
- Chromosome, Start, End, Strand, Transcription_Start_Site, Transcript_type
|
|
148
|
+
Index is Gene name.
|
|
149
|
+
chromsizes
|
|
150
|
+
Chromosome sizes (if available from NCBI):
|
|
151
|
+
- Chromosome, Start (0), End
|
|
152
|
+
|
|
153
|
+
Examples
|
|
154
|
+
--------
|
|
155
|
+
>>> import deepscenic as ds
|
|
156
|
+
>>> annot, chromsizes = ds.fetch_gene_annotation(
|
|
157
|
+
... species="mmusculus",
|
|
158
|
+
... biomart_host="http://nov2020.archive.ensembl.org/",
|
|
159
|
+
... )
|
|
160
|
+
>>> annot.head()
|
|
161
|
+
Chromosome Start End Strand Transcription_Start_Site
|
|
162
|
+
Gene
|
|
163
|
+
0610005C13Rik chr7 ...
|
|
164
|
+
"""
|
|
165
|
+
cache_dir = _get_cache_dir()
|
|
166
|
+
cache_key = _get_cache_key(species, biomart_host, use_ucsc_chromosome_style, transcript_type)
|
|
167
|
+
|
|
168
|
+
annot_path = cache_dir / f"{cache_key}.parquet"
|
|
169
|
+
chromsizes_path = cache_dir / f"{cache_key}_chromsizes.parquet"
|
|
170
|
+
|
|
171
|
+
# Try to load from cache
|
|
172
|
+
if not force_download and annot_path.exists():
|
|
173
|
+
log.info(f"Loading cached gene annotation from {annot_path}")
|
|
174
|
+
annot = pd.read_parquet(annot_path)
|
|
175
|
+
chromsizes = pd.read_parquet(chromsizes_path) if chromsizes_path.exists() else None
|
|
176
|
+
log.info(f"Loaded annotation for {len(annot)} genes from cache")
|
|
177
|
+
return annot, chromsizes
|
|
178
|
+
|
|
179
|
+
# Fetch from Biomart (this imports pybiomart)
|
|
180
|
+
log.info(f"Fetching gene annotation from {biomart_host}...")
|
|
181
|
+
annot, chromsizes = _fetch_from_biomart(
|
|
182
|
+
species=species,
|
|
183
|
+
biomart_host=biomart_host,
|
|
184
|
+
use_ucsc_chromosome_style=use_ucsc_chromosome_style,
|
|
185
|
+
transcript_type=transcript_type,
|
|
186
|
+
)
|
|
187
|
+
|
|
188
|
+
# Cache results as parquet
|
|
189
|
+
log.info(f"Caching gene annotation to {annot_path}")
|
|
190
|
+
annot.to_parquet(annot_path)
|
|
191
|
+
if chromsizes is not None:
|
|
192
|
+
chromsizes.to_parquet(chromsizes_path)
|
|
193
|
+
|
|
194
|
+
return annot, chromsizes
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
def _fetch_from_biomart(
|
|
198
|
+
species: str,
|
|
199
|
+
biomart_host: str,
|
|
200
|
+
use_ucsc_chromosome_style: bool,
|
|
201
|
+
transcript_type: str | None,
|
|
202
|
+
) -> tuple[pd.DataFrame, pd.DataFrame | None]:
|
|
203
|
+
"""
|
|
204
|
+
Fetch gene annotation from Biomart.
|
|
205
|
+
|
|
206
|
+
This function imports pybiomart, which creates a .pybiomart.sqlite file.
|
|
207
|
+
It is isolated here so the main fetch_gene_annotation() can avoid
|
|
208
|
+
importing pybiomart when loading from cache.
|
|
209
|
+
|
|
210
|
+
The requests_cache is redirected to our cache directory to prevent
|
|
211
|
+
.pybiomart.sqlite appearing in the user's working directory.
|
|
212
|
+
"""
|
|
213
|
+
# Redirect pybiomart's cache to our cache directory BEFORE importing
|
|
214
|
+
import requests_cache
|
|
215
|
+
|
|
216
|
+
cache_dir = _get_cache_dir()
|
|
217
|
+
_original_install_cache = requests_cache.install_cache
|
|
218
|
+
|
|
219
|
+
def _patched_install_cache(cache_name: str = "http_cache", **kwargs):
|
|
220
|
+
if cache_name == ".pybiomart":
|
|
221
|
+
cache_name = str(cache_dir / "pybiomart_requests")
|
|
222
|
+
return _original_install_cache(cache_name, **kwargs)
|
|
223
|
+
|
|
224
|
+
requests_cache.install_cache = _patched_install_cache # type: ignore[assignment]
|
|
225
|
+
|
|
226
|
+
# Now import pybiomart (will use our patched install_cache)
|
|
227
|
+
import pybiomart as pbm
|
|
228
|
+
|
|
229
|
+
dataset_name = f"{species}_gene_ensembl"
|
|
230
|
+
|
|
231
|
+
log.info(f"Connecting to Biomart host: {biomart_host}")
|
|
232
|
+
server = pbm.Server(host=biomart_host, use_cache=False)
|
|
233
|
+
mart = server["ENSEMBL_MART_ENSEMBL"]
|
|
234
|
+
|
|
235
|
+
if dataset_name not in mart.list_datasets()["name"].to_numpy():
|
|
236
|
+
raise ValueError(f"Dataset '{dataset_name}' not found. Check species name or Biomart host.")
|
|
237
|
+
|
|
238
|
+
dataset = mart[dataset_name]
|
|
239
|
+
|
|
240
|
+
# Handle different Biomart attribute names across versions
|
|
241
|
+
external_gene_name_query = (
|
|
242
|
+
"external_gene_name" if "external_gene_name" in dataset.attributes.keys() else "hgnc_symbol"
|
|
243
|
+
)
|
|
244
|
+
tss_query = (
|
|
245
|
+
"transcription_start_site" if "transcription_start_site" in dataset.attributes.keys() else "transcript_start"
|
|
246
|
+
)
|
|
247
|
+
|
|
248
|
+
log.info(f"Querying gene annotation for {species}...")
|
|
249
|
+
annot = pd.DataFrame(
|
|
250
|
+
dataset.query(
|
|
251
|
+
attributes=[
|
|
252
|
+
"chromosome_name",
|
|
253
|
+
"start_position",
|
|
254
|
+
"end_position",
|
|
255
|
+
"strand",
|
|
256
|
+
external_gene_name_query,
|
|
257
|
+
tss_query,
|
|
258
|
+
"transcript_biotype",
|
|
259
|
+
]
|
|
260
|
+
)
|
|
261
|
+
)
|
|
262
|
+
annot.columns = [
|
|
263
|
+
"Chromosome",
|
|
264
|
+
"Start",
|
|
265
|
+
"End",
|
|
266
|
+
"Strand",
|
|
267
|
+
"Gene",
|
|
268
|
+
"Transcription_Start_Site",
|
|
269
|
+
"Transcript_type",
|
|
270
|
+
]
|
|
271
|
+
|
|
272
|
+
# Filter for transcript type
|
|
273
|
+
if transcript_type:
|
|
274
|
+
annot = pd.DataFrame(annot[annot.Transcript_type == transcript_type].copy())
|
|
275
|
+
|
|
276
|
+
# Convert strand from numeric to +/-
|
|
277
|
+
annot["Strand"] = ["+" if strand == 1 else "-" for strand in annot["Strand"]]
|
|
278
|
+
|
|
279
|
+
# Try to get chromosome sizes from NCBI
|
|
280
|
+
chromsizes = None
|
|
281
|
+
try:
|
|
282
|
+
chromsizes, annot = _fetch_chromsizes_and_filter(annot, dataset, use_ucsc_chromosome_style)
|
|
283
|
+
except NCBIError as e:
|
|
284
|
+
log.warning(f"Could not fetch chromosome info from NCBI: {e}")
|
|
285
|
+
log.warning("Returning annotation without chromosome filtering.")
|
|
286
|
+
|
|
287
|
+
# Remove duplicate genes (keep first occurrence)
|
|
288
|
+
annot = annot.drop_duplicates(subset="Gene", keep="first")
|
|
289
|
+
|
|
290
|
+
# Remove rows with empty gene names
|
|
291
|
+
annot = annot[annot["Gene"].notna() & (annot["Gene"] != "")]
|
|
292
|
+
|
|
293
|
+
# Set Gene as index for easy lookup
|
|
294
|
+
annot = annot.set_index("Gene")
|
|
295
|
+
|
|
296
|
+
log.info(f"Downloaded annotation for {len(annot)} genes")
|
|
297
|
+
|
|
298
|
+
return annot, chromsizes
|
|
299
|
+
|
|
300
|
+
|
|
301
|
+
def _fetch_chromsizes_and_filter(
|
|
302
|
+
annot: pd.DataFrame,
|
|
303
|
+
dataset,
|
|
304
|
+
use_ucsc_style: bool,
|
|
305
|
+
) -> tuple[pd.DataFrame, pd.DataFrame]:
|
|
306
|
+
"""Fetch chromosome sizes from NCBI and filter annotation."""
|
|
307
|
+
import xml.etree.ElementTree as xml_tree
|
|
308
|
+
|
|
309
|
+
# Get assembly name from dataset
|
|
310
|
+
regex_display = re.search(r"\((.*?)\)", dataset.display_name)
|
|
311
|
+
if regex_display is None:
|
|
312
|
+
raise NCBIError("Could not find assembly from Biomart display name")
|
|
313
|
+
|
|
314
|
+
ncbi_search_term = regex_display.group(1)
|
|
315
|
+
log.info(f"Using genome assembly: {ncbi_search_term}")
|
|
316
|
+
|
|
317
|
+
def _get_with_retries(url: str, params: dict | None = None) -> requests.Response:
|
|
318
|
+
for _ in range(_NCBI_MAX_RETRIES):
|
|
319
|
+
resp = requests.get(url, params=params)
|
|
320
|
+
if resp.ok:
|
|
321
|
+
return resp
|
|
322
|
+
time.sleep(0.5)
|
|
323
|
+
raise NCBIError(f"Failed to fetch from {url} after {_NCBI_MAX_RETRIES} retries")
|
|
324
|
+
|
|
325
|
+
# Search NCBI assembly database
|
|
326
|
+
esearch_resp = _get_with_retries(
|
|
327
|
+
"https://eutils.ncbi.nlm.nih.gov/entrez/eutils/esearch.fcgi",
|
|
328
|
+
params={"db": "assembly", "term": f"{ncbi_search_term}[Assembly Name]"},
|
|
329
|
+
)
|
|
330
|
+
|
|
331
|
+
id_list = xml_tree.fromstring(esearch_resp.content).find("IdList")
|
|
332
|
+
id_elem = id_list.find("Id") if id_list is not None else None
|
|
333
|
+
|
|
334
|
+
if id_elem is None:
|
|
335
|
+
raise NCBIError(f"No assembly found for: {ncbi_search_term}")
|
|
336
|
+
|
|
337
|
+
assembly_id = id_elem.text
|
|
338
|
+
log.info(f"Found NCBI assembly ID: {assembly_id}")
|
|
339
|
+
|
|
340
|
+
# Get assembly summary
|
|
341
|
+
esummary_resp = _get_with_retries(
|
|
342
|
+
"https://eutils.ncbi.nlm.nih.gov/entrez/eutils/esummary.fcgi",
|
|
343
|
+
params={"db": "assembly", "id": assembly_id},
|
|
344
|
+
)
|
|
345
|
+
|
|
346
|
+
doc_summary = xml_tree.fromstring(esummary_resp.content).find(".//DocumentSummary")
|
|
347
|
+
if doc_summary is None:
|
|
348
|
+
raise NCBIError("No DocumentSummary in NCBI response")
|
|
349
|
+
|
|
350
|
+
ftp_path = doc_summary.find("FtpPath_Assembly_rpt")
|
|
351
|
+
if ftp_path is None or ftp_path.text is None:
|
|
352
|
+
raise NCBIError("No FTP path for assembly report")
|
|
353
|
+
|
|
354
|
+
log.info(f"Downloading assembly report from: {ftp_path.text}")
|
|
355
|
+
|
|
356
|
+
# Load assembly report
|
|
357
|
+
assembly_report = pd.read_csv(
|
|
358
|
+
ftp_path.text,
|
|
359
|
+
comment="#",
|
|
360
|
+
names=[
|
|
361
|
+
"Sequence-Name",
|
|
362
|
+
"Sequence-Role",
|
|
363
|
+
"Assigned-Molecule",
|
|
364
|
+
"Assigned-Molecule-Location/Type",
|
|
365
|
+
"GenBank-Accn",
|
|
366
|
+
"Relationship",
|
|
367
|
+
"RefSeq-Accn",
|
|
368
|
+
"Assembly-Unit",
|
|
369
|
+
"Sequence-Length",
|
|
370
|
+
"UCSC-style-name",
|
|
371
|
+
],
|
|
372
|
+
sep="\t",
|
|
373
|
+
)
|
|
374
|
+
|
|
375
|
+
# Filter to assembled chromosomes only
|
|
376
|
+
assembly_report = assembly_report[assembly_report["Sequence-Role"] == "assembled-molecule"].copy()
|
|
377
|
+
|
|
378
|
+
assembled_molecules = assembly_report["Sequence-Name"].tolist()
|
|
379
|
+
log.info(f"Found {len(assembled_molecules)} assembled chromosomes")
|
|
380
|
+
|
|
381
|
+
# Convert chromosome column to string for comparison
|
|
382
|
+
# (Ensembl returns mixed types: int for 1-22, str for X/Y/MT)
|
|
383
|
+
annot["Chromosome"] = annot["Chromosome"].astype(str)
|
|
384
|
+
|
|
385
|
+
# Filter annotation to assembled chromosomes
|
|
386
|
+
annot = pd.DataFrame(annot[annot["Chromosome"].isin(assembled_molecules)].copy())
|
|
387
|
+
|
|
388
|
+
# Build chromsizes DataFrame
|
|
389
|
+
chromsizes = pd.DataFrame(
|
|
390
|
+
{
|
|
391
|
+
"Chromosome": assembly_report["Sequence-Name"].tolist(),
|
|
392
|
+
"Start": 0,
|
|
393
|
+
"End": assembly_report["Sequence-Length"].tolist(),
|
|
394
|
+
}
|
|
395
|
+
)
|
|
396
|
+
|
|
397
|
+
# Convert to UCSC style if requested
|
|
398
|
+
if use_ucsc_style:
|
|
399
|
+
ensembl_to_ucsc = dict(
|
|
400
|
+
zip(
|
|
401
|
+
assembly_report["Sequence-Name"],
|
|
402
|
+
assembly_report["UCSC-style-name"],
|
|
403
|
+
strict=False,
|
|
404
|
+
)
|
|
405
|
+
)
|
|
406
|
+
annot["Chromosome"] = [ensembl_to_ucsc.get(c, c) for c in annot["Chromosome"]]
|
|
407
|
+
chromsizes["Chromosome"] = [ensembl_to_ucsc.get(c, c) for c in chromsizes["Chromosome"]]
|
|
408
|
+
log.info("Converted chromosome names to UCSC style")
|
|
409
|
+
|
|
410
|
+
return chromsizes, annot
|
|
411
|
+
|
|
412
|
+
|
|
413
|
+
def clear_cache(pattern: str | None = None) -> list[str]:
|
|
414
|
+
"""
|
|
415
|
+
Clear cached data files.
|
|
416
|
+
|
|
417
|
+
Parameters
|
|
418
|
+
----------
|
|
419
|
+
pattern
|
|
420
|
+
Glob pattern to match. If None, clears all cache.
|
|
421
|
+
Example: "gene_annot_*" to clear only gene annotations.
|
|
422
|
+
|
|
423
|
+
Returns
|
|
424
|
+
-------
|
|
425
|
+
List of deleted file paths.
|
|
426
|
+
|
|
427
|
+
Examples
|
|
428
|
+
--------
|
|
429
|
+
>>> import deepscenic as ds
|
|
430
|
+
>>> # Clear only gene annotation caches
|
|
431
|
+
>>> ds.clear_cache("gene_annot_*")
|
|
432
|
+
['/home/user/.cache/deepscenic/gene_annot_mmusculus_abc123_ucsc_protein_coding.parquet']
|
|
433
|
+
>>> # Clear all cache files
|
|
434
|
+
>>> ds.clear_cache()
|
|
435
|
+
[]
|
|
436
|
+
"""
|
|
437
|
+
cache_dir = pooch.os_cache("deepscenic")
|
|
438
|
+
if not cache_dir.exists():
|
|
439
|
+
return []
|
|
440
|
+
|
|
441
|
+
if pattern is None:
|
|
442
|
+
pattern = "*"
|
|
443
|
+
|
|
444
|
+
deleted = []
|
|
445
|
+
for f in cache_dir.glob(pattern):
|
|
446
|
+
if f.is_file():
|
|
447
|
+
f.unlink()
|
|
448
|
+
deleted.append(str(f))
|
|
449
|
+
|
|
450
|
+
return deleted
|
|
451
|
+
|
|
452
|
+
|
|
453
|
+
def get_cache_info() -> dict:
|
|
454
|
+
"""
|
|
455
|
+
Get information about cached files.
|
|
456
|
+
|
|
457
|
+
Returns
|
|
458
|
+
-------
|
|
459
|
+
dict
|
|
460
|
+
Dictionary containing:
|
|
461
|
+
- cache_dir: Path to cache directory
|
|
462
|
+
- files: List of dicts with name and size_mb
|
|
463
|
+
- total_size_mb: Total size of all cached files
|
|
464
|
+
|
|
465
|
+
Examples
|
|
466
|
+
--------
|
|
467
|
+
>>> import deepscenic as ds
|
|
468
|
+
>>> info = ds.get_cache_info()
|
|
469
|
+
>>> info['cache_dir']
|
|
470
|
+
'/home/user/.cache/deepscenic'
|
|
471
|
+
>>> info['total_size_mb']
|
|
472
|
+
5.2
|
|
473
|
+
"""
|
|
474
|
+
cache_dir = pooch.os_cache("deepscenic")
|
|
475
|
+
files = []
|
|
476
|
+
if cache_dir.exists():
|
|
477
|
+
for f in cache_dir.iterdir():
|
|
478
|
+
if f.is_file():
|
|
479
|
+
files.append(
|
|
480
|
+
{
|
|
481
|
+
"name": f.name,
|
|
482
|
+
"size_mb": f.stat().st_size / (1024 * 1024),
|
|
483
|
+
}
|
|
484
|
+
)
|
|
485
|
+
|
|
486
|
+
return {"cache_dir": str(cache_dir), "files": files, "total_size_mb": sum(f["size_mb"] for f in files)}
|