deep-cnv 0.0.1__tar.gz
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.
- deep_cnv-0.0.1/PKG-INFO +15 -0
- deep_cnv-0.0.1/README.md +0 -0
- deep_cnv-0.0.1/pyproject.toml +45 -0
- deep_cnv-0.0.1/pyproject.toml.orig +40 -0
- deep_cnv-0.0.1/src/deep_cnv/__init__.py +1 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/__init__.py +0 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/data_utils.py +38 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/preprocess.py +281 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/region_processing/__init__.py +0 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/region_processing/job_logic.py +116 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/region_processing/regions.py +181 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/run_preprocess.py +24 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/setup/__init__.py +0 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/setup/config.py +19 -0
- deep_cnv-0.0.1/src/deep_cnv/data_processing/setup/schemas.py +32 -0
- deep_cnv-0.0.1/src/deep_cnv/inference/__init__.py +0 -0
- deep_cnv-0.0.1/src/deep_cnv/inference/predict.py +66 -0
- deep_cnv-0.0.1/src/deep_cnv/inference/predict_batched.py +9 -0
- deep_cnv-0.0.1/src/deep_cnv/py.typed +0 -0
- deep_cnv-0.0.1/src/deep_cnv/training/__init__.py +0 -0
- deep_cnv-0.0.1/src/deep_cnv/training/data_augmentation.py +70 -0
- deep_cnv-0.0.1/src/deep_cnv/training/streaming_server.py +375 -0
- deep_cnv-0.0.1/src/deep_cnv/training/streaming_utils.py +0 -0
deep_cnv-0.0.1/PKG-INFO
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
Metadata-Version: 2.3
|
|
2
|
+
Name: deep-cnv
|
|
3
|
+
Version: 0.0.1
|
|
4
|
+
Summary: Deep learning for Copy Number Variation detection from SNP array data
|
|
5
|
+
Requires-Dist: eir-dl>=0.25.0
|
|
6
|
+
Requires-Dist: fastapi>=0.128.0
|
|
7
|
+
Requires-Dist: numpy>=1.26.0
|
|
8
|
+
Requires-Dist: pandas>=2.0.0
|
|
9
|
+
Requires-Dist: polars>=1.37.1
|
|
10
|
+
Requires-Dist: pydantic>=2.12.5
|
|
11
|
+
Requires-Dist: pyyaml>=6.0.3
|
|
12
|
+
Requires-Dist: uvicorn>=0.40.0
|
|
13
|
+
Requires-Python: >=3.13
|
|
14
|
+
Description-Content-Type: text/markdown
|
|
15
|
+
|
deep_cnv-0.0.1/README.md
ADDED
|
File without changes
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
[project]
|
|
2
|
+
name = "deep-cnv"
|
|
3
|
+
version = "0.0.1"
|
|
4
|
+
description = "Deep learning for Copy Number Variation detection from SNP array data"
|
|
5
|
+
readme = "README.md"
|
|
6
|
+
requires-python = ">=3.13"
|
|
7
|
+
dependencies = [
|
|
8
|
+
"eir-dl>=0.25.0",
|
|
9
|
+
"fastapi>=0.128.0",
|
|
10
|
+
"numpy>=1.26.0",
|
|
11
|
+
"pandas>=2.0.0",
|
|
12
|
+
"polars>=1.37.1",
|
|
13
|
+
"pydantic>=2.12.5",
|
|
14
|
+
"pyyaml>=6.0.3",
|
|
15
|
+
"uvicorn>=0.40.0",
|
|
16
|
+
]
|
|
17
|
+
|
|
18
|
+
[project.scripts]
|
|
19
|
+
deep-cnv-server = "deep_cnv.training.streaming_server:main"
|
|
20
|
+
deep-cnv-preprocess = "deep_cnv.data_processing.run_preprocess:main"
|
|
21
|
+
|
|
22
|
+
[dependency-groups]
|
|
23
|
+
dev = [
|
|
24
|
+
"pytest>=8.0.0",
|
|
25
|
+
"ruff>=0.4.0",
|
|
26
|
+
]
|
|
27
|
+
|
|
28
|
+
[build-system]
|
|
29
|
+
requires = ["uv_build>=0.12.19,<0.13"]
|
|
30
|
+
build-backend = "uv_build"
|
|
31
|
+
|
|
32
|
+
[tool.ruff]
|
|
33
|
+
line-length = 100
|
|
34
|
+
target-version = "py313"
|
|
35
|
+
|
|
36
|
+
[tool.ruff.lint]
|
|
37
|
+
select = [
|
|
38
|
+
"E",
|
|
39
|
+
"F",
|
|
40
|
+
"I",
|
|
41
|
+
"UP",
|
|
42
|
+
]
|
|
43
|
+
|
|
44
|
+
[tool.pytest.ini_options]
|
|
45
|
+
testpaths = ["tests"]
|
|
@@ -0,0 +1,40 @@
|
|
|
1
|
+
[project]
|
|
2
|
+
name = "deep-cnv"
|
|
3
|
+
version = "0.0.1"
|
|
4
|
+
description = "Deep learning for Copy Number Variation detection from SNP array data"
|
|
5
|
+
readme = "README.md"
|
|
6
|
+
requires-python = ">=3.13"
|
|
7
|
+
dependencies = [
|
|
8
|
+
"eir-dl>=0.25.0",
|
|
9
|
+
"fastapi>=0.128.0",
|
|
10
|
+
"numpy>=1.26.0",
|
|
11
|
+
"pandas>=2.0.0",
|
|
12
|
+
"polars>=1.37.1",
|
|
13
|
+
"pydantic>=2.12.5",
|
|
14
|
+
"pyyaml>=6.0.3",
|
|
15
|
+
"uvicorn>=0.40.0",
|
|
16
|
+
]
|
|
17
|
+
|
|
18
|
+
[dependency-groups]
|
|
19
|
+
dev = [
|
|
20
|
+
"pytest>=8.0.0",
|
|
21
|
+
"ruff>=0.4.0",
|
|
22
|
+
]
|
|
23
|
+
|
|
24
|
+
[project.scripts]
|
|
25
|
+
deep-cnv-server = "deep_cnv.training.streaming_server:main"
|
|
26
|
+
deep-cnv-preprocess = "deep_cnv.data_processing.run_preprocess:main"
|
|
27
|
+
|
|
28
|
+
[build-system]
|
|
29
|
+
requires = ["uv_build>=0.12.19,<0.13"]
|
|
30
|
+
build-backend = "uv_build"
|
|
31
|
+
|
|
32
|
+
[tool.ruff]
|
|
33
|
+
line-length = 100
|
|
34
|
+
target-version = "py313"
|
|
35
|
+
|
|
36
|
+
[tool.ruff.lint]
|
|
37
|
+
select = ["E", "F", "I", "UP"]
|
|
38
|
+
|
|
39
|
+
[tool.pytest.ini_options]
|
|
40
|
+
testpaths = ["tests"]
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
__version__ = "0.0.1"
|
|
File without changes
|
|
@@ -0,0 +1,38 @@
|
|
|
1
|
+
import gzip
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
|
|
4
|
+
import pandas as pd
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def load_cnvs(file_path: Path) -> pd.DataFrame:
|
|
8
|
+
df = pd.read_csv(file_path, sep="\t")
|
|
9
|
+
return df
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def read_tabix(
|
|
13
|
+
tabix_path: Path,
|
|
14
|
+
snppos: set[int] | None,
|
|
15
|
+
) -> pd.DataFrame:
|
|
16
|
+
with gzip.open(tabix_path, "rb") as f:
|
|
17
|
+
df = pd.read_csv(
|
|
18
|
+
f,
|
|
19
|
+
sep="\t",
|
|
20
|
+
header=None,
|
|
21
|
+
names=["chr", "start", "end", "lrr_raw", "baf", "lrr_adj"],
|
|
22
|
+
)
|
|
23
|
+
if snppos is not None:
|
|
24
|
+
df = df[df["start"].isin(snppos)]
|
|
25
|
+
|
|
26
|
+
return df.dropna(subset=["lrr_adj", "baf"]).reset_index(drop=True)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def load_samples_to_exclude_from_training(file_path: Path | None) -> set[str]:
|
|
30
|
+
if file_path is None:
|
|
31
|
+
return set()
|
|
32
|
+
|
|
33
|
+
if not file_path.exists():
|
|
34
|
+
raise FileNotFoundError(f"File not found: {file_path}")
|
|
35
|
+
|
|
36
|
+
df_test_set_to_skip = pd.read_csv(file_path, sep="\t")
|
|
37
|
+
|
|
38
|
+
return set(df_test_set_to_skip["sample_ID"])
|
|
@@ -0,0 +1,281 @@
|
|
|
1
|
+
from concurrent.futures import ProcessPoolExecutor
|
|
2
|
+
from dataclasses import dataclass, replace
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import pandas as pd
|
|
7
|
+
from eir.utils.logging import get_logger
|
|
8
|
+
|
|
9
|
+
from deep_cnv.data_processing.data_utils import (
|
|
10
|
+
load_cnvs,
|
|
11
|
+
load_samples_to_exclude_from_training,
|
|
12
|
+
read_tabix,
|
|
13
|
+
)
|
|
14
|
+
from deep_cnv.data_processing.region_processing.job_logic import (
|
|
15
|
+
CnvJob,
|
|
16
|
+
NormalJob,
|
|
17
|
+
build_cnv_jobs,
|
|
18
|
+
build_custom_cn2_jobs,
|
|
19
|
+
build_random_normal_jobs,
|
|
20
|
+
)
|
|
21
|
+
from deep_cnv.data_processing.region_processing.regions import (
|
|
22
|
+
CnvRegion,
|
|
23
|
+
RegionResult,
|
|
24
|
+
extract_cnv_region,
|
|
25
|
+
extract_random_normal_region,
|
|
26
|
+
save_region,
|
|
27
|
+
)
|
|
28
|
+
from deep_cnv.data_processing.setup.config import DataProcessingConfig, DataSourceConfig
|
|
29
|
+
|
|
30
|
+
logger = get_logger(name=__name__)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _extract_cnvs_from_sample(job: CnvJob) -> list[CnvRegion]:
|
|
34
|
+
try:
|
|
35
|
+
df_tabix = read_tabix(
|
|
36
|
+
tabix_path=job.tabix_path,
|
|
37
|
+
snppos=None,
|
|
38
|
+
)
|
|
39
|
+
except FileNotFoundError:
|
|
40
|
+
logger.error(f"Tabix file not found for sample {job.sample_id}: {job.tabix_path}")
|
|
41
|
+
return []
|
|
42
|
+
|
|
43
|
+
regions = []
|
|
44
|
+
|
|
45
|
+
for row in job.cnv_rows:
|
|
46
|
+
region = extract_cnv_region(
|
|
47
|
+
df_tabix=df_tabix,
|
|
48
|
+
chrom=int(row["chr"]),
|
|
49
|
+
cnv_start=int(row["start"]),
|
|
50
|
+
cnv_end=int(row["end"]),
|
|
51
|
+
cn_state=int(row["CN"]),
|
|
52
|
+
flank_snps=job.flank_snps,
|
|
53
|
+
)
|
|
54
|
+
if region is not None:
|
|
55
|
+
regions.append(region)
|
|
56
|
+
return regions
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def _extract_normal_regions_from_sample(job: NormalJob) -> list[CnvRegion]:
|
|
60
|
+
rng = np.random.default_rng(seed=job.seed)
|
|
61
|
+
try:
|
|
62
|
+
tabix_df = read_tabix(
|
|
63
|
+
tabix_path=job.tabix_path,
|
|
64
|
+
snppos=None,
|
|
65
|
+
)
|
|
66
|
+
except FileNotFoundError:
|
|
67
|
+
logger.error(f"Tabix file not found for sample {job.sample_id}: {job.tabix_path}.")
|
|
68
|
+
return []
|
|
69
|
+
|
|
70
|
+
regions = []
|
|
71
|
+
|
|
72
|
+
# informed CNV 2 case, these are visually confirmed to be 2
|
|
73
|
+
if job.cnv_rows is not None:
|
|
74
|
+
for row in job.cnv_rows:
|
|
75
|
+
region = extract_cnv_region(
|
|
76
|
+
df_tabix=tabix_df,
|
|
77
|
+
chrom=int(row["chr"]),
|
|
78
|
+
cnv_start=int(row["start"]),
|
|
79
|
+
cnv_end=int(row["end"]),
|
|
80
|
+
cn_state=2,
|
|
81
|
+
flank_snps=job.region_size // 2,
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
# this is a bit of a hack, but the streaming server expects cn_start and
|
|
85
|
+
# cn_end to be None for normal regions, i.e. we have this check
|
|
86
|
+
# "cnv" if pd.notna(row["cnv_start"]) else "normal" there
|
|
87
|
+
if region is not None:
|
|
88
|
+
regions.append(replace(region, cnv_start=None, cnv_end=None))
|
|
89
|
+
|
|
90
|
+
# purely random case
|
|
91
|
+
for _ in range(job.normal_regions_per_sample):
|
|
92
|
+
assert job.cnv_rows is None
|
|
93
|
+
region = extract_random_normal_region(
|
|
94
|
+
tabix_df=tabix_df,
|
|
95
|
+
region_size=job.region_size,
|
|
96
|
+
rng=rng,
|
|
97
|
+
)
|
|
98
|
+
if region is not None:
|
|
99
|
+
regions.append(region)
|
|
100
|
+
return regions
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
@dataclass
|
|
104
|
+
class SourceData:
|
|
105
|
+
cnvs: pd.DataFrame
|
|
106
|
+
samples: pd.DataFrame
|
|
107
|
+
samples_to_exclude: set[str]
|
|
108
|
+
all_sample_ids: set[str]
|
|
109
|
+
cnv_ids: set[str]
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def set_up_source_data(data_config: DataSourceConfig) -> SourceData:
|
|
113
|
+
dc = data_config
|
|
114
|
+
|
|
115
|
+
cnvs = load_cnvs(file_path=dc.cnvs_file)
|
|
116
|
+
cnvs["PN"] = cnvs["sample_ID"].str.split("_").str[0]
|
|
117
|
+
|
|
118
|
+
samples = pd.read_csv(filepath_or_buffer=dc.samples_file, sep="\t")
|
|
119
|
+
|
|
120
|
+
samples_to_exclude = load_samples_to_exclude_from_training(file_path=dc.test_set_file)
|
|
121
|
+
|
|
122
|
+
logger.info(f"Will exclude {len(samples_to_exclude)} PNs from training")
|
|
123
|
+
|
|
124
|
+
samples_filtered = samples[~samples["PN"].isin(samples_to_exclude)]
|
|
125
|
+
cnvs_filtered = cnvs[~cnvs["PN"].isin(samples_to_exclude)]
|
|
126
|
+
|
|
127
|
+
logger.info(f"Samples after filtering: {len(samples_filtered)}")
|
|
128
|
+
logger.info(f"CNVs after filtering: {len(cnvs_filtered)}")
|
|
129
|
+
|
|
130
|
+
all_sample_ids = set(samples_filtered["PN"])
|
|
131
|
+
cnv_ids = set(cnvs_filtered["PN"].astype(str).unique())
|
|
132
|
+
|
|
133
|
+
logger.info(
|
|
134
|
+
f"Samples with CNV (PN in CNV file): {len(cnv_ids)}, total samples: {len(all_sample_ids)}"
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
return SourceData(
|
|
138
|
+
cnvs=cnvs_filtered,
|
|
139
|
+
samples=samples_filtered,
|
|
140
|
+
samples_to_exclude=samples_to_exclude,
|
|
141
|
+
all_sample_ids=all_sample_ids,
|
|
142
|
+
cnv_ids=cnv_ids,
|
|
143
|
+
)
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def run_cnv_multiprocessing_extraction(
|
|
147
|
+
n_workers: int,
|
|
148
|
+
cnv_jobs: list[CnvJob],
|
|
149
|
+
output_dir: Path,
|
|
150
|
+
) -> list[RegionResult]:
|
|
151
|
+
region_idx = 0
|
|
152
|
+
|
|
153
|
+
records: list[RegionResult] = []
|
|
154
|
+
|
|
155
|
+
with ProcessPoolExecutor(max_workers=n_workers) as executor:
|
|
156
|
+
job_fn_mp_iter = zip(cnv_jobs, executor.map(_extract_cnvs_from_sample, cnv_jobs))
|
|
157
|
+
|
|
158
|
+
for job, regions in job_fn_mp_iter:
|
|
159
|
+
for region in regions:
|
|
160
|
+
region_id = f"cnv_{region_idx:06d}"
|
|
161
|
+
record = save_region(
|
|
162
|
+
region=region,
|
|
163
|
+
output_dir=output_dir / "train" / "cnv",
|
|
164
|
+
region_id=region_id,
|
|
165
|
+
sample_id=job.sample_id,
|
|
166
|
+
)
|
|
167
|
+
records.append(record)
|
|
168
|
+
region_idx += 1
|
|
169
|
+
|
|
170
|
+
logger.info(f"{region_idx} CNV regions saved")
|
|
171
|
+
|
|
172
|
+
return records
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
def run_normal_multiprocessing_extraction(
|
|
176
|
+
n_workers: int,
|
|
177
|
+
normal_jobs: list[NormalJob],
|
|
178
|
+
output_dir: Path,
|
|
179
|
+
) -> list[RegionResult]:
|
|
180
|
+
region_idx = 0
|
|
181
|
+
|
|
182
|
+
records: list[RegionResult] = []
|
|
183
|
+
|
|
184
|
+
with ProcessPoolExecutor(max_workers=n_workers) as executor:
|
|
185
|
+
job_fn_mp_iter = zip(
|
|
186
|
+
normal_jobs, executor.map(_extract_normal_regions_from_sample, normal_jobs)
|
|
187
|
+
)
|
|
188
|
+
|
|
189
|
+
for job, regions in job_fn_mp_iter:
|
|
190
|
+
for region in regions:
|
|
191
|
+
region_id = f"normal_{region_idx:06d}"
|
|
192
|
+
record = save_region(
|
|
193
|
+
region=region,
|
|
194
|
+
output_dir=output_dir / "train" / "normal",
|
|
195
|
+
region_id=region_id,
|
|
196
|
+
sample_id=job.sample_id,
|
|
197
|
+
)
|
|
198
|
+
records.append(record)
|
|
199
|
+
region_idx += 1
|
|
200
|
+
|
|
201
|
+
logger.info(f"{region_idx} normal regions saved")
|
|
202
|
+
|
|
203
|
+
return records
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def save_metadata(records: list[RegionResult], output_dir: Path) -> None:
|
|
207
|
+
df = pd.DataFrame(records)
|
|
208
|
+
df.to_csv(output_dir / "train" / "metadata.csv", index=False)
|
|
209
|
+
logger.info(f"Train: {len(records)} regions")
|
|
210
|
+
|
|
211
|
+
if len(df) > 0:
|
|
212
|
+
logger.info(f"cn_state counts:\n{df['cn_state'].value_counts().sort_index().to_string()}")
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
def process_dataset(
|
|
216
|
+
output_dir: Path,
|
|
217
|
+
data_config: DataSourceConfig,
|
|
218
|
+
processing_config: DataProcessingConfig,
|
|
219
|
+
) -> None:
|
|
220
|
+
dc = data_config
|
|
221
|
+
pc = processing_config
|
|
222
|
+
|
|
223
|
+
rng = np.random.default_rng(seed=pc.random_seed)
|
|
224
|
+
|
|
225
|
+
source_data = set_up_source_data(data_config=dc)
|
|
226
|
+
sd = source_data
|
|
227
|
+
|
|
228
|
+
(output_dir / "train" / "cnv").mkdir(parents=True, exist_ok=True)
|
|
229
|
+
(output_dir / "train" / "normal").mkdir(parents=True, exist_ok=True)
|
|
230
|
+
|
|
231
|
+
logger.info("Extracting CNV regions...")
|
|
232
|
+
|
|
233
|
+
id_to_path_map = dict(zip(sd.samples["PN"], sd.samples["file_path_tabix"]))
|
|
234
|
+
cnv_jobs = build_cnv_jobs(
|
|
235
|
+
cnvs=sd.cnvs,
|
|
236
|
+
id_to_path_map=id_to_path_map,
|
|
237
|
+
flank_snps=pc.flank_snps,
|
|
238
|
+
)
|
|
239
|
+
|
|
240
|
+
cnv_records = run_cnv_multiprocessing_extraction(
|
|
241
|
+
n_workers=pc.n_workers,
|
|
242
|
+
cnv_jobs=cnv_jobs,
|
|
243
|
+
output_dir=output_dir,
|
|
244
|
+
)
|
|
245
|
+
|
|
246
|
+
logger.info(
|
|
247
|
+
f"Extracting normal regions from {len(sd.all_sample_ids) - len(sd.cnv_ids)} samples..."
|
|
248
|
+
)
|
|
249
|
+
normal_sample_ids = sorted([sid for sid in sd.all_sample_ids if sid not in sd.cnv_ids])
|
|
250
|
+
|
|
251
|
+
random_jobs = build_random_normal_jobs(
|
|
252
|
+
sample_ids=normal_sample_ids,
|
|
253
|
+
id_to_path_map=id_to_path_map,
|
|
254
|
+
normal_regions_per_sample=pc.normal_regions_per_sample,
|
|
255
|
+
region_size=2 * pc.flank_snps,
|
|
256
|
+
rng=rng,
|
|
257
|
+
)
|
|
258
|
+
custom_jobs = build_custom_cn2_jobs(
|
|
259
|
+
id_to_path_map=id_to_path_map,
|
|
260
|
+
samples_to_exclude=sd.samples_to_exclude,
|
|
261
|
+
flank_snps=pc.flank_snps,
|
|
262
|
+
samples_file=dc.samples_file,
|
|
263
|
+
custom_normal_samples_file=dc.custom_normal_samples_file,
|
|
264
|
+
)
|
|
265
|
+
normal_jobs = random_jobs + custom_jobs
|
|
266
|
+
|
|
267
|
+
logger.info(
|
|
268
|
+
f"Built {len(random_jobs)} random normal jobs "
|
|
269
|
+
f"({pc.normal_regions_per_sample} regions each) "
|
|
270
|
+
f"and {len(custom_jobs)} custom CN2 jobs for multiprocessing..."
|
|
271
|
+
)
|
|
272
|
+
|
|
273
|
+
normal_records = run_normal_multiprocessing_extraction(
|
|
274
|
+
n_workers=pc.n_workers,
|
|
275
|
+
normal_jobs=normal_jobs,
|
|
276
|
+
output_dir=output_dir,
|
|
277
|
+
)
|
|
278
|
+
|
|
279
|
+
all_records = cnv_records + normal_records
|
|
280
|
+
|
|
281
|
+
save_metadata(records=all_records, output_dir=output_dir)
|
|
File without changes
|
|
@@ -0,0 +1,116 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
import pandas as pd
|
|
6
|
+
from eir.utils.logging import get_logger
|
|
7
|
+
|
|
8
|
+
logger = get_logger(name=__name__)
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
@dataclass(frozen=True)
|
|
12
|
+
class CnvJob:
|
|
13
|
+
sample_id: str
|
|
14
|
+
tabix_path: Path
|
|
15
|
+
cnv_rows: list[dict]
|
|
16
|
+
flank_snps: int
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def build_cnv_jobs(
|
|
20
|
+
cnvs: pd.DataFrame,
|
|
21
|
+
id_to_path_map: dict[str, str],
|
|
22
|
+
flank_snps: int,
|
|
23
|
+
) -> list[CnvJob]:
|
|
24
|
+
jobs = [
|
|
25
|
+
CnvJob(
|
|
26
|
+
sample_id=str(sid),
|
|
27
|
+
tabix_path=Path(id_to_path_map[sid].replace(".tbi", "")),
|
|
28
|
+
cnv_rows=sample_cnvs.to_dict("records"),
|
|
29
|
+
flank_snps=flank_snps,
|
|
30
|
+
)
|
|
31
|
+
for sid, sample_cnvs in cnvs.groupby("PN")
|
|
32
|
+
]
|
|
33
|
+
|
|
34
|
+
logger.info(f"Built {len(jobs)} CNV jobs for multiprocessing...")
|
|
35
|
+
|
|
36
|
+
return jobs
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass(frozen=True)
|
|
40
|
+
class NormalJob:
|
|
41
|
+
sample_id: str
|
|
42
|
+
tabix_path: Path
|
|
43
|
+
cnv_rows: list[dict] | None
|
|
44
|
+
normal_regions_per_sample: int
|
|
45
|
+
region_size: int
|
|
46
|
+
seed: int | None
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def build_random_normal_jobs(
|
|
50
|
+
sample_ids: list[str],
|
|
51
|
+
id_to_path_map: dict[str, str],
|
|
52
|
+
normal_regions_per_sample: int,
|
|
53
|
+
region_size: int,
|
|
54
|
+
rng: np.random.Generator,
|
|
55
|
+
) -> list[NormalJob]:
|
|
56
|
+
|
|
57
|
+
seeds = rng.integers(0, 2**31, size=len(sample_ids)).tolist()
|
|
58
|
+
|
|
59
|
+
random_normal_jobs = [
|
|
60
|
+
NormalJob(
|
|
61
|
+
sample_id=sid,
|
|
62
|
+
tabix_path=Path(id_to_path_map[sid].replace(".tbi", "")),
|
|
63
|
+
cnv_rows=None,
|
|
64
|
+
normal_regions_per_sample=normal_regions_per_sample,
|
|
65
|
+
region_size=region_size,
|
|
66
|
+
seed=int(seed),
|
|
67
|
+
)
|
|
68
|
+
for sid, seed in zip(sample_ids, seeds)
|
|
69
|
+
]
|
|
70
|
+
|
|
71
|
+
return random_normal_jobs
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def build_custom_cn2_jobs(
|
|
76
|
+
id_to_path_map: dict[str, str],
|
|
77
|
+
samples_to_exclude: set[str],
|
|
78
|
+
flank_snps: int,
|
|
79
|
+
samples_file: Path,
|
|
80
|
+
custom_normal_samples_file: Path | None,
|
|
81
|
+
) -> list[NormalJob]:
|
|
82
|
+
custom_file = custom_normal_samples_file
|
|
83
|
+
|
|
84
|
+
if custom_file is None:
|
|
85
|
+
logger.info("No custom CNV2 file passed in - will only use randomly drawn samples.")
|
|
86
|
+
return []
|
|
87
|
+
|
|
88
|
+
custom_cnv2_regions = pd.read_csv(custom_file, sep="\t")
|
|
89
|
+
custom_cnv2_regions["PN"] = custom_cnv2_regions["sample_ID"].str.split("_").str[0]
|
|
90
|
+
|
|
91
|
+
is_excluded = custom_cnv2_regions["PN"].isin(samples_to_exclude)
|
|
92
|
+
logger.info(f"Excluding {is_excluded.sum()} custom CNV2 regions belonging to test set PNs")
|
|
93
|
+
|
|
94
|
+
custom_cnv2_regions = custom_cnv2_regions[~is_excluded]
|
|
95
|
+
|
|
96
|
+
missing_pns = set(custom_cnv2_regions["PN"]) - set(id_to_path_map)
|
|
97
|
+
if missing_pns:
|
|
98
|
+
raise ValueError(
|
|
99
|
+
f"{len(missing_pns)} PNs in {custom_file} not found in {samples_file}: "
|
|
100
|
+
f"{sorted(missing_pns)[:10]}"
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
jobs = []
|
|
104
|
+
for sid, sample_cnv2_rows in custom_cnv2_regions.groupby("PN"):
|
|
105
|
+
jobs.append(
|
|
106
|
+
NormalJob(
|
|
107
|
+
sample_id=sid,
|
|
108
|
+
tabix_path=Path(id_to_path_map[sid].replace(".tbi", "")),
|
|
109
|
+
cnv_rows=sample_cnv2_rows.to_dict("records"),
|
|
110
|
+
normal_regions_per_sample=0,
|
|
111
|
+
region_size=2 * flank_snps,
|
|
112
|
+
seed=None,
|
|
113
|
+
)
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
return jobs
|
|
@@ -0,0 +1,181 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
import pandas as pd
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
@dataclass
|
|
9
|
+
class CnvRegion:
|
|
10
|
+
chrom: int
|
|
11
|
+
cn_state: int
|
|
12
|
+
signals: np.ndarray # shape (N, 2): LRR, BAF
|
|
13
|
+
positions: np.ndarray # shape (N,): genomic position per SNP
|
|
14
|
+
cnv_start: int | None = None # None for normal (CN2) regions
|
|
15
|
+
cnv_end: int | None = None # None for normal (CN2) regions
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def extract_cnv_region(
|
|
19
|
+
df_tabix: pd.DataFrame,
|
|
20
|
+
chrom: int,
|
|
21
|
+
cnv_start: int,
|
|
22
|
+
cnv_end: int,
|
|
23
|
+
cn_state: int,
|
|
24
|
+
flank_snps: int = 1000,
|
|
25
|
+
use_adjusted_lrr: bool = True,
|
|
26
|
+
) -> CnvRegion | None:
|
|
27
|
+
chrom_df = df_tabix[df_tabix["chr"] == chrom].sort_values("start")
|
|
28
|
+
if chrom_df.empty:
|
|
29
|
+
return None
|
|
30
|
+
|
|
31
|
+
positions = chrom_df["start"].values
|
|
32
|
+
|
|
33
|
+
cnv_start_idx = np.searchsorted(positions, cnv_start)
|
|
34
|
+
cnv_end_idx = np.searchsorted(positions, cnv_end, side="right")
|
|
35
|
+
|
|
36
|
+
region_start_idx = max(0, cnv_start_idx - flank_snps)
|
|
37
|
+
region_end_idx = min(len(positions), cnv_end_idx + flank_snps)
|
|
38
|
+
|
|
39
|
+
region_df = chrom_df.iloc[region_start_idx:region_end_idx]
|
|
40
|
+
lrr_col = "lrr_adj" if use_adjusted_lrr else "lrr_raw"
|
|
41
|
+
signals = region_df[[lrr_col, "baf"]].values.astype(np.float32)
|
|
42
|
+
|
|
43
|
+
return CnvRegion(
|
|
44
|
+
chrom=chrom,
|
|
45
|
+
cn_state=cn_state,
|
|
46
|
+
signals=signals,
|
|
47
|
+
positions=region_df["start"].values,
|
|
48
|
+
cnv_start=cnv_start,
|
|
49
|
+
cnv_end=cnv_end,
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def extract_random_normal_region(
|
|
54
|
+
tabix_df: pd.DataFrame,
|
|
55
|
+
region_size: int,
|
|
56
|
+
rng: np.random.Generator,
|
|
57
|
+
use_adjusted_lrr: bool = True,
|
|
58
|
+
) -> CnvRegion | None:
|
|
59
|
+
if tabix_df.empty:
|
|
60
|
+
return None
|
|
61
|
+
|
|
62
|
+
chromosomes = tabix_df["chr"].unique()
|
|
63
|
+
chrom = int(rng.choice(chromosomes))
|
|
64
|
+
|
|
65
|
+
chrom_df = tabix_df[tabix_df["chr"] == chrom].sort_values("start")
|
|
66
|
+
if len(chrom_df) < region_size:
|
|
67
|
+
return None
|
|
68
|
+
|
|
69
|
+
max_start_idx = len(chrom_df) - region_size
|
|
70
|
+
start_idx = int(rng.integers(0, max_start_idx + 1))
|
|
71
|
+
|
|
72
|
+
region_df = chrom_df.iloc[start_idx : start_idx + region_size]
|
|
73
|
+
lrr_col = "lrr_adj" if use_adjusted_lrr else "lrr_raw"
|
|
74
|
+
signals = region_df[[lrr_col, "baf"]].values.astype(np.float32)
|
|
75
|
+
|
|
76
|
+
return CnvRegion(
|
|
77
|
+
chrom=chrom,
|
|
78
|
+
cn_state=2,
|
|
79
|
+
signals=signals,
|
|
80
|
+
positions=region_df["start"].values,
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def load_region(region_path: Path, metadata_row: dict) -> CnvRegion:
|
|
85
|
+
data = np.load(file=region_path)
|
|
86
|
+
cnv_start = metadata_row["cnv_start"]
|
|
87
|
+
cnv_end = metadata_row["cnv_end"]
|
|
88
|
+
|
|
89
|
+
cnv_start = int(cnv_start) if cnv_start is not None and not np.isnan(cnv_start) else None
|
|
90
|
+
cnv_end = int(cnv_end) if cnv_end is not None and not np.isnan(cnv_end) else None
|
|
91
|
+
|
|
92
|
+
return CnvRegion(
|
|
93
|
+
chrom=int(metadata_row["chrom"]),
|
|
94
|
+
cn_state=int(metadata_row["cn_state"]),
|
|
95
|
+
signals=data["signals"],
|
|
96
|
+
positions=data["positions"],
|
|
97
|
+
cnv_start=cnv_start,
|
|
98
|
+
cnv_end=cnv_end,
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
def sample_window(
|
|
103
|
+
region: CnvRegion,
|
|
104
|
+
window_size: int,
|
|
105
|
+
rng: np.random.Generator,
|
|
106
|
+
) -> tuple[np.ndarray, np.ndarray] | None:
|
|
107
|
+
if len(region.positions) < window_size:
|
|
108
|
+
return None
|
|
109
|
+
|
|
110
|
+
max_start = len(region.positions) - window_size
|
|
111
|
+
start_idx = int(rng.integers(0, max_start + 1))
|
|
112
|
+
|
|
113
|
+
return (
|
|
114
|
+
region.signals[start_idx : start_idx + window_size],
|
|
115
|
+
region.positions[start_idx : start_idx + window_size],
|
|
116
|
+
)
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
def make_labels(
|
|
120
|
+
positions: np.ndarray,
|
|
121
|
+
cnv_start: int | None,
|
|
122
|
+
cnv_end: int | None,
|
|
123
|
+
cn_state: int,
|
|
124
|
+
) -> np.ndarray:
|
|
125
|
+
labels = np.full(len(positions), 2, dtype=np.int8)
|
|
126
|
+
|
|
127
|
+
if cnv_start is not None and cnv_end is not None:
|
|
128
|
+
mask = (positions >= cnv_start) & (positions <= cnv_end)
|
|
129
|
+
labels[mask] = cn_state
|
|
130
|
+
|
|
131
|
+
return labels
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
@dataclass
|
|
135
|
+
class RegionResult:
|
|
136
|
+
region_id: str
|
|
137
|
+
sample_id: str
|
|
138
|
+
chrom: int
|
|
139
|
+
cnv_start: int | None
|
|
140
|
+
cnv_end: int | None
|
|
141
|
+
cn_state: int
|
|
142
|
+
n_snps: int
|
|
143
|
+
n_flank_left: int | None
|
|
144
|
+
n_flank_right: int | None
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
def save_region(
|
|
148
|
+
region: CnvRegion,
|
|
149
|
+
output_dir: Path,
|
|
150
|
+
region_id: str,
|
|
151
|
+
sample_id: str,
|
|
152
|
+
) -> RegionResult:
|
|
153
|
+
np.savez_compressed(
|
|
154
|
+
output_dir / f"{region_id}.npz",
|
|
155
|
+
signals=region.signals,
|
|
156
|
+
positions=region.positions,
|
|
157
|
+
)
|
|
158
|
+
|
|
159
|
+
if region.cnv_start is not None and region.cnv_end is not None:
|
|
160
|
+
n_flank_left = int(np.searchsorted(region.positions, region.cnv_start))
|
|
161
|
+
n_flank_right = int(
|
|
162
|
+
len(region.positions) - np.searchsorted(region.positions, region.cnv_end, side="right")
|
|
163
|
+
)
|
|
164
|
+
|
|
165
|
+
else:
|
|
166
|
+
n_flank_left = None
|
|
167
|
+
n_flank_right = None
|
|
168
|
+
|
|
169
|
+
result = RegionResult(
|
|
170
|
+
region_id=region_id,
|
|
171
|
+
sample_id=sample_id,
|
|
172
|
+
chrom=region.chrom,
|
|
173
|
+
cnv_start=region.cnv_start,
|
|
174
|
+
cnv_end=region.cnv_end,
|
|
175
|
+
cn_state=region.cn_state,
|
|
176
|
+
n_snps=len(region.positions),
|
|
177
|
+
n_flank_left=n_flank_left,
|
|
178
|
+
n_flank_right=n_flank_right,
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
return result
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
|
|
4
|
+
from deep_cnv.data_processing.preprocess import process_dataset
|
|
5
|
+
from deep_cnv.data_processing.setup.config import load_data_config, load_processing_config
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
def main() -> None:
|
|
9
|
+
parser = argparse.ArgumentParser()
|
|
10
|
+
parser.add_argument("--data-config", type=Path, required=True)
|
|
11
|
+
parser.add_argument("--processing-config", type=Path, required=True)
|
|
12
|
+
parser.add_argument("--output-dir", type=Path, required=True)
|
|
13
|
+
args = parser.parse_args()
|
|
14
|
+
|
|
15
|
+
data_config = load_data_config(config_path=args.data_config)
|
|
16
|
+
processing_config = load_processing_config(config_path=args.processing_config)
|
|
17
|
+
|
|
18
|
+
process_dataset(
|
|
19
|
+
output_dir=args.output_dir, data_config=data_config, processing_config=processing_config
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
if __name__ == "__main__":
|
|
24
|
+
main()
|
|
File without changes
|
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
from yaml import safe_load
|
|
2
|
+
|
|
3
|
+
from deep_cnv.data_processing.setup.schemas import DataProcessingConfig, DataSourceConfig
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def load_data_config(config_path: str) -> DataSourceConfig:
|
|
7
|
+
with open(config_path) as f:
|
|
8
|
+
config_dict = safe_load(
|
|
9
|
+
stream=f,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
return DataSourceConfig(**config_dict)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def load_processing_config(config_path: str) -> DataProcessingConfig:
|
|
16
|
+
with open(config_path) as f:
|
|
17
|
+
config_dict = safe_load(stream=f)
|
|
18
|
+
|
|
19
|
+
return DataProcessingConfig(**config_dict)
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
@dataclass
|
|
6
|
+
class DataSourceConfig:
|
|
7
|
+
cnvs_file: Path
|
|
8
|
+
samples_file: Path
|
|
9
|
+
test_set_file: Path | None = None
|
|
10
|
+
custom_normal_samples_file: Path | None = None
|
|
11
|
+
|
|
12
|
+
def __post_init__(self) -> None:
|
|
13
|
+
str_to_path_keys = [
|
|
14
|
+
"cnvs_file",
|
|
15
|
+
"samples_file",
|
|
16
|
+
"test_set_file",
|
|
17
|
+
"custom_normal_samples_file",
|
|
18
|
+
]
|
|
19
|
+
for key in str_to_path_keys:
|
|
20
|
+
value = getattr(self, key)
|
|
21
|
+
if isinstance(value, str):
|
|
22
|
+
setattr(self, key, Path(value))
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@dataclass
|
|
26
|
+
class DataProcessingConfig:
|
|
27
|
+
""" """
|
|
28
|
+
|
|
29
|
+
flank_snps: int = 1000
|
|
30
|
+
normal_regions_per_sample: int = 3
|
|
31
|
+
random_seed: int = 42
|
|
32
|
+
n_workers: int = 8
|
|
File without changes
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
import base64
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
import pandas as pd
|
|
6
|
+
import requests
|
|
7
|
+
|
|
8
|
+
from deep_cnv.data_processing.region_processing.regions import CnvRegion, load_region
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def _encode_array(array: np.ndarray) -> str:
|
|
12
|
+
return base64.b64encode(array.tobytes()).decode("utf-8")
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _decode_array(encoded: str, dtype: np.dtype, shape: tuple) -> np.ndarray:
|
|
16
|
+
return np.frombuffer(base64.b64decode(encoded), dtype=dtype).reshape(shape)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def predict(
|
|
20
|
+
signals_batch: list[np.ndarray],
|
|
21
|
+
url: str = "http://localhost:8000/predict",
|
|
22
|
+
output_name: str = "snp_signals",
|
|
23
|
+
) -> list[np.ndarray]:
|
|
24
|
+
payload = []
|
|
25
|
+
for signals in signals_batch:
|
|
26
|
+
payload.append(
|
|
27
|
+
{
|
|
28
|
+
"snp_signals": _encode_array(array=signals),
|
|
29
|
+
}
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
response = requests.post(url=url, json=payload)
|
|
33
|
+
response.raise_for_status()
|
|
34
|
+
results = response.json()["result"]
|
|
35
|
+
|
|
36
|
+
predictions = []
|
|
37
|
+
for item in results:
|
|
38
|
+
raw = item[output_name]
|
|
39
|
+
n_bytes = len(base64.b64decode(raw))
|
|
40
|
+
n_floats = n_bytes // 4
|
|
41
|
+
window_size = signals_batch[0].shape[-1]
|
|
42
|
+
if n_floats == 5 * window_size:
|
|
43
|
+
shape = (5, window_size)
|
|
44
|
+
elif n_floats == 5 * 1 * window_size:
|
|
45
|
+
shape = (5, 1, window_size)
|
|
46
|
+
else:
|
|
47
|
+
raise ValueError(
|
|
48
|
+
f"Unexpected prediction size: {n_floats} floats (window_size={window_size})"
|
|
49
|
+
)
|
|
50
|
+
arr = _decode_array(encoded=raw, dtype=np.float32, shape=shape)
|
|
51
|
+
predictions.append(arr)
|
|
52
|
+
|
|
53
|
+
return predictions
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def load_regions_with_metadata(regions_dir: Path) -> list[tuple[CnvRegion, dict]]:
|
|
57
|
+
metadata = pd.read_csv(regions_dir / "metadata.csv")
|
|
58
|
+
regions = []
|
|
59
|
+
for _, row in metadata.iterrows():
|
|
60
|
+
subdir = "cnv" if pd.notna(row["cnv_start"]) else "normal"
|
|
61
|
+
path = regions_dir / subdir / f"{row['region_id']}.npz"
|
|
62
|
+
if not path.exists():
|
|
63
|
+
continue
|
|
64
|
+
region = load_region(region_path=path, metadata_row=row.to_dict())
|
|
65
|
+
regions.append((region, row.to_dict()))
|
|
66
|
+
return regions
|
|
File without changes
|
|
File without changes
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def augment_signals(
|
|
5
|
+
signals: np.ndarray,
|
|
6
|
+
rng: np.random.Generator,
|
|
7
|
+
lrr_shift_std: float = 0.05,
|
|
8
|
+
lrr_noise_std: float = 0.03,
|
|
9
|
+
lrr_scale_range: tuple[float, float] = (0.85, 1.15),
|
|
10
|
+
baf_jitter_std: float = 0.015,
|
|
11
|
+
probe_dropout_rate: float = 0.02,
|
|
12
|
+
lrr_mask_rate: float = 0.0,
|
|
13
|
+
baf_mask_rate: float = 0.0,
|
|
14
|
+
) -> np.ndarray:
|
|
15
|
+
signals = signals.copy()
|
|
16
|
+
n = len(signals)
|
|
17
|
+
lrr = signals[:, 0]
|
|
18
|
+
baf = signals[:, 1]
|
|
19
|
+
|
|
20
|
+
if lrr_mask_rate > 0 and rng.random() < lrr_mask_rate:
|
|
21
|
+
lrr[:] = rng.normal(scale=0.15, size=n).astype(np.float32)
|
|
22
|
+
else:
|
|
23
|
+
lrr += rng.normal(scale=lrr_shift_std)
|
|
24
|
+
lrr += rng.normal(scale=lrr_noise_std, size=n).astype(np.float32)
|
|
25
|
+
lrr *= rng.uniform(low=lrr_scale_range[0], high=lrr_scale_range[1])
|
|
26
|
+
|
|
27
|
+
if baf_mask_rate > 0 and rng.random() < baf_mask_rate:
|
|
28
|
+
baf[:] = rng.uniform(low=0.0, high=1.0, size=n).astype(np.float32)
|
|
29
|
+
else:
|
|
30
|
+
baf += rng.normal(scale=baf_jitter_std, size=n).astype(np.float32)
|
|
31
|
+
np.clip(baf, 0.0, 1.0, out=baf)
|
|
32
|
+
|
|
33
|
+
if probe_dropout_rate > 0:
|
|
34
|
+
drop_mask = rng.random(size=n) < probe_dropout_rate
|
|
35
|
+
lrr[drop_mask] = 0.0
|
|
36
|
+
baf[drop_mask] = 0.5
|
|
37
|
+
|
|
38
|
+
return signals
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def inject_fake_lrr_bump(
|
|
42
|
+
signals: np.ndarray,
|
|
43
|
+
rng: np.random.Generator,
|
|
44
|
+
min_snps: int = 50,
|
|
45
|
+
max_snps: int = 300,
|
|
46
|
+
shift_range: tuple[float, float] = (0.1, 0.4),
|
|
47
|
+
taper_snps: int = 10,
|
|
48
|
+
) -> np.ndarray:
|
|
49
|
+
signals = signals.copy()
|
|
50
|
+
n = len(signals)
|
|
51
|
+
lrr = signals[:, 0]
|
|
52
|
+
|
|
53
|
+
width = rng.integers(low=min_snps, high=min(max_snps, n) + 1)
|
|
54
|
+
start = rng.integers(low=0, high=n - width + 1)
|
|
55
|
+
end = start + width
|
|
56
|
+
|
|
57
|
+
shift = rng.uniform(low=shift_range[0], high=shift_range[1])
|
|
58
|
+
|
|
59
|
+
bump = np.zeros(n, dtype=np.float32)
|
|
60
|
+
bump[start:end] = shift
|
|
61
|
+
|
|
62
|
+
left_taper = min(taper_snps, width // 2)
|
|
63
|
+
right_taper = min(taper_snps, width // 2)
|
|
64
|
+
if left_taper > 0:
|
|
65
|
+
bump[start : start + left_taper] *= np.linspace(0, 1, left_taper, dtype=np.float32)
|
|
66
|
+
if right_taper > 0:
|
|
67
|
+
bump[end - right_taper : end] *= np.linspace(1, 0, right_taper, dtype=np.float32)
|
|
68
|
+
|
|
69
|
+
lrr += bump
|
|
70
|
+
return signals
|
|
@@ -0,0 +1,375 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
import base64
|
|
3
|
+
from collections.abc import Iterator
|
|
4
|
+
from concurrent.futures import ThreadPoolExecutor
|
|
5
|
+
from contextlib import asynccontextmanager
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from pathlib import Path
|
|
8
|
+
from threading import Lock
|
|
9
|
+
|
|
10
|
+
import numpy as np
|
|
11
|
+
import polars as pl
|
|
12
|
+
import uvicorn
|
|
13
|
+
from eir.setup.streaming_data_setup.protocol import PROTOCOL_VERSION
|
|
14
|
+
from eir.utils.logging import get_logger
|
|
15
|
+
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
|
16
|
+
from pydantic import BaseModel
|
|
17
|
+
|
|
18
|
+
from deep_cnv.data_processing.region_processing.regions import (
|
|
19
|
+
CnvRegion,
|
|
20
|
+
load_region,
|
|
21
|
+
make_labels,
|
|
22
|
+
sample_window,
|
|
23
|
+
)
|
|
24
|
+
from deep_cnv.training.data_augmentation import augment_signals, inject_fake_lrr_bump
|
|
25
|
+
|
|
26
|
+
logger = get_logger(name=__name__)
|
|
27
|
+
|
|
28
|
+
app = FastAPI()
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
@dataclass
|
|
32
|
+
class LoadArgs:
|
|
33
|
+
region_id: str
|
|
34
|
+
region_path: Path
|
|
35
|
+
metadata_row: dict
|
|
36
|
+
window_size: int
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _load_region_entry(load_args: LoadArgs) -> tuple[str, CnvRegion] | None:
|
|
40
|
+
la = load_args
|
|
41
|
+
|
|
42
|
+
if not la.region_path.exists():
|
|
43
|
+
return None
|
|
44
|
+
|
|
45
|
+
loaded_region = load_region(region_path=la.region_path, metadata_row=la.metadata_row)
|
|
46
|
+
|
|
47
|
+
if len(loaded_region.positions) < la.window_size:
|
|
48
|
+
return None
|
|
49
|
+
|
|
50
|
+
return (la.region_id, loaded_region)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def _build_load_args(
|
|
54
|
+
df_metadata: pl.DataFrame, regions_dir: Path, window_size: int
|
|
55
|
+
) -> Iterator[LoadArgs]:
|
|
56
|
+
for row in df_metadata.iter_rows(named=True):
|
|
57
|
+
region_id = row["region_id"]
|
|
58
|
+
cnv_start = row["cnv_start"]
|
|
59
|
+
|
|
60
|
+
region_folder = "cnv" if cnv_start is not None else "normal"
|
|
61
|
+
region_path = regions_dir / region_folder / f"{region_id}.npz"
|
|
62
|
+
|
|
63
|
+
yield LoadArgs(
|
|
64
|
+
region_id=region_id,
|
|
65
|
+
region_path=region_path,
|
|
66
|
+
metadata_row=row,
|
|
67
|
+
window_size=window_size,
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
class RegionDataset:
|
|
72
|
+
def __init__(
|
|
73
|
+
self,
|
|
74
|
+
regions_dir: Path,
|
|
75
|
+
window_size: int = 1024,
|
|
76
|
+
cnv_fraction: float = 0.7,
|
|
77
|
+
seed: int = 42,
|
|
78
|
+
n_workers: int = 4,
|
|
79
|
+
) -> None:
|
|
80
|
+
self.window_size = window_size
|
|
81
|
+
self.cnv_fraction = cnv_fraction
|
|
82
|
+
self.rng = np.random.default_rng(seed=seed)
|
|
83
|
+
self._lock = Lock()
|
|
84
|
+
self.position = 0
|
|
85
|
+
self.validation_ids: set[str] = set()
|
|
86
|
+
|
|
87
|
+
logger.info(f"Loading regions from {regions_dir}...")
|
|
88
|
+
metadata = pl.read_csv(source=regions_dir / "metadata.csv")
|
|
89
|
+
|
|
90
|
+
load_args_iterable = _build_load_args(
|
|
91
|
+
df_metadata=metadata,
|
|
92
|
+
regions_dir=regions_dir,
|
|
93
|
+
window_size=window_size,
|
|
94
|
+
)
|
|
95
|
+
load_args = list(load_args_iterable)
|
|
96
|
+
|
|
97
|
+
self.cnv_regions = []
|
|
98
|
+
self.normal_regions = []
|
|
99
|
+
with ThreadPoolExecutor(max_workers=n_workers) as executor:
|
|
100
|
+
for entry in executor.map(_load_region_entry, load_args):
|
|
101
|
+
if entry is None:
|
|
102
|
+
continue
|
|
103
|
+
_, region = entry
|
|
104
|
+
if region.cnv_start is not None:
|
|
105
|
+
self.cnv_regions.append(entry)
|
|
106
|
+
else:
|
|
107
|
+
self.normal_regions.append(entry)
|
|
108
|
+
|
|
109
|
+
skipped = sum(1 for a in load_args if not a.region_path.exists())
|
|
110
|
+
logger.info(
|
|
111
|
+
f"Loaded {len(self.cnv_regions)} CNV + {len(self.normal_regions)} normal regions "
|
|
112
|
+
f"({skipped} skipped), {len(metadata)} total in metadata"
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
def reset(self) -> None:
|
|
116
|
+
with self._lock:
|
|
117
|
+
self.position = 0
|
|
118
|
+
|
|
119
|
+
def get_batch(self, batch_size: int) -> list[dict]:
|
|
120
|
+
batch = []
|
|
121
|
+
attempts = 0
|
|
122
|
+
max_attempts = batch_size * 10
|
|
123
|
+
|
|
124
|
+
with self._lock:
|
|
125
|
+
while len(batch) < batch_size and attempts < max_attempts:
|
|
126
|
+
attempts += 1
|
|
127
|
+
sample_id = f"sample_{self.position}"
|
|
128
|
+
self.position += 1
|
|
129
|
+
|
|
130
|
+
if sample_id in self.validation_ids:
|
|
131
|
+
continue
|
|
132
|
+
|
|
133
|
+
if self.rng.random() < self.cnv_fraction:
|
|
134
|
+
pool = self.cnv_regions
|
|
135
|
+
else:
|
|
136
|
+
pool = self.normal_regions
|
|
137
|
+
region_id, region = pool[int(self.rng.integers(0, len(pool)))]
|
|
138
|
+
|
|
139
|
+
result = sample_window(
|
|
140
|
+
region=region,
|
|
141
|
+
window_size=self.window_size,
|
|
142
|
+
rng=self.rng,
|
|
143
|
+
)
|
|
144
|
+
if result is None:
|
|
145
|
+
continue
|
|
146
|
+
|
|
147
|
+
signals, positions = result
|
|
148
|
+
labels = make_labels(
|
|
149
|
+
positions=positions,
|
|
150
|
+
cnv_start=region.cnv_start,
|
|
151
|
+
cnv_end=region.cnv_end,
|
|
152
|
+
cn_state=region.cn_state,
|
|
153
|
+
)
|
|
154
|
+
|
|
155
|
+
signals = signals.astype(np.float32)
|
|
156
|
+
|
|
157
|
+
if self.rng.random() < 0.5:
|
|
158
|
+
signals = signals[::-1].copy()
|
|
159
|
+
labels = labels[::-1].copy()
|
|
160
|
+
|
|
161
|
+
signals = augment_signals(
|
|
162
|
+
signals=signals,
|
|
163
|
+
rng=self.rng,
|
|
164
|
+
lrr_mask_rate=_config["lrr_mask_rate"],
|
|
165
|
+
baf_mask_rate=_config["baf_mask_rate"],
|
|
166
|
+
)
|
|
167
|
+
|
|
168
|
+
fake_dup_rate = _config["fake_dup_rate"]
|
|
169
|
+
is_cn2 = np.all(labels == 2)
|
|
170
|
+
if fake_dup_rate > 0 and is_cn2 and self.rng.random() < fake_dup_rate:
|
|
171
|
+
signals = inject_fake_lrr_bump(
|
|
172
|
+
signals=signals,
|
|
173
|
+
rng=self.rng,
|
|
174
|
+
)
|
|
175
|
+
|
|
176
|
+
labels = labels.astype(np.float32)
|
|
177
|
+
labels_encoded = base64.b64encode(labels.tobytes()).decode("utf-8")
|
|
178
|
+
|
|
179
|
+
combined = np.transpose(signals, (1, 0))[:, np.newaxis, :]
|
|
180
|
+
inputs = {
|
|
181
|
+
"snp_signals": base64.b64encode(combined.tobytes()).decode("utf-8"),
|
|
182
|
+
}
|
|
183
|
+
|
|
184
|
+
output_name = _config["output_name"]
|
|
185
|
+
batch.append(
|
|
186
|
+
{
|
|
187
|
+
"inputs": inputs,
|
|
188
|
+
"target_labels": {
|
|
189
|
+
output_name: {
|
|
190
|
+
output_name: labels_encoded,
|
|
191
|
+
},
|
|
192
|
+
},
|
|
193
|
+
"sample_id": sample_id,
|
|
194
|
+
}
|
|
195
|
+
)
|
|
196
|
+
|
|
197
|
+
return batch
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
# ---------------------------------------------------------------------------
|
|
201
|
+
# Protocol models
|
|
202
|
+
# ---------------------------------------------------------------------------
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
class InputInfo(BaseModel):
|
|
206
|
+
type: str
|
|
207
|
+
shape: list[int] | None = None
|
|
208
|
+
|
|
209
|
+
|
|
210
|
+
class OutputInfo(BaseModel):
|
|
211
|
+
type: str
|
|
212
|
+
shape: list[int] | None = None
|
|
213
|
+
|
|
214
|
+
|
|
215
|
+
class DatasetInfo(BaseModel):
|
|
216
|
+
inputs: dict[str, InputInfo]
|
|
217
|
+
outputs: dict[str, OutputInfo]
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
# ---------------------------------------------------------------------------
|
|
221
|
+
# Server
|
|
222
|
+
# ---------------------------------------------------------------------------
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
_config: dict = {
|
|
226
|
+
"regions_dir": Path("data", "real", "regions", "train"),
|
|
227
|
+
"window_size": 1024,
|
|
228
|
+
"cnv_fraction": 0.7,
|
|
229
|
+
"output_name": "snp_signals",
|
|
230
|
+
"lrr_mask_rate": 0.0,
|
|
231
|
+
"baf_mask_rate": 0.0,
|
|
232
|
+
"fake_dup_rate": 0.0,
|
|
233
|
+
"n_workers": 4,
|
|
234
|
+
}
|
|
235
|
+
dataset: RegionDataset | None = None
|
|
236
|
+
|
|
237
|
+
|
|
238
|
+
@asynccontextmanager
|
|
239
|
+
async def lifespan(app: FastAPI):
|
|
240
|
+
global dataset
|
|
241
|
+
dataset = RegionDataset(
|
|
242
|
+
regions_dir=_config["regions_dir"],
|
|
243
|
+
window_size=_config["window_size"],
|
|
244
|
+
cnv_fraction=_config["cnv_fraction"],
|
|
245
|
+
n_workers=_config["n_workers"],
|
|
246
|
+
)
|
|
247
|
+
yield
|
|
248
|
+
|
|
249
|
+
|
|
250
|
+
app = FastAPI(lifespan=lifespan)
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
async def _handshake(websocket: WebSocket) -> bool:
|
|
254
|
+
await websocket.accept()
|
|
255
|
+
msg = await websocket.receive_json()
|
|
256
|
+
if msg.get("type") != "handshake" or msg.get("version") != PROTOCOL_VERSION:
|
|
257
|
+
await websocket.send_json(
|
|
258
|
+
{"type": "error", "payload": {"message": "Incompatible protocol version"}}
|
|
259
|
+
)
|
|
260
|
+
await websocket.close()
|
|
261
|
+
return False
|
|
262
|
+
await websocket.send_json({"type": "handshake", "version": PROTOCOL_VERSION})
|
|
263
|
+
return True
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
@app.websocket("/ws")
|
|
267
|
+
async def websocket_endpoint(websocket: WebSocket) -> None:
|
|
268
|
+
if not await _handshake(websocket=websocket):
|
|
269
|
+
return
|
|
270
|
+
|
|
271
|
+
try:
|
|
272
|
+
while True:
|
|
273
|
+
data = await websocket.receive_json()
|
|
274
|
+
msg_type = data.get("type")
|
|
275
|
+
|
|
276
|
+
if msg_type == "getInfo":
|
|
277
|
+
window_size = _config["window_size"]
|
|
278
|
+
output_name = _config["output_name"]
|
|
279
|
+
|
|
280
|
+
inputs = {
|
|
281
|
+
"snp_signals": InputInfo(type="array", shape=[2, 1, window_size]),
|
|
282
|
+
}
|
|
283
|
+
|
|
284
|
+
info = DatasetInfo(
|
|
285
|
+
inputs=inputs,
|
|
286
|
+
outputs={output_name: OutputInfo(type="array", shape=[1, 1, window_size])},
|
|
287
|
+
)
|
|
288
|
+
|
|
289
|
+
await websocket.send_json({"type": "info", "payload": info.model_dump()})
|
|
290
|
+
|
|
291
|
+
elif msg_type == "getData":
|
|
292
|
+
batch_size = data.get("payload", {}).get("batch_size", 32)
|
|
293
|
+
batch = dataset.get_batch(batch_size=batch_size)
|
|
294
|
+
|
|
295
|
+
if not batch:
|
|
296
|
+
await websocket.send_json({"type": "data", "payload": ["terminate"]})
|
|
297
|
+
break
|
|
298
|
+
|
|
299
|
+
await websocket.send_json({"type": "data", "payload": batch})
|
|
300
|
+
|
|
301
|
+
elif msg_type == "reset":
|
|
302
|
+
dataset.reset()
|
|
303
|
+
|
|
304
|
+
await websocket.send_json(
|
|
305
|
+
{"type": "resetConfirmation", "payload": {"message": "Reset successful"}}
|
|
306
|
+
)
|
|
307
|
+
|
|
308
|
+
await websocket.send_json(
|
|
309
|
+
{"type": "reset", "payload": {"message": "Reset command received"}}
|
|
310
|
+
)
|
|
311
|
+
|
|
312
|
+
elif msg_type == "setValidationIds":
|
|
313
|
+
validation_ids = data.get("payload", {}).get("validation_ids", [])
|
|
314
|
+
dataset.validation_ids = set(validation_ids)
|
|
315
|
+
|
|
316
|
+
await websocket.send_json(
|
|
317
|
+
{
|
|
318
|
+
"type": "validationIdsConfirmation",
|
|
319
|
+
"payload": {"message": f"Received {len(validation_ids)} validation IDs"},
|
|
320
|
+
}
|
|
321
|
+
)
|
|
322
|
+
|
|
323
|
+
elif msg_type == "status":
|
|
324
|
+
await websocket.send_json(
|
|
325
|
+
{
|
|
326
|
+
"type": "status",
|
|
327
|
+
"payload": {"current_position": dataset.position},
|
|
328
|
+
}
|
|
329
|
+
)
|
|
330
|
+
|
|
331
|
+
elif msg_type == "heartbeat":
|
|
332
|
+
await websocket.send_json({"type": "heartbeat"})
|
|
333
|
+
|
|
334
|
+
else:
|
|
335
|
+
logger.warning(f"Unknown message type: {msg_type}")
|
|
336
|
+
|
|
337
|
+
except WebSocketDisconnect:
|
|
338
|
+
logger.info("Client disconnected")
|
|
339
|
+
finally:
|
|
340
|
+
dataset.reset()
|
|
341
|
+
|
|
342
|
+
|
|
343
|
+
# ---------------------------------------------------------------------------
|
|
344
|
+
# Entry point
|
|
345
|
+
# ---------------------------------------------------------------------------
|
|
346
|
+
|
|
347
|
+
|
|
348
|
+
def main() -> None:
|
|
349
|
+
parser = argparse.ArgumentParser()
|
|
350
|
+
parser.add_argument("--host", default="0.0.0.0")
|
|
351
|
+
parser.add_argument("--port", type=int, default=8000)
|
|
352
|
+
parser.add_argument("--regions-dir", type=Path, default=_config["regions_dir"])
|
|
353
|
+
parser.add_argument("--window-size", type=int, default=_config["window_size"])
|
|
354
|
+
parser.add_argument("--cnv-fraction", type=float, default=_config["cnv_fraction"])
|
|
355
|
+
parser.add_argument("--output-name", type=str, default=_config["output_name"])
|
|
356
|
+
parser.add_argument("--lrr-mask-rate", type=float, default=_config["lrr_mask_rate"])
|
|
357
|
+
parser.add_argument("--baf-mask-rate", type=float, default=_config["baf_mask_rate"])
|
|
358
|
+
parser.add_argument("--fake-dup-rate", type=float, default=_config["fake_dup_rate"])
|
|
359
|
+
parser.add_argument("--n-workers", type=int, default=_config["n_workers"])
|
|
360
|
+
args = parser.parse_args()
|
|
361
|
+
|
|
362
|
+
_config["regions_dir"] = args.regions_dir
|
|
363
|
+
_config["window_size"] = args.window_size
|
|
364
|
+
_config["cnv_fraction"] = args.cnv_fraction
|
|
365
|
+
_config["output_name"] = args.output_name
|
|
366
|
+
_config["lrr_mask_rate"] = args.lrr_mask_rate
|
|
367
|
+
_config["baf_mask_rate"] = args.baf_mask_rate
|
|
368
|
+
_config["fake_dup_rate"] = args.fake_dup_rate
|
|
369
|
+
_config["n_workers"] = args.n_workers
|
|
370
|
+
|
|
371
|
+
uvicorn.run(app=app, host=args.host, port=args.port, ws_ping_timeout=3600)
|
|
372
|
+
|
|
373
|
+
|
|
374
|
+
if __name__ == "__main__":
|
|
375
|
+
main()
|
|
File without changes
|