vhamster 1.3.1__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.
- install_models.py +251 -0
- src/__init__.py +1 -0
- src/data.py +293 -0
- src/features.py +378 -0
- src/genomad_markers.py +515 -0
- src/models.py +822 -0
- src/predictor.py +108 -0
- src/sequences.py +298 -0
- vhamster-1.3.1.dist-info/METADATA +379 -0
- vhamster-1.3.1.dist-info/RECORD +15 -0
- vhamster-1.3.1.dist-info/WHEEL +5 -0
- vhamster-1.3.1.dist-info/entry_points.txt +4 -0
- vhamster-1.3.1.dist-info/licenses/LICENSE +19 -0
- vhamster-1.3.1.dist-info/top_level.txt +3 -0
- vhamster.py +1171 -0
install_models.py
ADDED
|
@@ -0,0 +1,251 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""
|
|
3
|
+
VHAMSTeR Model Installation Script
|
|
4
|
+
|
|
5
|
+
Downloads VHAMSTeR models from HuggingFace and the geNomad marker database
|
|
6
|
+
from Zenodo.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import gzip
|
|
10
|
+
import hashlib
|
|
11
|
+
import os
|
|
12
|
+
import shutil
|
|
13
|
+
import subprocess as sp
|
|
14
|
+
import sys
|
|
15
|
+
import sysconfig
|
|
16
|
+
import tarfile
|
|
17
|
+
import tempfile
|
|
18
|
+
import urllib.request
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
|
|
21
|
+
from huggingface_hub import snapshot_download
|
|
22
|
+
from loguru import logger
|
|
23
|
+
import click
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
HF_REPO_ID = "DOEJGI/vhamster-models"
|
|
27
|
+
|
|
28
|
+
# Only the MMseqs2 database and metadata are needed — vhamster does not use
|
|
29
|
+
# the HMM or MSA files from the full geNomad bundle.
|
|
30
|
+
GENOMAD_ZENODO_BASE = "https://zenodo.org/records/14886553/files"
|
|
31
|
+
GENOMAD_FILES = [
|
|
32
|
+
{
|
|
33
|
+
"filename": "genomad_db_v1.9.tar.gz",
|
|
34
|
+
"md5": "67244b528bb8bed464d1ca147136d33e",
|
|
35
|
+
"size_mb": 842,
|
|
36
|
+
},
|
|
37
|
+
{
|
|
38
|
+
"filename": "genomad_metadata_v1.9.tsv.gz",
|
|
39
|
+
"md5": "d4fa26b7a77017543bd80fb6bb4ee9d0",
|
|
40
|
+
"size_mb": 7,
|
|
41
|
+
},
|
|
42
|
+
]
|
|
43
|
+
|
|
44
|
+
REQUIRED_MODEL_FILES = ["fold_0", "fold_1", "fold_2", "fold_3", "fold_4"]
|
|
45
|
+
REQUIRED_ROOT_FILES = ["proportional_vector_scaling_scalar_nll_notclassbalanced_posthoc_fungi_nolength.json"]
|
|
46
|
+
DEFAULT_MODEL_DIRNAME = "vhamster_models_v1.3.0"
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def configure_logging(debug: bool = False):
|
|
50
|
+
logger.remove()
|
|
51
|
+
logger.add(sys.stderr, level="DEBUG" if debug else "INFO")
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def get_default_model_dir() -> str:
|
|
55
|
+
purelib = sysconfig.get_path("purelib")
|
|
56
|
+
if purelib is None:
|
|
57
|
+
script_dir = os.path.dirname(os.path.abspath(__file__))
|
|
58
|
+
logger.warning("Could not resolve environment site-packages; falling back to script directory")
|
|
59
|
+
return os.path.join(script_dir, DEFAULT_MODEL_DIRNAME)
|
|
60
|
+
env_model_dir = os.path.join(purelib, DEFAULT_MODEL_DIRNAME)
|
|
61
|
+
logger.info(f"Using environment-scoped default model location: {env_model_dir}")
|
|
62
|
+
return env_model_dir
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def check_model_installation(model_dir: str) -> bool:
|
|
66
|
+
for file_name in REQUIRED_ROOT_FILES:
|
|
67
|
+
file_path = os.path.join(model_dir, file_name)
|
|
68
|
+
if not os.path.isfile(file_path):
|
|
69
|
+
logger.warning(f"Required file missing: {file_path}")
|
|
70
|
+
return False
|
|
71
|
+
for fold_name in REQUIRED_MODEL_FILES:
|
|
72
|
+
fold_path = os.path.join(model_dir, fold_name)
|
|
73
|
+
if not os.path.isdir(fold_path):
|
|
74
|
+
logger.warning(f"Fold directory missing: {fold_path}")
|
|
75
|
+
return False
|
|
76
|
+
logger.info("All required model files are present")
|
|
77
|
+
return True
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def check_genomad_installation(genomad_db_dir: str) -> bool:
|
|
81
|
+
return os.path.isfile(os.path.join(genomad_db_dir, "genomad_marker_metadata.tsv")) and \
|
|
82
|
+
os.path.isfile(os.path.join(genomad_db_dir, "genomad_db"))
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def _md5(path: str) -> str:
|
|
86
|
+
h = hashlib.md5()
|
|
87
|
+
with open(path, "rb") as f:
|
|
88
|
+
for chunk in iter(lambda: f.read(1024 * 1024), b""):
|
|
89
|
+
h.update(chunk)
|
|
90
|
+
return h.hexdigest()
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def _download_file(url: str, dest: str, size_mb: int):
|
|
94
|
+
"""Download a file with a simple progress indicator."""
|
|
95
|
+
def reporthook(count, block_size, total_size):
|
|
96
|
+
if total_size > 0:
|
|
97
|
+
pct = min(count * block_size * 100 // total_size, 100)
|
|
98
|
+
sys.stdout.write(f"\r {pct}%")
|
|
99
|
+
sys.stdout.flush()
|
|
100
|
+
|
|
101
|
+
try:
|
|
102
|
+
logger.info(f"Downloading {os.path.basename(dest)} (~{size_mb} MB)")
|
|
103
|
+
urllib.request.urlretrieve(url, dest, reporthook)
|
|
104
|
+
sys.stdout.write("\n")
|
|
105
|
+
sys.stdout.flush()
|
|
106
|
+
except Exception as e:
|
|
107
|
+
logger.error(f"Download failed: {e}")
|
|
108
|
+
sys.exit(f"Could not download {url}\n{e}")
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def get_models_huggingface(model_dir: str):
|
|
112
|
+
abs_path = os.path.abspath(model_dir)
|
|
113
|
+
logger.info(f"Downloading VHAMSTeR models from HuggingFace ({HF_REPO_ID})")
|
|
114
|
+
try:
|
|
115
|
+
snapshot_download(repo_id=HF_REPO_ID, repo_type="model", local_dir=abs_path, revision='v1.3.0')
|
|
116
|
+
except Exception as e:
|
|
117
|
+
logger.error(f"Download failed: {e}")
|
|
118
|
+
sys.exit(f"Coul d not download models from HuggingFace.\n{e}")
|
|
119
|
+
logger.info("Model download complete.")
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def download_genomad_from_zenodo(genomad_db_dir: str):
|
|
123
|
+
"""Download and extract the geNomad marker database from Zenodo."""
|
|
124
|
+
os.makedirs(genomad_db_dir, exist_ok=True)
|
|
125
|
+
|
|
126
|
+
with tempfile.TemporaryDirectory() as tmp:
|
|
127
|
+
for entry in GENOMAD_FILES:
|
|
128
|
+
filename = entry["filename"]
|
|
129
|
+
url = f"{GENOMAD_ZENODO_BASE}/{filename}?download=1"
|
|
130
|
+
tmp_path = os.path.join(tmp, filename)
|
|
131
|
+
|
|
132
|
+
_download_file(url, tmp_path, entry["size_mb"])
|
|
133
|
+
|
|
134
|
+
logger.info(f"Verifying checksum for {filename}")
|
|
135
|
+
actual = _md5(tmp_path)
|
|
136
|
+
if actual != entry["md5"]:
|
|
137
|
+
sys.exit(
|
|
138
|
+
f"MD5 mismatch for {filename}.\n"
|
|
139
|
+
f" expected: {entry['md5']}\n"
|
|
140
|
+
f" got: {actual}"
|
|
141
|
+
)
|
|
142
|
+
|
|
143
|
+
if filename.endswith(".tar.gz"):
|
|
144
|
+
logger.info(f"Extracting {filename}")
|
|
145
|
+
with tarfile.open(tmp_path, "r:gz") as tar:
|
|
146
|
+
# Extract to a temp subdir so we can inspect the structure
|
|
147
|
+
extract_dir = os.path.join(tmp, "extracted")
|
|
148
|
+
os.makedirs(extract_dir, exist_ok=True)
|
|
149
|
+
for member in tar.getmembers():
|
|
150
|
+
if member.name.startswith("/") or ".." in member.name:
|
|
151
|
+
continue
|
|
152
|
+
tar.extract(member, path=extract_dir)
|
|
153
|
+
|
|
154
|
+
# If everything landed in a single top-level directory, use its contents
|
|
155
|
+
extracted_items = os.listdir(extract_dir)
|
|
156
|
+
if len(extracted_items) == 1 and os.path.isdir(os.path.join(extract_dir, extracted_items[0])):
|
|
157
|
+
extract_dir = os.path.join(extract_dir, extracted_items[0])
|
|
158
|
+
|
|
159
|
+
for item in os.listdir(extract_dir):
|
|
160
|
+
src = os.path.join(extract_dir, item)
|
|
161
|
+
dst = os.path.join(genomad_db_dir, item)
|
|
162
|
+
if os.path.exists(dst):
|
|
163
|
+
if os.path.isdir(dst):
|
|
164
|
+
shutil.rmtree(dst)
|
|
165
|
+
else:
|
|
166
|
+
os.remove(dst)
|
|
167
|
+
shutil.move(src, dst)
|
|
168
|
+
|
|
169
|
+
elif filename.endswith(".tsv.gz"):
|
|
170
|
+
out_name = filename.replace("_v1.9", "").replace(".gz", "")
|
|
171
|
+
# metadata file should be named genomad_marker_metadata.tsv
|
|
172
|
+
if "metadata" in filename:
|
|
173
|
+
out_name = "genomad_marker_metadata.tsv"
|
|
174
|
+
out_path = os.path.join(genomad_db_dir, out_name)
|
|
175
|
+
logger.info(f"Decompressing {filename} -> {os.path.basename(out_path)}")
|
|
176
|
+
with gzip.open(tmp_path, "rb") as f_in, open(out_path, "wb") as f_out:
|
|
177
|
+
shutil.copyfileobj(f_in, f_out)
|
|
178
|
+
|
|
179
|
+
logger.info(f"geNomad database installed at: {genomad_db_dir}")
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
def instantiate_install(model_dir: str, force: bool = False):
|
|
183
|
+
abs_path = os.path.abspath(model_dir)
|
|
184
|
+
os.makedirs(abs_path, exist_ok=True)
|
|
185
|
+
logger.info(f"Model installation directory: {abs_path}")
|
|
186
|
+
|
|
187
|
+
if check_model_installation(model_dir) and not force:
|
|
188
|
+
logger.info(f"All VHAMSTeR models already present in: {abs_path}")
|
|
189
|
+
else:
|
|
190
|
+
if force:
|
|
191
|
+
logger.info("Force reinstall requested.")
|
|
192
|
+
get_models_huggingface(model_dir)
|
|
193
|
+
|
|
194
|
+
|
|
195
|
+
def install_genomad(model_dir: str, force: bool = False):
|
|
196
|
+
genomad_db_dir = os.path.join(model_dir, "genomad_db")
|
|
197
|
+
|
|
198
|
+
if check_genomad_installation(genomad_db_dir) and not force:
|
|
199
|
+
logger.info(f"geNomad database already present at: {genomad_db_dir}")
|
|
200
|
+
return
|
|
201
|
+
|
|
202
|
+
download_genomad_from_zenodo(genomad_db_dir)
|
|
203
|
+
|
|
204
|
+
|
|
205
|
+
@click.command()
|
|
206
|
+
@click.option("-o", "--outdir", type=click.Path(path_type=str), default=None,
|
|
207
|
+
help="Directory to install models into (default: environment site-packages).")
|
|
208
|
+
@click.option("-f", "--force", is_flag=True, default=False,
|
|
209
|
+
help="Force reinstallation even if models already exist.")
|
|
210
|
+
@click.option("--debug", is_flag=True, default=False,
|
|
211
|
+
help="Enable verbose debug logging.")
|
|
212
|
+
def main(outdir, force, debug):
|
|
213
|
+
"""Download and install VHAMSTeR models from HuggingFace and the geNomad
|
|
214
|
+
marker database from Zenodo."""
|
|
215
|
+
configure_logging(debug)
|
|
216
|
+
|
|
217
|
+
model_dir = os.path.abspath(outdir) if outdir else get_default_model_dir()
|
|
218
|
+
logger.info(f"Model installation directory: {model_dir}")
|
|
219
|
+
|
|
220
|
+
instantiate_install(model_dir, force)
|
|
221
|
+
install_genomad(model_dir, force)
|
|
222
|
+
|
|
223
|
+
logger.info("\n" + "=" * 60)
|
|
224
|
+
logger.info("INSTALLATION SUMMARY")
|
|
225
|
+
logger.info("=" * 60)
|
|
226
|
+
|
|
227
|
+
fold_dirs = [d for d in os.listdir(model_dir) if d.startswith("fold_")]
|
|
228
|
+
if len(fold_dirs) == 5:
|
|
229
|
+
logger.info(f"✓ Fold models installed: {len(fold_dirs)}/5")
|
|
230
|
+
else:
|
|
231
|
+
logger.warning(f"✗ Fold models incomplete. Found: {len(fold_dirs)}/5")
|
|
232
|
+
|
|
233
|
+
calibration_file = os.path.join(model_dir, REQUIRED_ROOT_FILES[0])
|
|
234
|
+
if os.path.exists(calibration_file):
|
|
235
|
+
logger.info(f"✓ Calibration file present")
|
|
236
|
+
else:
|
|
237
|
+
logger.warning(f"✗ Calibration file not found: {calibration_file}")
|
|
238
|
+
|
|
239
|
+
genomad_db_dir = os.path.join(model_dir, "genomad_db")
|
|
240
|
+
if check_genomad_installation(genomad_db_dir):
|
|
241
|
+
logger.info(f"✓ geNomad database present")
|
|
242
|
+
else:
|
|
243
|
+
logger.warning(f"✗ geNomad database incomplete at: {genomad_db_dir}")
|
|
244
|
+
|
|
245
|
+
logger.info(f"\nTo run vhamster:")
|
|
246
|
+
logger.info(f" vhamster --fasta <input.fasta> --output <output_dir> --ensemble-dir {model_dir}")
|
|
247
|
+
logger.info("=" * 60)
|
|
248
|
+
|
|
249
|
+
|
|
250
|
+
if __name__ == "__main__":
|
|
251
|
+
main()
|
src/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
"""vhamster source package."""
|
src/data.py
ADDED
|
@@ -0,0 +1,293 @@
|
|
|
1
|
+
#!/usr/bin/env python3
|
|
2
|
+
"""
|
|
3
|
+
modules to handle data
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
# imports
|
|
7
|
+
import pathlib
|
|
8
|
+
import pandas as pd
|
|
9
|
+
import multiprocessing
|
|
10
|
+
import concurrent.futures
|
|
11
|
+
from Bio import SeqIO
|
|
12
|
+
from sklearn.model_selection import train_test_split
|
|
13
|
+
from typing import Tuple, List, Dict, Optional
|
|
14
|
+
import torch
|
|
15
|
+
from tqdm.auto import tqdm
|
|
16
|
+
from loguru import logger
|
|
17
|
+
import numpy as np
|
|
18
|
+
import random
|
|
19
|
+
import sequences
|
|
20
|
+
from torch.utils.data import Sampler
|
|
21
|
+
from features import ARCH_FEATURE_NAMES
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class EpochRedrawProkaryoteSampler(Sampler):
|
|
25
|
+
"""
|
|
26
|
+
Dynamically undersamples the prokaryotic class at each epoch.
|
|
27
|
+
At each epoch, randomly selects a subset of prokaryotic indices and combines with all other indices.
|
|
28
|
+
"""
|
|
29
|
+
def __init__(self, labels, prokaryote_idx, max_prop=0.5, random_seed=42):
|
|
30
|
+
self.labels = np.array(labels)
|
|
31
|
+
self.prokaryote_idx = prokaryote_idx
|
|
32
|
+
self.max_prop = max_prop
|
|
33
|
+
self.random_seed = random_seed
|
|
34
|
+
self.other_indices = np.where(self.labels != self.prokaryote_idx)[0]
|
|
35
|
+
self.prok_indices = np.where(self.labels == self.prokaryote_idx)[0]
|
|
36
|
+
self.n_other = len(self.other_indices)
|
|
37
|
+
self.n_prok_desired = int(self.n_other * self.max_prop / (1 - self.max_prop)) if self.max_prop < 1.0 else len(self.prok_indices)
|
|
38
|
+
self.epoch = 0
|
|
39
|
+
self._lat_indices = None
|
|
40
|
+
|
|
41
|
+
def set_epoch(self, epoch):
|
|
42
|
+
self.epoch = epoch
|
|
43
|
+
|
|
44
|
+
def __iter__(self):
|
|
45
|
+
random.seed(self.random_seed + self.epoch)
|
|
46
|
+
if self.n_prok_desired < len(self.prok_indices):
|
|
47
|
+
prok_sampled = random.sample(list(self.prok_indices), self.n_prok_desired)
|
|
48
|
+
else:
|
|
49
|
+
prok_sampled = list(self.prok_indices)
|
|
50
|
+
indices = list(self.other_indices) + prok_sampled
|
|
51
|
+
random.shuffle(indices)
|
|
52
|
+
self._lat_indices = indices
|
|
53
|
+
return iter(indices)
|
|
54
|
+
|
|
55
|
+
def __len__(self):
|
|
56
|
+
return len(self.other_indices) + min(self.n_prok_desired, len(self.prok_indices))
|
|
57
|
+
|
|
58
|
+
def get_category_counts(self):
|
|
59
|
+
"""Return a dictionary of class counts for the most recent sampled indices."""
|
|
60
|
+
if self._lat_indices is None:
|
|
61
|
+
raise ValueError("Sampler has not been iterated yet; no indices available.")
|
|
62
|
+
labels_sampled = self.labels[self._lat_indices]
|
|
63
|
+
unique, counts = np.unique(labels_sampled, return_counts=True)
|
|
64
|
+
return dict(zip(unique, counts))
|
|
65
|
+
|
|
66
|
+
def load_epoch_fold_data(epoch_dir, use_features=True, xgb_features_filename="xgb_probabilities.tsv"):
|
|
67
|
+
"""
|
|
68
|
+
Loads FASTA, labels, precomputed XGBoost probabilities, and raw arch/marker
|
|
69
|
+
features from a given epoch/fold directory.
|
|
70
|
+
|
|
71
|
+
The raw features (structural/architectural + marker frequencies) come from the
|
|
72
|
+
``*features.tsv`` file that is NOT ``xgb_probabilities.tsv``. They are used as
|
|
73
|
+
gate inputs in the ``GenomeClassifier`` to decouple gating from the model's own
|
|
74
|
+
XGBoost predictions and prevent training-set data leakage.
|
|
75
|
+
|
|
76
|
+
Returns:
|
|
77
|
+
seqs, labels, accs, idx_to_label, feature_names, features,
|
|
78
|
+
raw_feature_names, raw_features
|
|
79
|
+
(raw_feature_names and raw_features are None when no raw file is found)
|
|
80
|
+
"""
|
|
81
|
+
epoch_dir = pathlib.Path(epoch_dir)
|
|
82
|
+
fasta_file = next(epoch_dir.glob("*.fasta"))
|
|
83
|
+
labels_file = next(epoch_dir.glob("*labels.tsv"))
|
|
84
|
+
features_file = epoch_dir / xgb_features_filename
|
|
85
|
+
if use_features and not features_file.exists():
|
|
86
|
+
raise FileNotFoundError(
|
|
87
|
+
f"Expected precomputed {xgb_features_filename} in {epoch_dir} because use_features=True. "
|
|
88
|
+
"Run hierarchical XGBoost stacking first or disable feature usage explicitly."
|
|
89
|
+
)
|
|
90
|
+
if not features_file.exists():
|
|
91
|
+
features_file = None
|
|
92
|
+
|
|
93
|
+
# Locate raw arch/marker features file (any *features.tsv that isn't the XGB probs file)
|
|
94
|
+
_raw_feature_candidates = sorted(
|
|
95
|
+
p for p in epoch_dir.glob("*features.tsv")
|
|
96
|
+
if p.name != xgb_features_filename and p.name != "xgb_probabilities.tsv"
|
|
97
|
+
)
|
|
98
|
+
raw_features_file = _raw_feature_candidates[0] if _raw_feature_candidates else None
|
|
99
|
+
|
|
100
|
+
seqs, accs = sequences.load_fasta_sequences(fasta_file)
|
|
101
|
+
label_map = {}
|
|
102
|
+
with open(labels_file) as f:
|
|
103
|
+
header = f.readline()
|
|
104
|
+
for line in f:
|
|
105
|
+
parts = line.strip().split('\t')
|
|
106
|
+
if len(parts) < 2:
|
|
107
|
+
continue
|
|
108
|
+
acc, label = parts[0], parts[1]
|
|
109
|
+
label_map[acc] = label
|
|
110
|
+
labels = [label_map[a] for a in accs]
|
|
111
|
+
unique_labels = sorted(set(labels))
|
|
112
|
+
idx_to_label = {i: l for i, l in enumerate(unique_labels)}
|
|
113
|
+
label_to_idx = {l: i for i, l in idx_to_label.items()}
|
|
114
|
+
labels = [label_to_idx[l] for l in labels]
|
|
115
|
+
|
|
116
|
+
feature_names, features = None, None
|
|
117
|
+
if use_features and features_file is not None:
|
|
118
|
+
df = pd.read_csv(features_file, sep='\t')
|
|
119
|
+
id_col = None
|
|
120
|
+
for candidate in ("accession", "chunk_id", "id"):
|
|
121
|
+
if candidate in df.columns:
|
|
122
|
+
id_col = candidate
|
|
123
|
+
break
|
|
124
|
+
if id_col is None:
|
|
125
|
+
id_col = df.columns[0]
|
|
126
|
+
|
|
127
|
+
feature_names = [c for c in df.columns if c != id_col]
|
|
128
|
+
if not feature_names:
|
|
129
|
+
raise ValueError(
|
|
130
|
+
f"Precomputed feature file has no usable feature columns: {features_file}"
|
|
131
|
+
)
|
|
132
|
+
|
|
133
|
+
feature_df = df.set_index(id_col)[feature_names]
|
|
134
|
+
feature_df.index = feature_df.index.astype(str)
|
|
135
|
+
feature_df = feature_df.apply(pd.to_numeric, errors="coerce").fillna(0.0)
|
|
136
|
+
aligned = feature_df.reindex(accs).fillna(0.0)
|
|
137
|
+
features = aligned.to_numpy(dtype=np.float32)
|
|
138
|
+
|
|
139
|
+
# Load raw biological features for gate input (optional — not required to exist)
|
|
140
|
+
raw_feature_names, raw_features = None, None
|
|
141
|
+
if raw_features_file is not None:
|
|
142
|
+
try:
|
|
143
|
+
raw_df = pd.read_csv(raw_features_file, sep='\t')
|
|
144
|
+
raw_id_col = None
|
|
145
|
+
for candidate in ("accession", "chunk_id", "id"):
|
|
146
|
+
if candidate in raw_df.columns:
|
|
147
|
+
raw_id_col = candidate
|
|
148
|
+
break
|
|
149
|
+
if raw_id_col is None:
|
|
150
|
+
raw_id_col = raw_df.columns[0]
|
|
151
|
+
|
|
152
|
+
# Strictly use the 12 canonical arch features — in canonical order — as gate input.
|
|
153
|
+
# This guarantees training and inference receive the identical feature vector.
|
|
154
|
+
raw_feature_names = [c for c in ARCH_FEATURE_NAMES if c in raw_df.columns]
|
|
155
|
+
if raw_feature_names:
|
|
156
|
+
raw_feat_df = raw_df.set_index(raw_id_col)[raw_feature_names]
|
|
157
|
+
raw_feat_df.index = raw_feat_df.index.astype(str)
|
|
158
|
+
raw_feat_df = raw_feat_df.apply(pd.to_numeric, errors="coerce").fillna(0.0)
|
|
159
|
+
raw_aligned = raw_feat_df.reindex(accs).fillna(0.0)
|
|
160
|
+
raw_features = raw_aligned.to_numpy(dtype=np.float32)
|
|
161
|
+
else:
|
|
162
|
+
raw_feature_names = None
|
|
163
|
+
except Exception:
|
|
164
|
+
raw_feature_names, raw_features = None, None
|
|
165
|
+
|
|
166
|
+
return seqs, labels, accs, idx_to_label, feature_names, features, raw_feature_names, raw_features
|
|
167
|
+
|
|
168
|
+
def load_labels(labels_file: pathlib.Path) -> Tuple[Dict[str, str], List[str]]:
|
|
169
|
+
"""Load labels from TSV/CSV file."""
|
|
170
|
+
try:
|
|
171
|
+
df = pd.read_csv(labels_file, sep='\t', header=0)
|
|
172
|
+
except:
|
|
173
|
+
df = pd.read_csv(labels_file, header=0)
|
|
174
|
+
|
|
175
|
+
if len(df.columns) < 2:
|
|
176
|
+
raise ValueError(f"Labels file must have at least 2 columns (accession, label)")
|
|
177
|
+
|
|
178
|
+
accession_col = df.columns[0]
|
|
179
|
+
label_col = df.columns[1]
|
|
180
|
+
|
|
181
|
+
label_dict = dict(zip(df[accession_col].astype(str), df[label_col].astype(str)))
|
|
182
|
+
label_names = sorted(df[label_col].unique())
|
|
183
|
+
|
|
184
|
+
logger.info(f"Loaded {len(label_dict)} labels with {len(label_names)} classes: {label_names}")
|
|
185
|
+
return label_dict, label_names
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def load_data(
|
|
189
|
+
fasta_file: pathlib.Path,
|
|
190
|
+
labels_file: pathlib.Path,
|
|
191
|
+
test_size: float = 0.2,
|
|
192
|
+
random_seed: int = 42,
|
|
193
|
+
use_rv: bool = False,
|
|
194
|
+
use_features: bool = True,
|
|
195
|
+
use_class_weights: bool = False,
|
|
196
|
+
features_file: pathlib.Path = None,
|
|
197
|
+
output_dir: Optional[pathlib.Path] = None,
|
|
198
|
+
) -> Tuple[List[str], List[int], List[str], List[str], List[int], List[str], Dict[int, str], List[str], List[List[float]], List[List[float]], torch.Tensor]:
|
|
199
|
+
"""Load sequences and labels, split into train/val sets."""
|
|
200
|
+
label_dict, label_names = load_labels(labels_file)
|
|
201
|
+
label_to_idx = {name: idx for idx, name in enumerate(label_names)}
|
|
202
|
+
idx_to_label = {idx: name for name, idx in label_to_idx.items()}
|
|
203
|
+
|
|
204
|
+
sequences = []
|
|
205
|
+
labels = []
|
|
206
|
+
accessions = []
|
|
207
|
+
skipped = []
|
|
208
|
+
|
|
209
|
+
logger.info(f"Loading sequences from {fasta_file}...")
|
|
210
|
+
for record in tqdm(SeqIO.parse(str(fasta_file), "fasta"), desc="Loading sequences"):
|
|
211
|
+
accession = record.id.split()[0]
|
|
212
|
+
if accession not in label_dict:
|
|
213
|
+
skipped.append(accession)
|
|
214
|
+
continue
|
|
215
|
+
sequences.append(str(record.seq).upper())
|
|
216
|
+
labels.append(label_to_idx[label_dict[accession]])
|
|
217
|
+
accessions.append(accession)
|
|
218
|
+
|
|
219
|
+
if skipped:
|
|
220
|
+
logger.warning(f"Skipped {len(skipped)} sequences without labels")
|
|
221
|
+
|
|
222
|
+
if len(sequences) == 0:
|
|
223
|
+
raise ValueError("No sequences with labels found!")
|
|
224
|
+
|
|
225
|
+
logger.info(f"Loaded {len(sequences)} sequences with labels")
|
|
226
|
+
|
|
227
|
+
label_counts = pd.Series(labels).value_counts().sort_index()
|
|
228
|
+
logger.info("Class distribution:")
|
|
229
|
+
for idx, count in label_counts.items():
|
|
230
|
+
logger.info(f" {idx_to_label[idx]}: {count} ({100*count/len(labels):.1f}%)")
|
|
231
|
+
|
|
232
|
+
if use_class_weights:
|
|
233
|
+
total_samples = len(labels)
|
|
234
|
+
num_classes = len(idx_to_label)
|
|
235
|
+
weights = torch.tensor([total_samples / (num_classes * label_counts[idx]) for idx in sorted(idx_to_label.keys())], dtype=torch.float)
|
|
236
|
+
logger.info(f"Using class weights: {weights}")
|
|
237
|
+
else:
|
|
238
|
+
weights = None
|
|
239
|
+
|
|
240
|
+
# Feature Extraction
|
|
241
|
+
features_list = None
|
|
242
|
+
feature_names = []
|
|
243
|
+
if use_features:
|
|
244
|
+
if features_file and features_file.is_file():
|
|
245
|
+
logger.info(f"Loading pre-computed features from {features_file}...")
|
|
246
|
+
features_df = pd.read_csv(features_file)
|
|
247
|
+
feature_names = [col for col in features_df.columns if col != 'accession']
|
|
248
|
+
features_dict = {row['accession']: [row[name] for name in feature_names] for _, row in features_df.iterrows()}
|
|
249
|
+
features_list = []
|
|
250
|
+
for acc in accessions:
|
|
251
|
+
if acc not in features_dict:
|
|
252
|
+
raise ValueError(f"Accession {acc} not found in features file")
|
|
253
|
+
features_list.append(features_dict[acc])
|
|
254
|
+
else:
|
|
255
|
+
# Note: requires sequences.extract_features_worker to be importable here if running standard,
|
|
256
|
+
# but usually called via fine_tune_glm context.
|
|
257
|
+
# Assuming 'from features import extract_features_worker' logic in calling script or similar
|
|
258
|
+
# Since data.py imports 'sequences', make sure 'extract_features_worker' is available or import from features.py
|
|
259
|
+
from features import extract_features_worker, ARCH_FEATURE_NAMES_NO_FRAGMENT
|
|
260
|
+
|
|
261
|
+
feature_names = ARCH_FEATURE_NAMES_NO_FRAGMENT
|
|
262
|
+
logger.info(f"Extracting {len(feature_names)} features from {len(sequences)} sequences...")
|
|
263
|
+
|
|
264
|
+
num_workers = min(multiprocessing.cpu_count(), len(sequences))
|
|
265
|
+
logger.info(f"Using {num_workers} workers for parallel feature extraction...")
|
|
266
|
+
|
|
267
|
+
with concurrent.futures.ProcessPoolExecutor(max_workers=num_workers) as executor:
|
|
268
|
+
args_list = [(seq, acc, use_rv, feature_names, 10000) for seq, acc in zip(sequences, accessions)]
|
|
269
|
+
features_list = list(tqdm(
|
|
270
|
+
executor.map(extract_features_worker, args_list),
|
|
271
|
+
desc="Extracting features",
|
|
272
|
+
total=len(sequences)
|
|
273
|
+
))
|
|
274
|
+
else:
|
|
275
|
+
logger.info("Skipping feature extraction (--no-use-features)")
|
|
276
|
+
|
|
277
|
+
# Split
|
|
278
|
+
if use_features:
|
|
279
|
+
train_seqs, val_seqs, train_labels, val_labels, train_accs, val_accs, train_features, val_features = train_test_split(
|
|
280
|
+
sequences, labels, accessions, features_list,
|
|
281
|
+
test_size=test_size, random_state=random_seed, stratify=labels,
|
|
282
|
+
)
|
|
283
|
+
else:
|
|
284
|
+
train_seqs, val_seqs, train_labels, val_labels, train_accs, val_accs = train_test_split(
|
|
285
|
+
sequences, labels, accessions,
|
|
286
|
+
test_size=test_size, random_state=random_seed, stratify=labels,
|
|
287
|
+
)
|
|
288
|
+
train_features = None
|
|
289
|
+
val_features = None
|
|
290
|
+
|
|
291
|
+
logger.info(f"Split: {len(train_seqs)} train, {len(val_seqs)} validation")
|
|
292
|
+
|
|
293
|
+
return train_seqs, train_labels, train_accs, val_seqs, val_labels, val_accs, idx_to_label, feature_names, train_features, val_features, weights
|