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 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