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.
@@ -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
+
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"
@@ -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)
@@ -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()
@@ -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
@@ -0,0 +1,9 @@
1
+ _config: dict = {
2
+ "serve_url": "http://localhost:8001/predict",
3
+ }
4
+
5
+ if __name__ == "__main__":
6
+ import argparse
7
+
8
+ parser = argparse.ArgumentParser()
9
+ parser.add_argument("--serve-url", type=str, default=_config["serve_url"])
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()