diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 00000000..55221c61 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,46 @@ +# geniml -- Genomic Interval Machine Learning + +Builds vector embeddings and ML models from genomic interval data (BED files), enabling similarity search, clustering, and classification of genomic region sets. + +## Install + +``` +pip install geniml # base (no ML deps) +pip install geniml[ml] # torch, transformers, gensim +pip install geniml[sc] # scanpy, anndata +pip install geniml[search] # qdrant-client, fastembed +pip install geniml[all] # everything +``` + +## Quick start + +```python +from geniml.region2vec import Region2VecExModel + +model = Region2VecExModel("databio/r2v-ChIP-atlas-hg38") +vec = model.encode("peaks.bed") +``` + +## Submodules + +- **region2vec** -- Embed BED files into vectors (Region2VecExModel) +- **scembed** -- Single-cell embeddings (wraps region2vec for AnnData) +- **search** -- Vector similarity search (BED2BED, Text2BED; HNSW/Qdrant backends) +- **bbclient** -- Download BED files from BEDbase (BBClient) +- **atacformer** -- Transformer for scATAC-seq +- **geneformer** -- Transformer for scRNA-seq +- **craft** -- Contrastive model for gene activity +- **bedspace** -- StarSpace-based BED embedding (requires external binary) +- **assess/likelihood** -- Region set overlap, distance, and likelihood stats +- **io** -- BedSet class. For single region sets, prefer `gtars.models.RegionSet` + +## Dependency gating + +Heavy submodules (region2vec, scembed, atacformer, etc.) are NOT imported by `import geniml`. Import them directly: `from geniml.region2vec import Region2VecExModel`. They will fail with ImportError if the matching optional dep group is not installed. + +## Deprecated -- do not use + +- `geniml.io.Region`, `geniml.io.RegionSet` -- use `gtars.models.Region/RegionSet` +- `geniml.region2vec.main_legacy` -- replaced by Region2VecExModel +- `geniml.text2bednn` -- use `search.Text2BEDSearchInterface` +- `geniml.nn` -- internal utilities, not public API diff --git a/README.md b/README.md index 0068c22b..188deb8d 100644 --- a/README.md +++ b/README.md @@ -21,11 +21,14 @@ or install the latest version from the GitHub repository: pip install git+https://github.com/databio/geniml.git ``` -### To install Machine learning dependencies use this command: +### Optional dependency groups -From pypi: ``` -pip install geniml[ml] +pip install geniml[ml] # torch, transformers, gensim — for embeddings and ML models +pip install geniml[sc] # scanpy, anndata — for single-cell data processing +pip install geniml[search] # qdrant-client, fastembed — for vector search +pip install geniml[all] # everything (ml + sc + search) +pip install geniml[ml,sc] # ML + single-cell (for scembed, geneformer) ``` diff --git a/geniml/__init__.py b/geniml/__init__.py index fbaa69d9..f1c02e3a 100644 --- a/geniml/__init__.py +++ b/geniml/__init__.py @@ -1,6 +1,16 @@ +# Submodules with heavy optional dependencies (torch, scanpy, gensim) +# are NOT imported here. Import them directly: +# from geniml.region2vec import Region2VecExModel +# from geniml.scembed import ScEmbed +# Optional dependency groups: +# pip install geniml[ml] # torch, transformers, gensim — embeddings & ML models +# pip install geniml[sc] # scanpy, anndata — single-cell data processing +# pip install geniml[search] # qdrant-client, fastembed — vector search +# pip install geniml[all] # everything (ml + sc + search) + from logging import getLogger -from ._version import __version__ +from ._version import __version__ # noqa: F401 from .const import PKG_NAME _LOGGER = getLogger(PKG_NAME) diff --git a/geniml/assess/__init__.py b/geniml/assess/__init__.py index 20c0a5bd..2fd9886f 100644 --- a/geniml/assess/__init__.py +++ b/geniml/assess/__init__.py @@ -1 +1 @@ -from .cli import build_subparser +from .cli import build_subparser # noqa: F401 diff --git a/geniml/assess/likelihood.py b/geniml/assess/likelihood.py index efcbadd7..8fb176e4 100644 --- a/geniml/assess/likelihood.py +++ b/geniml/assess/likelihood.py @@ -275,10 +275,10 @@ def likelihood_flexible_universe( # likelihood of part of the genome after the last region res += background_likelihood( empty_start, - chr_size, - prob_start, - prob_core, - prob_end, + chr_size, # noqa: F821 + prob_start, # noqa: F821 + prob_core, # noqa: F821 + prob_end, # noqa: F821 ) current_chrom = i[0] done_chroms.append(current_chrom) diff --git a/geniml/assess/utils.py b/geniml/assess/utils.py index 2f1d7f30..2192c7fa 100644 --- a/geniml/assess/utils.py +++ b/geniml/assess/utils.py @@ -33,9 +33,9 @@ def check_if_uni_sorted(universe): def check_if_uni_flexible(universe): with open(universe) as u: - l = u.readline() - l = l.split("\t") - if len(l) < 6: + line = u.readline() + line = line.split("\t") + if len(line) < 6: raise Exception("Universe is not flexible") diff --git a/geniml/atacformer/modeling_atacformer.py b/geniml/atacformer/modeling_atacformer.py index 3edcfb02..a9bf9fc8 100644 --- a/geniml/atacformer/modeling_atacformer.py +++ b/geniml/atacformer/modeling_atacformer.py @@ -484,7 +484,6 @@ def forward( attention_mask_negative: Optional[torch.Tensor] = None, return_dict: Optional[bool] = None, ) -> Union[Tuple[torch.Tensor], BaseModelOutput]: - if attention_mask_anchor is None: attention_mask_anchor = torch.ones_like(input_ids_anchor, dtype=torch.bool) if attention_mask_positive is None: diff --git a/geniml/atacformer/training_utils.py b/geniml/atacformer/training_utils.py index e941378f..4ba01e24 100644 --- a/geniml/atacformer/training_utils.py +++ b/geniml/atacformer/training_utils.py @@ -1,11 +1,15 @@ +from __future__ import annotations + import math import subprocess from functools import partial -from typing import List, Dict +from typing import List, Dict, TYPE_CHECKING from collections import defaultdict import torch -import scanpy as sc + +if TYPE_CHECKING: + import scanpy as sc import numpy as np from torch.optim import Optimizer @@ -325,8 +329,7 @@ def __init__( ): super().__init__() try: - from sklearn.metrics import adjusted_rand_score - from sklearn.cluster import KMeans + import sklearn # noqa: F401 except ImportError: raise ImportError( "scikit-learn is required for AdjustedRandIndexCallback. Please install it with `pip install scikit-learn`." diff --git a/geniml/bedshift/__init__.py b/geniml/bedshift/__init__.py index d5ee4338..c1fa7e47 100644 --- a/geniml/bedshift/__init__.py +++ b/geniml/bedshift/__init__.py @@ -2,8 +2,8 @@ import logmuse -from .bedshift import Bedshift -from .yaml_handler import BedshiftYAMLHandler +from .bedshift import Bedshift # noqa: F401 +from .yaml_handler import BedshiftYAMLHandler # noqa: F401 __classes__ = ["Bedshift"] __all__ = __classes__ + [] diff --git a/geniml/bedshift/bedshift.py b/geniml/bedshift/bedshift.py index a8601990..4ce74da8 100644 --- a/geniml/bedshift/bedshift.py +++ b/geniml/bedshift/bedshift.py @@ -1,11 +1,12 @@ """Perturb regions in bedfiles""" import logging +import os import random +import tempfile -import genomicranges as gr import numpy as np -import pandas as pd +from gtars.models import RegionSet from .yaml_handler import BedshiftYAMLHandler @@ -14,29 +15,42 @@ __all__ = ["Bedshift"] +def _list_to_regionset(regions): + """Convert a list of lists to a RegionSet via a temporary BED file. + + Args: + regions (list): A list of lists, each containing [chrom, start, end, ...]. + + Returns: + RegionSet: A RegionSet constructed from the regions. + """ + with tempfile.NamedTemporaryFile(mode="w", suffix=".bed", delete=False) as f: + for r in regions: + f.write(f"{r[0]}\t{r[1]}\t{r[2]}\n") + tmp_path = f.name + rs = RegionSet(tmp_path) + os.unlink(tmp_path) + return rs + + class Bedshift(object): """The bedshift object with methods to perturb regions.""" - def __init__(self, bedfile_path, chrom_sizes=None, delimiter="\t"): - """Read in a .bed file to pandas DataFrame format. + def __init__(self, bedfile_path, chrom_sizes=None): + """Read in a .bed file to a list of lists. Args: bedfile_path (str): The path to the BED file. chrom_sizes (str): The path to the chrom.sizes file. - delimiter (str): The delimiter used in the BED file. """ self.bedfile_path = bedfile_path self.chrom_lens = {} if chrom_sizes: self._read_chromsizes(chrom_sizes) - df = self.read_bed(bedfile_path, delimiter=delimiter) - self.original_num_regions = df.shape[0] - self.bed = ( - df.astype({0: "object", 1: "int64", 2: "int64", 3: "object"}) - .sort_values([0, 1, 2]) - .reset_index(drop=True) - ) - self.original_bed = self.bed.copy() + self.bed = self.read_bed(bedfile_path) + self.original_num_regions = len(self.bed) + self.bed.sort(key=lambda r: (r[0], r[1], r[2])) + self.original_bed = [row[:] for row in self.bed] # deep copy def _read_chromsizes(self, fp): """Read chromosome sizes file. @@ -61,7 +75,7 @@ def _read_chromsizes(self, fp): def reset_bed(self): """Reset the stored bedfile to the state before perturbations.""" - self.bed = self.original_bed.copy() + self.bed = [row[:] for row in self.original_bed] def _precheck(self, rate, requiresChromLens=False, isAdd=False): """Check if the rate of perturbation is too high or low. @@ -87,6 +101,23 @@ def _precheck(self, rate, requiresChromLens=False, isAdd=False): _LOGGER.error(msg) raise FileNotFoundError(msg) + def _validate_region(self, start, end): + """Return True if the region is valid (start < end and start >= 0).""" + return start >= 0 and start < end + + def _remove_invalid_regions(self): + """Remove any regions where start >= end.""" + before = len(self.bed) + self.bed = [r for r in self.bed if r[1] < r[2]] + removed = before - len(self.bed) + if removed > 0: + _LOGGER.warning(f"Removed {removed} invalid regions (start >= end)") + return removed + + def _sort_bed(self): + """Sort bed by chromosome, start, end.""" + self.bed.sort(key=lambda r: (r[0], r[1], r[2])) + def pick_random_chroms(self, n): """Utility function to pick a random chromosome. @@ -100,7 +131,7 @@ def pick_random_chroms(self, n): chrom_lens = [self.chrom_lens[chrom_str] for chrom_str in chrom_strs] return zip(chrom_strs, chrom_lens) - def add(self, addrate, addmean, addstdev, valid_bed=None, delimiter="\t"): + def add(self, addrate, addmean, addstdev, valid_bed=None): """Add regions. Args: @@ -108,7 +139,6 @@ def add(self, addrate, addmean, addstdev, valid_bed=None, delimiter="\t"): addmean (float): The mean length of added regions. addstdev (float): The standard deviation of the length of added regions. valid_bed (str): The file with valid regions where new regions can be added. - delimiter (str): The delimiter used in valid_bed. Returns: int: The number of regions added. @@ -118,69 +148,71 @@ def add(self, addrate, addmean, addstdev, valid_bed=None, delimiter="\t"): else: self._precheck(addrate, requiresChromLens=True, isAdd=True) - rows = self.bed.shape[0] + rows = len(self.bed) num_add = int(rows * addrate) - new_regions = {0: [], 1: [], 2: [], 3: []} + new_rows = [] + if valid_bed: - valid_regions = self.read_bed(valid_bed, delimiter) - valid_regions[3] = valid_regions[2] - valid_regions[1] - total_bp = valid_regions[3].sum() - valid_regions[4] = valid_regions[3].apply(lambda x: x / total_bp) + valid_regions = self.read_bed(valid_bed) + total_bp = sum(r[2] - r[1] for r in valid_regions) + weights = [(r[2] - r[1]) / total_bp for r in valid_regions] add_rows = random.choices( list(range(len(valid_regions))), - weights=list(valid_regions[4]), + weights=weights, k=num_add, ) for row in add_rows: - data = valid_regions.loc[row] + data = valid_regions[row] chrom = data[0] start = random.randint(data[1], data[2]) - end = start + int(np.random.normal(addmean, addstdev)) - new_regions[0].append(chrom) - new_regions[1].append(start) - new_regions[2].append(end) - new_regions[3].append("A") + length = max(1, abs(int(np.random.normal(addmean, addstdev)))) + end = min(start + length, data[2]) + if end <= start: + end = start + 1 + new_rows.append([chrom, start, end, "A"]) else: random_chroms = self.pick_random_chroms(num_add) for chrom_str, chrom_len in random_chroms: start = random.randint(1, chrom_len) - # ensure chromosome length is not exceeded - end = min(start + int(np.random.normal(addmean, addstdev)), chrom_len) - new_regions[0].append(chrom_str) - new_regions[1].append(start) - new_regions[2].append(end) - new_regions[3].append("A") - self.bed = pd.concat([self.bed, pd.DataFrame(new_regions)], ignore_index=True) + length = max(1, abs(int(np.random.normal(addmean, addstdev)))) + end = min(start + length, chrom_len) + if end <= start: + end = start + 1 + new_rows.append([chrom_str, start, end, "A"]) + + self.bed.extend(new_rows) + self._sort_bed() return num_add - def add_from_file(self, fp, addrate, delimiter="\t"): + def add_from_file(self, fp, addrate): """Add regions from another bedfile to this perturbed bedfile. Args: fp (str): The filepath to the other bedfile. addrate (float): The rate to add regions. - delimiter (str): The delimiter used in the bedfile. Returns: int: The number of regions added. """ self._precheck(addrate, requiresChromLens=False, isAdd=True) - rows = self.bed.shape[0] + rows = len(self.bed) num_add = int(rows * addrate) - df = self.read_bed(fp, delimiter=delimiter) - dflen = len(df) - if num_add > dflen: + regions = self.read_bed(fp) + reglen = len(regions) + if num_add > reglen: _LOGGER.warning( "Number of regions to be added ({}) is larger than the provided bedfile size ({}). Adding {} regions.".format( - num_add, dflen, dflen + num_add, reglen, reglen ) ) - num_add = dflen - add_rows = random.sample(list(range(dflen)), num_add) - add_df = df.loc[add_rows].reset_index(drop=True) - add_df[3] = pd.Series(["A"] * add_df.shape[0]) - self.bed = pd.concat([self.bed, add_df], ignore_index=True) + num_add = reglen + add_indices = random.sample(list(range(reglen)), num_add) + for i in add_indices: + row = regions[i][:] + row[3] = "A" + self.bed.append(row) + self._sort_bed() return num_add def shift(self, shiftrate, shiftmean, shiftstdev, shift_rows=[]): @@ -197,7 +229,7 @@ def shift(self, shiftrate, shiftmean, shiftstdev, shift_rows=[]): """ self._precheck(shiftrate, requiresChromLens=True) - rows = self.bed.shape[0] + rows = len(self.bed) if len(shift_rows) == 0: shift_rows = random.sample(list(range(rows)), int(rows * shiftrate)) new_row_list = [] @@ -205,21 +237,20 @@ def shift(self, shiftrate, shiftmean, shiftstdev, shift_rows=[]): num_shifted = 0 invalid_shifted = 0 for row in shift_rows: - drop_row, new_region = self._shift( - row, shiftmean, shiftstdev - ) # shifted rows display a 1 + drop_row, new_region = self._shift(row, shiftmean, shiftstdev) if drop_row is not None and new_region: num_shifted += 1 new_row_list.append(new_region) to_drop.append(drop_row) else: invalid_shifted += 1 - self.bed = self.bed.drop(to_drop) - self.bed = pd.concat([self.bed, pd.DataFrame(new_row_list)], ignore_index=True) - self.bed = self.bed.reset_index(drop=True) + for idx in sorted(to_drop, reverse=True): + del self.bed[idx] + self.bed.extend(new_row_list) + self._sort_bed() if invalid_shifted > 0: _LOGGER.warning( - f"{invalid_shifted} regions were prevented from being shifted outside of chromosome boundaries. Reported regions shifted will be less than expected." + f"{invalid_shifted} regions were prevented from being shifted outside of chromosome boundaries." ) return num_shifted @@ -232,20 +263,23 @@ def _shift(self, row, mean, stdev): stdev (float): The standard deviation of the shift distance. Returns: - tuple: A tuple of (row_index, shifted_region_dict) or (None, None) if shift is invalid. + tuple: A tuple of (row_index, shifted_region_list) or (None, None) if shift is invalid. """ theshift = int(np.random.normal(mean, stdev)) - chrom = self.bed.loc[row][0] - start = self.bed.loc[row][1] - end = self.bed.loc[row][2] - if start + theshift < 0 or end + theshift > self.chrom_lens[str(chrom)]: - # check if the region is shifted out of chromosome length bounds + chrom = self.bed[row][0] + start = self.bed[row][1] + end = self.bed[row][2] + new_start = start + theshift + new_end = end + theshift + if new_start < 0 or new_end > self.chrom_lens[str(chrom)]: + return None, None + if new_start >= new_end: return None, None - return row, {0: chrom, 1: start + theshift, 2: end + theshift, 3: "S"} + return row, [chrom, new_start, new_end, "S"] - def shift_from_file(self, fp, shiftrate, shiftmean, shiftstdev, delimiter="\t"): + def shift_from_file(self, fp, shiftrate, shiftmean, shiftstdev): """Shift regions that overlap the specified file's regions. Args: @@ -253,38 +287,33 @@ def shift_from_file(self, fp, shiftrate, shiftmean, shiftstdev, delimiter="\t"): shiftrate (float): The rate to shift regions (both the start and end are shifted by the same amount). shiftmean (float): The mean shift distance. shiftstdev (float): The standard deviation of the shift distance. - delimiter (str): The delimiter used in fp. Returns: int: The number of regions shifted. """ self._precheck(shiftrate, requiresChromLens=True) - rows = self.bed.shape[0] + rows = len(self.bed) num_shift = int(rows * shiftrate) intersect_regions = self._find_overlap(fp) - original_colnames = self.bed.columns - intersect_regions.columns = [str(col) for col in intersect_regions.columns] - self.bed.columns = [str(col) for col in self.bed.columns] - indices_of_overlap_regions = self.bed.reset_index().merge(intersect_regions)["index"] - self.bed.columns = [int(col) for col in self.bed.columns] + intersect_set = {(r[0], r[1], r[2]) for r in intersect_regions} + indices_of_overlap = [ + i for i, r in enumerate(self.bed) if (r[0], r[1], r[2]) in intersect_set + ] - interlen = len(indices_of_overlap_regions) + interlen = len(indices_of_overlap) if num_shift > interlen: _LOGGER.warning( "Desired regions shifted ({}) is greater than the number of overlaps found ({}). Shifting {} regions.".format( num_shift, interlen, interlen ) ) - num_shift = len(indices_of_overlap_regions) - + num_shift = interlen elif interlen > num_shift: - indices_of_overlap_regions = indices_of_overlap_regions.sample(n=num_shift) - - indices_of_overlap_regions = indices_of_overlap_regions.to_list() + indices_of_overlap = random.sample(indices_of_overlap, num_shift) - return self.shift(shiftrate, shiftmean, shiftstdev, indices_of_overlap_regions) + return self.shift(shiftrate, shiftmean, shiftstdev, indices_of_overlap) def cut(self, cutrate): """Cut regions to create two new regions. @@ -297,18 +326,22 @@ def cut(self, cutrate): """ self._precheck(cutrate) - rows = self.bed.shape[0] + rows = len(self.bed) cut_rows = random.sample(list(range(rows)), int(rows * cutrate)) new_row_list = [] to_drop = [] + num_cut = 0 for row in cut_rows: - drop_row, new_regions = self._cut(row) # cut rows display a 2 - new_row_list.extend(new_regions) - to_drop.append(drop_row) - self.bed = self.bed.drop(to_drop) - self.bed = pd.concat([self.bed, pd.DataFrame(new_row_list)], ignore_index=True) - self.bed = self.bed.reset_index(drop=True) - return len(cut_rows) + drop_row, new_regions = self._cut(row) + if drop_row is not None and new_regions: + new_row_list.extend(new_regions) + to_drop.append(drop_row) + num_cut += 1 + for idx in sorted(to_drop, reverse=True): + del self.bed[idx] + self.bed.extend(new_row_list) + self._sort_bed() + return num_cut def _cut(self, row): """Cut a single region into two regions. @@ -317,30 +350,23 @@ def _cut(self, row): row (int): The index of the row to cut. Returns: - tuple: A tuple of (row_index, list_of_two_new_regions). - """ - chrom = self.bed.loc[row][0] - start = self.bed.loc[row][1] - end = self.bed.loc[row][2] - - # choose where to cut the region - thecut = (start + end) // 2 # int(np.random.normal((start+end)/2, (end - start)/6)) - if thecut <= start: - thecut = start + 10 - if thecut >= end: - thecut = end - 10 - - """ may add in later, this makes the api confusing! - # adjust the cut regions using the shift function - new_segs = self.__shift(new_segs, 0, meanshift, stdevshift) - new_segs = self.__shift(new_segs, 1, meanshift, stdevshift) + tuple: A tuple of (row_index, list_of_two_new_regions) or (None, None) if region is too small. """ + chrom = self.bed[row][0] + start = self.bed[row][1] + end = self.bed[row][2] + + # Region must be at least 2bp to cut into two valid regions + if end - start < 2: + return None, None + + thecut = random.randint(start + 1, end - 1) return ( row, [ - {0: chrom, 1: start, 2: thecut, 3: "C"}, - {0: chrom, 1: thecut, 2: end, 3: "C"}, + [chrom, start, thecut, "C"], + [chrom, thecut, end, "C"], ], ) @@ -355,7 +381,8 @@ def merge(self, mergerate): """ self._precheck(mergerate) - rows = self.bed.shape[0] + self._sort_bed() + rows = len(self.bed) merge_rows = random.sample(list(range(rows)), int(rows * mergerate)) to_add = [] to_drop = [] @@ -364,9 +391,10 @@ def merge(self, mergerate): if drop_rows and add_row: to_add.append(add_row) to_drop.extend(drop_rows) - self.bed = self.bed.drop(to_drop) - self.bed = pd.concat([self.bed, pd.DataFrame(to_add)], ignore_index=True) - self.bed = self.bed.reset_index(drop=True) + for idx in sorted(set(to_drop), reverse=True): + del self.bed[idx] + self.bed.extend(to_add) + self._sort_bed() return len(to_drop) def _merge(self, row): @@ -376,16 +404,15 @@ def _merge(self, row): row (int): The index of the row to merge. Returns: - tuple: A tuple of (list_of_rows_to_drop, merged_region_dict) or (None, None) if merge is invalid. + tuple: A tuple of (list_of_rows_to_drop, merged_region_list) or (None, None) if merge is invalid. """ - # check if the regions being merged are on the same chromosome - if row + 1 not in self.bed.index or self.bed.loc[row][0] != self.bed.loc[row + 1][0]: + if row + 1 >= len(self.bed) or self.bed[row][0] != self.bed[row + 1][0]: return None, None - chrom = self.bed.loc[row][0] - start = self.bed.loc[row][1] - end = self.bed.loc[row + 1][2] - return [row, row + 1], {0: chrom, 1: start, 2: end, 3: "M"} + chrom = self.bed[row][0] + start = min(self.bed[row][1], self.bed[row + 1][1]) + end = max(self.bed[row][2], self.bed[row + 1][2]) + return [row, row + 1], [chrom, start, end, "M"] def drop(self, droprate): """Drop regions. @@ -398,49 +425,48 @@ def drop(self, droprate): """ self._precheck(droprate) - rows = self.bed.shape[0] + rows = len(self.bed) drop_rows = random.sample(list(range(rows)), int(rows * droprate)) - self.bed = self.bed.drop(drop_rows) - self.bed = self.bed.reset_index(drop=True) + for idx in sorted(drop_rows, reverse=True): + del self.bed[idx] + self._sort_bed() return len(drop_rows) - def drop_from_file(self, fp, droprate, delimiter="\t"): + def drop_from_file(self, fp, droprate): """Drop regions that overlap between the reference bedfile and the provided bedfile. Args: fp (str): The filepath to the other bedfile containing regions to be dropped. droprate (float): The rate to drop regions. - delimiter (str): The delimiter used in the bedfile. Returns: int: The number of regions dropped. """ self._precheck(droprate) - rows = self.bed.shape[0] + rows = len(self.bed) num_drop = int(rows * droprate) - drop_bed = self.read_bed(fp, delimiter=delimiter) + drop_bed = self.read_bed(fp) intersect_regions = self._find_overlap(drop_bed) - # original_colnames = self.bed.columns - intersect_regions.columns = [str(col) for col in intersect_regions.columns] - self.bed.columns = [str(col) for col in self.bed.columns] - indices_of_overlap_regions = self.bed.reset_index().merge(intersect_regions)["index"] - self.bed.columns = [int(col) for col in self.bed.columns] + intersect_set = {(r[0], r[1], r[2]) for r in intersect_regions} + indices_of_overlap = [ + i for i, r in enumerate(self.bed) if (r[0], r[1], r[2]) in intersect_set + ] - interlen = len(indices_of_overlap_regions) + interlen = len(indices_of_overlap) if num_drop > interlen: _LOGGER.warning( "Desired regions dropped ({}) is greater than the number of overlaps found ({}). Dropping {} regions.".format( num_drop, interlen, interlen ) ) - num_drop = len(indices_of_overlap_regions) + num_drop = interlen elif interlen > num_drop: - indices_of_overlap_regions = indices_of_overlap_regions.sample(n=num_drop) - indices_of_overlap_regions = indices_of_overlap_regions.to_list() + indices_of_overlap = random.sample(indices_of_overlap, num_drop) - self.bed = self.bed.drop(indices_of_overlap_regions) + for idx in sorted(indices_of_overlap, reverse=True): + del self.bed[idx] return num_drop def set_seed(self, seednum): @@ -465,44 +491,47 @@ def _find_overlap(self, fp, reference=None): """Find intersecting regions between the reference bedfile and the comparison file. Args: - fp (str or pd.DataFrame): Path to file, or pandas DataFrame, for comparison. - reference (str or pd.DataFrame): Path to file, or pandas DataFrame, for reference. If None, then defaults to the original BED file provided to the Bedshift constructor. + fp (str or list): Path to file, or list of lists, for comparison. + reference (str or list): Path to file, or list of lists, for reference. + If None, then defaults to the original BED file provided to the Bedshift constructor. Returns: - pd.DataFrame: A DataFrame of overlapping regions. + list: A list of [chrom, start, end] lists representing overlapping regions. """ + # Build reference region data if reference is None: - reference_bed = self.original_bed.copy() + ref_data = self.original_bed + elif isinstance(reference, list): + ref_data = reference + elif isinstance(reference, str): + ref_data = self.read_bed(reference) else: - if isinstance(reference, pd.DataFrame): - reference_bed = reference.copy() - elif isinstance(reference, str): - reference_bed = self.read_bed(reference) - else: - raise Exception("unsupported input type: {}".format(type(reference))) - if isinstance(fp, pd.DataFrame): - comparison_bed = fp.copy() + raise Exception("unsupported input type: {}".format(type(reference))) + + # Build comparison region data + if isinstance(fp, list): + comp_data = fp elif isinstance(fp, str): - comparison_bed = self.read_bed(fp) + comp_data = self.read_bed(fp) else: - raise Exception("unsupported input type: {}".format(type(reference))) - reference_bed.columns = ["seqnames", "starts", "ends", "modifications"] - comparison_bed.columns = ["seqnames", "starts", "ends", "modifications"] + raise Exception("unsupported input type: {}".format(type(fp))) - reference_gr = gr.GenomicRanges.from_pandas(reference_bed) - comparison_gr = gr.GenomicRanges.from_pandas(comparison_bed) - intersection_gr = reference_gr.subset_by_overlaps(comparison_gr) - intersection = intersection_gr.to_pandas() + # Convert list-of-lists to RegionSet via tempfile + ref_rs = _list_to_regionset(ref_data) + comp_rs = _list_to_regionset(comp_data) - if len(intersection) == 0: - raise Exception( - "no intersection found between {} and {}".format(reference_bed, comparison_bed) - ) + # Use RegionSet overlap detection + overlap_rs = ref_rs.subset_by_overlaps(comp_rs) - intersection = intersection[["seqnames", "starts", "ends"]] - intersection.columns = [0, 1, 2] + if len(overlap_rs) == 0: + raise Exception("no intersection found") - return intersection + # Convert back to list of lists + result = [] + for i in range(len(overlap_rs)): + region = overlap_rs[i] + result.append([region.chr, region.start, region.end]) + return result def all_perturbations( self, @@ -569,47 +598,42 @@ def all_perturbations( else: n += self.drop(droprate) + self._remove_invalid_regions() return n def to_bed(self, outfile_name): - """Write a pandas dataframe back into BED file format. + """Write regions to a BED file. Args: outfile_name (str): The name of the output BED file. """ - self.bed.sort_values([0, 1, 2], inplace=True) - self.bed.to_csv(outfile_name, sep="\t", header=False, index=False, float_format="%.0f") + self._remove_invalid_regions() + self._sort_bed() + with open(outfile_name, "w") as f: + for row in self.bed: + f.write(f"{row[0]}\t{int(row[1])}\t{int(row[2])}\n") - def read_bed(self, bedfile_path, delimiter="\t"): - """Read a BED file into pandas dataframe. + def read_bed(self, bedfile_path): + """Read a BED file into a list of lists. Args: bedfile_path (str): The path to the BED file. - delimiter (str): The delimiter used in the BED file. Returns: - pd.DataFrame: The BED file as a pandas DataFrame. + list: A list of lists, each containing [chrom, start, end, mod_flag]. """ try: - df = pd.read_csv( - bedfile_path, - sep=delimiter, - header=None, - usecols=[0, 1, 2], - engine="python", - ) - except FileNotFoundError: - msg = "BED file path {} invalid".format(bedfile_path) - _LOGGER.error(msg) - raise FileNotFoundError(msg) - except: + rs = RegionSet(bedfile_path) + except Exception: msg = "File {} could not be read".format(bedfile_path) _LOGGER.error(msg) raise Exception(msg) - # if there is a header line in the table, remove it - if not str(df.iloc[0, 1]).isdigit(): - df = df[1:].reset_index(drop=True) + if len(rs) == 0: + raise Exception(f"File {bedfile_path} is empty") - df[3] = "-" # column indicating which modifications were made - return df + regions = [] + for i in range(len(rs)): + region = rs[i] + regions.append([region.chr, region.start, region.end, "-"]) + return regions diff --git a/geniml/bedshift/yaml_handler.py b/geniml/bedshift/yaml_handler.py index 8812422a..8cfcbefb 100644 --- a/geniml/bedshift/yaml_handler.py +++ b/geniml/bedshift/yaml_handler.py @@ -38,7 +38,6 @@ def _print_sample_config(self): - drop_from_file: file: tests/test.bed rate: 0.1 - delimiter: \\t - shift_from_file: file: bedshifted_test.bed rate: 0.3 @@ -78,6 +77,8 @@ def handle_yaml(self): """Perform perturbations from the YAML configuration file. Executes perturbations specified in the YAML config file in the order they were provided. + The 'delimiter' key is accepted in YAML configs for backwards compatibility but ignored + since RegionSet handles BED parsing internally. Returns: int: The total number of regions changed by all perturbations. @@ -87,48 +88,36 @@ def handle_yaml(self): num_changed = 0 for operation in operations: + # Strip 'delimiter' key if present (no longer used, RegionSet handles parsing) + op_keys = set(operation.keys()) - {"delimiter"} + ##### add ##### - if set(["add", "rate", "mean", "stdev"]) == set(list(operation.keys())): + if op_keys == {"add", "rate", "mean", "stdev"}: rate = operation["rate"] mean = operation["mean"] std = operation["stdev"] num_added = self.bedshifter.add(rate, mean, std) num_changed += num_added - ##### add_from_file with no delimiter provided ##### - elif set(["add_from_file", "file", "rate"]) == set(list(operation.keys())): - # fp = operation["file"] + ##### add_from_file ##### + elif op_keys == {"add_from_file", "file", "rate"}: fp = os.path.expandvars(operation["file"]) if os.path.isfile(fp): add_rate = operation["rate"] num_added = self.bedshifter.add_from_file(fp, add_rate) num_changed += num_added else: - self._logger.error("File '{}' does not exist.".format(fp)) - sys.exit(1) - - ##### add_from_file with delimiter provided ##### - elif set(["add_from_file", "file", "rate", "delimiter"]) == set( - list(operation.keys()) - ): - fp = os.path.expandvars(operation["file"]) - if os.path.isfile(fp): - add_rate = operation["rate"] - delimiter = operation["delimiter"] - num_added = self.bedshifter.add_from_file(fp, add_rate, delimiter) - num_changed += num_added - else: - self._logger.error("File '{}' does not exist.".format(fp)) + self._LOGGER.error("File '{}' does not exist.".format(fp)) sys.exit(1) ##### drop ##### - elif set(["drop", "rate"]) == set(list(operation.keys())): + elif op_keys == {"drop", "rate"}: rate = operation["rate"] num_dropped = self.bedshifter.drop(rate) num_changed += num_dropped - ##### drop_from_file with no delimiter provided ##### - elif set(["drop_from_file", "file", "rate"]) == set(list(operation.keys())): + ##### drop_from_file ##### + elif op_keys == {"drop_from_file", "file", "rate"}: fp = os.path.expandvars(operation["file"]) if os.path.isfile(fp): drop_rate = operation["rate"] @@ -138,22 +127,8 @@ def handle_yaml(self): self._LOGGER.error("File '{}' does not exist.".format(fp)) sys.exit(1) - ##### drop_from_file with delimiter provided ##### - elif set(["drop_from_file", "file", "rate", "delimiter"]) == set( - list(operation.keys()) - ): - fp = os.path.expandvars(operation["file"]) - if os.path.isfile(fp): - drop_rate = operation["rate"] - delimiter = operation["delimiter"] - num_dropped = self.bedshifter.drop_from_file(fp, drop_rate, delimiter) - num_changed += num_dropped - else: - self._LOGGER.error("File '{}' does not exist.".format(fp)) - sys.exit(1) - ##### shift ##### - elif set(["shift", "rate", "mean", "stdev"]) == set(list(operation.keys())): + elif op_keys == {"shift", "rate", "mean", "stdev"}: rate = operation["rate"] mean = operation["mean"] std = operation["stdev"] @@ -161,9 +136,7 @@ def handle_yaml(self): num_changed += num_shifted ##### shift_from_file ##### - elif set(["shift_from_file", "file", "rate", "mean", "stdev"]) == set( - list(operation.keys()) - ): + elif op_keys == {"shift_from_file", "file", "rate", "mean", "stdev"}: fp = os.path.expandvars(operation["file"]) if os.path.isfile(fp): rate = operation["rate"] @@ -175,30 +148,14 @@ def handle_yaml(self): self._LOGGER.error("File '{}' does not exist.".format(fp)) sys.exit(1) - ##### shift_from_file with delimiter provided ##### - elif set(["shift_from_file", "file", "rate", "mean", "stdev", "delimiter"]) == set( - list(operation.keys()) - ): - fp = os.path.expandvars(operation["file"]) - if os.path.isfile(fp): - rate = operation["rate"] - mean = operation["mean"] - std = operation["stdev"] - delimiter = operation["delimiter"] - num_shifted = self.bedshifter.shift_from_file(fp, rate, mean, std, delimiter) - num_changed += num_shifted - else: - self._LOGGER.error("File '{}' does not exist.".format(fp)) - sys.exit(1) - ##### cut ##### - elif set(["cut", "rate"]) == set(list(operation.keys())): + elif op_keys == {"cut", "rate"}: rate = operation["rate"] num_cut = self.bedshifter.cut(rate) num_changed += num_cut ##### merge ##### - elif set(["merge", "rate"]) == set(list(operation.keys())): + elif op_keys == {"merge", "rate"}: rate = operation["rate"] num_merged = self.bedshifter.merge(rate) num_changed += num_merged diff --git a/geniml/bedspace/__init__.py b/geniml/bedspace/__init__.py index 8b137891..a463f25b 100644 --- a/geniml/bedspace/__init__.py +++ b/geniml/bedspace/__init__.py @@ -1 +1,23 @@ +from .const import SearchType +from .search import run_scenario1, run_scenario2, run_scenario3 +from .helpers import ( + meta_preprocessing, + data_preparation, + bed2vec, + get_label_embedding, + get_embedding_matrix, + calculate_distance, +) +__all__ = [ + "SearchType", + "run_scenario1", + "run_scenario2", + "run_scenario3", + "meta_preprocessing", + "data_preparation", + "bed2vec", + "get_label_embedding", + "get_embedding_matrix", + "calculate_distance", +] diff --git a/geniml/bedspace/argparsers.py b/geniml/bedspace/argparsers.py index c46266be..c1e574b0 100644 --- a/geniml/bedspace/argparsers.py +++ b/geniml/bedspace/argparsers.py @@ -1,6 +1,6 @@ from ubiquerg import VersionInHelpParser -from .const import * +from .const import SearchType def build_preprocess_argparser( diff --git a/geniml/bedspace/cli.py b/geniml/bedspace/cli.py index 2f210439..242c0cbd 100644 --- a/geniml/bedspace/cli.py +++ b/geniml/bedspace/cli.py @@ -8,7 +8,7 @@ from .argparsers import build_preprocess_argparser as preprocess_subparser from .argparsers import build_search_argparser as search_subparser from .argparsers import build_train_argparser as train_subparser -from .const import * +from .const import DISTANCES_CMD, PKG_NAME, PREPROCESS_CMD, SEARCH_CMD, TRAIN_CMD global _LOGGER logging.basicConfig(level=logging.INFO) diff --git a/geniml/bedspace/helpers.py b/geniml/bedspace/helpers.py index 365a6b98..fb7ca29d 100644 --- a/geniml/bedspace/helpers.py +++ b/geniml/bedspace/helpers.py @@ -109,9 +109,9 @@ def get_label_embedding(path_word_embedding, label_prefix): # Filter rows that contain the label prefix vectors = word_embedding[word_embedding[0].str.contains(label_prefix)] # .reset_index() # Extract label vectors and labels - for l in range(len(vectors)): - label_vectors.append((list(vectors.iloc[l])[1:])) - labels.append(list(vectors.iloc[l])[0].replace(label_prefix, "")) + for idx in range(len(vectors)): + label_vectors.append((list(vectors.iloc[idx])[1:])) + labels.append(list(vectors.iloc[idx])[0].replace(label_prefix, "")) return label_vectors, labels diff --git a/geniml/bedspace/search.py b/geniml/bedspace/search.py index 1960e2ab..248635f7 100644 --- a/geniml/bedspace/search.py +++ b/geniml/bedspace/search.py @@ -12,7 +12,7 @@ def run_scenario1( distances: str, output: str, num_results: int = DEFAULT_NUM_SEARCH_RESULTS, -): +) -> None: """Run the search command for scenario 1: Give me a label, I'll return region sets. Args: @@ -60,7 +60,7 @@ def run_scenario2( distances: str, output: str, num_results: int = DEFAULT_NUM_SEARCH_RESULTS, -): +) -> None: """Run the search command for scenario 2: Give me a region set, I'll return labels. Args: @@ -73,7 +73,6 @@ def run_scenario2( _LOGGER.info("Running search...") # PLACE SEARCH CODE HERE - file = query distance = pd.read_csv(distances) distance.file_label = distance.file_label.str.lower() distance.search_term = distance.search_term.str.lower() @@ -105,7 +104,7 @@ def run_scenario3( distances: str, output: str, num_results: int = DEFAULT_NUM_SEARCH_RESULTS, -): +) -> None: """Run the search command for scenario 3: Give me a region set, I'll return region sets. Args: diff --git a/geniml/bedspace/visualization.py b/geniml/bedspace/visualization.py index a349a7be..38dd071d 100644 --- a/geniml/bedspace/visualization.py +++ b/geniml/bedspace/visualization.py @@ -8,7 +8,7 @@ matplotlib.rcParams["svg.fonttype"] = "none" matplotlib.rcParams["text.usetex"] = False -import matplotlib.pyplot as plt +import matplotlib.pyplot as plt # noqa: E402 # Label embedding @@ -18,8 +18,6 @@ def label_preprocessing(path_label_embedding, label_prefix, common_labels=[]): - labels = [] - label_vectors = [] label_embedding = pd.read_csv(path_label_embedding, sep="\t", header=None, skiprows=1) vectors = label_embedding[label_embedding[0].str.contains(label_prefix)] # .reset_index() @@ -43,7 +41,6 @@ def UMAP_plot( output_folder="", ): np.random.seed(3) - dp = 400 ump = umap.UMAP( a=None, @@ -111,7 +108,7 @@ def UMAP_plot( return fig -from scipy.cluster import hierarchy +from scipy.cluster import hierarchy # noqa: E402 nn = 5 target = "target" @@ -199,7 +196,7 @@ def retrieve_meta_test(): def Scenario1(path_simfile): - distance = pd.read_csv(file) + distance = pd.read_csv(path_simfile) distance.file_label = distance.file_label.str.lower() distance.search_term = distance.search_term.str.lower() distance = distance.drop_duplicates() @@ -278,7 +275,7 @@ def Scenario1(path_simfile): def Scenario2(path_simfile): - distance = pd.read_csv(file) + distance = pd.read_csv(path_simfile) distance.file_label = distance.file_label.str.lower() distance.search_term = distance.search_term.str.lower() distance = distance.drop_duplicates() diff --git a/geniml/cli.py b/geniml/cli.py index e63a6bda..1c85dead 100644 --- a/geniml/cli.py +++ b/geniml/cli.py @@ -553,7 +553,7 @@ def main(test_args=None): _LOGGER.info( "REGION COUNT | original: {}\tnew: {}\tchanged: {}\t\noutput file: {}".format( bedshifter.original_num_regions, - bedshifter.bed.shape[0], + len(bedshifter.bed), str(n), outfile_base, ) diff --git a/geniml/eval/gdst.py b/geniml/eval/gdst.py index dfcf01ea..f67ee4ce 100644 --- a/geniml/eval/gdst.py +++ b/geniml/eval/gdst.py @@ -28,7 +28,6 @@ def sample_from_vocab(vocab: List[str], num_samples: int, seed: int = 42) -> Lis """ chr_probs = {} region_dict = {} - num_vocab = len(vocab) # build stat from vocab for region in vocab: chr_str, position = region.split(":") @@ -231,7 +230,6 @@ def gdst_eval( mean_gds = [np.array(r).mean() for r in gds_res] std_gds = [np.array(r).std() for r in gds_res] - models = [t[0] for t in batch] for i in range(len(mean_gds)): print(f"{batch[i][0]}\n GDST score (std): {mean_gds[i]:.4f} ({std_gds[i]:.4f}) \n") gds_arr = [(batch[i][0], gds_res[i]) for i in range(len(batch))] diff --git a/geniml/eval/npt.py b/geniml/eval/npt.py index de726efa..75d97aeb 100644 --- a/geniml/eval/npt.py +++ b/geniml/eval/npt.py @@ -27,7 +27,6 @@ def get_topk_embed( tuple[np.ndarray, np.ndarray]: K indexes of nearest embeddings and the corresponding similarities. """ - num = len(embed) if dist == "cosine": nom = np.dot(embed[i : i + 1], embed.T) denom = np.linalg.norm(embed[i : i + 1]) * np.linalg.norm(embed, axis=1) @@ -366,9 +365,6 @@ def get_npt_score( count = count + 1 snprs = cal_snpr(avg_ratio, avg_ratio_ref) - ratio_msg = " ".join([f"{r:.6f}" for r in avg_ratio]) - ratio_ref_msg = " ".join([f"{r:.6f}" for r in avg_ratio_ref]) - snprs_msg = " ".join([f"{r:.6f}" for r in snprs]) result = { "K": K, "Avg_qNPR": avg_ratio, diff --git a/geniml/eval/rct.py b/geniml/eval/rct.py index 0f92061b..960cd77a 100644 --- a/geniml/eval/rct.py +++ b/geniml/eval/rct.py @@ -39,7 +39,6 @@ def get_rct_score( """ embed_rep, vocab = load_genomic_embeddings(path, embed_type) embed_bin, vocab_bin = load_genomic_embeddings(bin_path, "base") - region2idx = {r: i for i, r in enumerate(vocab)} region2idx_bin = {r: i for i, r in enumerate(vocab_bin)} # align embed_bin with embed_rep if out_dim <= 0: @@ -156,7 +155,6 @@ def rct_eval( assert res[0] == batch[i][0], "key == batch[i][0]" mean_rct = [np.array(r).mean() for r in rct_res] std_rct = [np.array(r).std() for r in rct_res] - models = [t[0] for t in batch] for i in range(len(mean_rct)): print(f"{batch[i][0]}\n RCT (std): {mean_rct[i]:.4f} ({std_rct[i]:.4f}) \n") rct_arr = [(batch[i][0], rct_res[i]) for i in range(len(batch))] diff --git a/geniml/io/io.py b/geniml/io/io.py index 9445dabd..b75af8b4 100644 --- a/geniml/io/io.py +++ b/geniml/io/io.py @@ -4,7 +4,7 @@ from typing_extensions import deprecated import os from hashlib import md5 -from typing import List, NoReturn, Union +from typing import Iterator, List, Union import numpy as np import pandas as pd @@ -50,7 +50,7 @@ def __init__(self, chr: str, start: int, stop: int): self.start = start self.end = stop - def __repr__(self): + def __repr__(self) -> str: return f"Region({self.chr}, {self.start}, {self.end})" @@ -204,16 +204,16 @@ def _read_file_pd(self, *args, **kwargs) -> pd.DataFrame: row_count += 1 raise BEDFileReadError("Cannot read bed file.") - def __len__(self): + def __len__(self) -> int: return self.length - def __getitem__(self, key): + def __getitem__(self, key) -> Region: if self.backed: raise NotImplementedError("Backed RegionSets do not currently support indexing.") else: return self.regions[key] - def __repr__(self): + def __repr__(self) -> str: if self.path: if self.backed: return f"RegionSet({self.path}, backed=True)" @@ -222,7 +222,7 @@ def __repr__(self): else: return f"RegionSet(n={self.length})" - def __iter__(self): + def __iter__(self) -> Iterator[Region]: if self.backed: # Open function depending on file type if self.is_gzipped: @@ -257,13 +257,14 @@ def __iter__(self): def identifier(self) -> str: return self.compute_bed_identifier() - def to_granges(self): + def to_granges(self) -> "genomicranges.GenomicRanges": # noqa: F821 """ Return GenomicRanges contained in this BED file. Returns: genomicranges.GenomicRanges: GenomicRanges object """ + import genomicranges seqnames, starts, ends = zip( *[(region.chr, region.start, region.end) for region in self.regions] @@ -362,21 +363,21 @@ def __init__( self._bedset_identifier = identifier - def __len__(self): + def __len__(self) -> int: return len(self.region_sets) - def __iter__(self): + def __iter__(self) -> Iterator[Union[RegionSet, GRegionSet]]: for region_set in self.region_sets: yield region_set - def __getitem__(self, indx: int): + def __getitem__(self, indx: int) -> Union[RegionSet, GRegionSet]: return self.region_sets[indx] @property def identifier(self) -> str: return self._bedset_identifier or self.compute_bedset_identifier() - def add(self, bedfile: RegionSet) -> NoReturn: + def add(self, bedfile: RegionSet) -> None: """ Add a BED file to the BED set. @@ -434,18 +435,18 @@ def __init__( self.strand = strand @property - def start(self): + def start(self) -> int: return self.start_position @property - def end(self): + def end(self) -> int: return self.end_position @property - def chr(self): + def chr(self) -> str: return self.chromosome - def to_region(self): + def to_region(self) -> Region: chr = self.chromosome start = int(self.start_position) end = int(self.end_position) @@ -455,10 +456,10 @@ def to_region(self): return Region(chr, start, end) - def __len__(self): + def __len__(self) -> int: return self.end - self.start - def __repr__(self): + def __repr__(self) -> str: return f"SNP({self.chromosome}, {self.start_position}, {self.end_position}, {self.strand})" @@ -467,7 +468,7 @@ class Maf: Python representation of a MAF file, only supports some columns for now """ - def _extract_value_from_col(self, col_name: str, line: str) -> any: + def _extract_value_from_col(self, col_name: str, line: str) -> Union[str, None]: """ Extract a value from a column in a line of a MAF file. @@ -476,7 +477,7 @@ def _extract_value_from_col(self, col_name: str, line: str) -> any: line (str): line from MAF file Returns: - any: value of column + Union[str, None]: value of column """ return line[self.col_positions[col_name]] if self.col_positions[col_name] else None @@ -556,16 +557,16 @@ def __init__( else: raise ValueError("mafs must be a path to a maf file") - def __len__(self): + def __len__(self) -> int: return self.length - def __getitem__(self, key): + def __getitem__(self, key) -> SNP: if self.backed: raise NotImplementedError("Backed MAFs do not currently support indexing.") else: return self.mafs[key] - def __iter__(self): + def __iter__(self) -> Iterator[SNP]: if self.backed: # Open function depending on file type open_func = gzip.open if is_gzipped(self.maf_file) else open @@ -594,7 +595,7 @@ def __iter__(self): for maf in self.mafs: yield maf - def __repr__(self): + def __repr__(self) -> str: return f"MAF({self.maf_file})" @@ -620,10 +621,10 @@ def __init__(self, region_sets: List[RegionSet] = None, file_globs: List[str] = for glob in file_globs: self.region_sets.extend([RegionSet(path) for path in glob.glob(glob)]) - def __getitem__(self, key): + def __getitem__(self, key) -> RegionSet: return self.region_sets[key] - def __len__(self): + def __len__(self) -> int: return len(self.region_sets) diff --git a/geniml/likelihood/__init__.py b/geniml/likelihood/__init__.py index 20c0a5bd..2fd9886f 100644 --- a/geniml/likelihood/__init__.py +++ b/geniml/likelihood/__init__.py @@ -1 +1 @@ -from .cli import build_subparser +from .cli import build_subparser # noqa: F401 diff --git a/geniml/models/__init__.py b/geniml/models/__init__.py index a7dc64ed..eeb292a2 100644 --- a/geniml/models/__init__.py +++ b/geniml/models/__init__.py @@ -1 +1 @@ -from .main import ExModel +from .main import ExModel # noqa: F401 diff --git a/geniml/nn/__init__.py b/geniml/nn/__init__.py index b23a4405..57091b16 100644 --- a/geniml/nn/__init__.py +++ b/geniml/nn/__init__.py @@ -1 +1 @@ -from .main import Attention, GradientReversal +from .main import Attention, GradientReversal # noqa: F401 diff --git a/geniml/region2vec/__init__.py b/geniml/region2vec/__init__.py index 22a7f78e..34b0b872 100644 --- a/geniml/region2vec/__init__.py +++ b/geniml/region2vec/__init__.py @@ -1,3 +1,12 @@ -# from .main import Region2Vec, Region2VecExModel -# from .main_legacy import region2vec -# +from .models import Region2Vec, RegionSet2Vec +from .main import Region2VecExModel +from .main_legacy import region2vec +from .utils import Region2VecDataset + +__all__ = [ + "Region2Vec", + "Region2VecExModel", + "RegionSet2Vec", + "Region2VecDataset", + "region2vec", +] diff --git a/geniml/region2vec/region2vec_train.py b/geniml/region2vec/region2vec_train.py index 35e89882..56436fee 100644 --- a/geniml/region2vec/region2vec_train.py +++ b/geniml/region2vec/region2vec_train.py @@ -11,7 +11,7 @@ from gensim.models.word2vec import LineSentence from . import utils -from .const import * +from .const import MAX_WAIT_TIME def find_dataset(data_folder: str) -> Union[str, int]: @@ -65,7 +65,7 @@ def main(args: argparse.Namespace) -> None: else: train_alg = 0 msg_model = "\033[94mUsing cbow, " - if args.hier_softmax == True or args.neg_samples == 0: + if args.hier_softmax or args.neg_samples == 0: hs = 1 msg_model += "hierarchical softmax\033[00m" else: diff --git a/geniml/region2vec/region_shuffling.py b/geniml/region2vec/region_shuffling.py index d38dace7..eb407660 100644 --- a/geniml/region2vec/region_shuffling.py +++ b/geniml/region2vec/region_shuffling.py @@ -47,7 +47,7 @@ def regions2sentences_sampling(self, src_path: str, dst_path: str) -> None: dst_path (str): The destination file that stores all the generated BED files; each line has regions sampled from a BED file. """ - with open(dst_fname, "w") as fout: + with open(dst_path, "w") as fout: for fname in self.filename_list: src_fname = os.path.join(src_path, fname) sentence = [] diff --git a/geniml/scembed/__init__.py b/geniml/scembed/__init__.py index 02ad2ece..bd62cb8f 100644 --- a/geniml/scembed/__init__.py +++ b/geniml/scembed/__init__.py @@ -1,4 +1,11 @@ -# from .annotation import * -# from .const import * -# from .main import * -# from .utils import * +from .main import ScEmbed +from .annotation import Annotator, AnnotationServer +from .exceptions import ScembedException, ModelNotTrainedError + +__all__ = [ + "ScEmbed", + "Annotator", + "AnnotationServer", + "ScembedException", + "ModelNotTrainedError", +] diff --git a/geniml/scembed/argparser.py b/geniml/scembed/argparser.py index 9eb315c7..3968b2dc 100644 --- a/geniml/scembed/argparser.py +++ b/geniml/scembed/argparser.py @@ -1,7 +1,7 @@ from ubiquerg import VersionInHelpParser from ._version import __version__ -from .const import * +from .const import MODULE_NAME def build_argparser(parser: VersionInHelpParser = None) -> VersionInHelpParser: @@ -18,7 +18,7 @@ def build_argparser(parser: VersionInHelpParser = None) -> VersionInHelpParser: ########################################################################### if parser is None: parser = VersionInHelpParser( - prog=PKG_NAME, + prog=MODULE_NAME, version=__version__, description="%(prog)s - embed single-cell data as region vectors", ) diff --git a/geniml/scembed/cli.py b/geniml/scembed/cli.py index 409615e4..f5eb4966 100644 --- a/geniml/scembed/cli.py +++ b/geniml/scembed/cli.py @@ -6,7 +6,6 @@ from ._version import __version__ from .argparser import build_argparser -from .const import * from .main import convert_anndata_to_documents, load_scanpy_data, shuffle_documents, train diff --git a/geniml/scembed/main.py b/geniml/scembed/main.py index 099b34be..d2b65bf6 100755 --- a/geniml/scembed/main.py +++ b/geniml/scembed/main.py @@ -27,7 +27,7 @@ load_local_region2vec_model, train_region2vec_model, ) -from ..tokenization.utils import tokenize_anndata +from ..tokenization.tokenize import tokenize_anndata from .const import MODULE_NAME _GENSIM_LOGGER = getLogger("gensim") @@ -75,7 +75,7 @@ def __init__( device if device else ("cuda" if torch.cuda.is_available() else "cpu") ) - def _init_tokenizer(self, tokenizer: Union[Tokenizer, str]): + def _init_tokenizer(self, tokenizer: Union[Tokenizer, str]) -> None: """ Initialize the tokenizer. @@ -92,7 +92,7 @@ def _init_tokenizer(self, tokenizer: Union[Tokenizer, str]): else: raise TypeError("tokenizer must be of type Tokenizer or str.") - def _init_model(self, tokenizer, **kwargs): + def _init_model(self, tokenizer, **kwargs) -> None: """ Initialize the core model. This will initialize the model from scratch. @@ -108,7 +108,7 @@ def _init_model(self, tokenizer, **kwargs): ) @property - def model(self): + def model(self) -> Region2Vec: """ Get the core Region2Vec model. @@ -117,7 +117,7 @@ def model(self): """ return self._model - def add_tokenizer(self, tokenizer: Tokenizer, **kwargs): + def add_tokenizer(self, tokenizer: Tokenizer, **kwargs) -> None: """ Add a tokenizer to the model. This should be use when the model is not initialized with a tokenizer. @@ -132,7 +132,7 @@ def add_tokenizer(self, tokenizer: Tokenizer, **kwargs): if not self.trained: self._init_model(**kwargs) - def _load_local_model(self, model_path: str, vocab_path: str, config_path: str): + def _load_local_model(self, model_path: str, vocab_path: str, config_path: str) -> None: """ Load the model from a checkpoint. @@ -157,7 +157,7 @@ def _init_from_huggingface( universe_file_name: str = UNIVERSE_FILE_NAME, config_file_name: str = CONFIG_FILE_NAME, **kwargs, - ): + ) -> None: """ Initialize the model from a huggingface model. @@ -275,7 +275,7 @@ def export( checkpoint_file: str = MODEL_FILE_NAME, universe_file: str = UNIVERSE_FILE_NAME, config_file: str = CONFIG_FILE_NAME, - ): + ) -> None: """ Function to facilitate exporting the model in a way that can be directly uploaded to huggingface. diff --git a/geniml/search/__init__.py b/geniml/search/__init__.py index ea603813..0f404e24 100644 --- a/geniml/search/__init__.py +++ b/geniml/search/__init__.py @@ -1,6 +1,6 @@ -from .backends import HNSWBackend, QdrantBackend -from .filebackend_tools import merge_backends -from .interfaces import BED2BEDSearchInterface, Text2BEDSearchInterface -from .query2vec import BED2Vec, Text2Vec -from .search_eval import anecdotal_search_from_hf_data -from .utils import rand_eval +from .backends import HNSWBackend, QdrantBackend # noqa: F401 +from .filebackend_tools import merge_backends # noqa: F401 +from .interfaces import BED2BEDSearchInterface, Text2BEDSearchInterface # noqa: F401 +from .query2vec import BED2Vec, Text2Vec # noqa: F401 +from .search_eval import anecdotal_search_from_hf_data # noqa: F401 +from .utils import rand_eval # noqa: F401 diff --git a/geniml/search/backends/__init__.py b/geniml/search/backends/__init__.py index 41343e65..285217f8 100644 --- a/geniml/search/backends/__init__.py +++ b/geniml/search/backends/__init__.py @@ -1,3 +1,3 @@ -from .bivecbackend import BiVectorBackend -from .dbbackend import QdrantBackend -from .filebackend import HNSWBackend +from .bivecbackend import BiVectorBackend # noqa: F401 +from .dbbackend import QdrantBackend # noqa: F401 +from .filebackend import HNSWBackend # noqa: F401 diff --git a/geniml/search/backends/dbbackend.py b/geniml/search/backends/dbbackend.py index 58883503..0c96f739 100644 --- a/geniml/search/backends/dbbackend.py +++ b/geniml/search/backends/dbbackend.py @@ -323,7 +323,7 @@ def retrieve_info( for id_ in ids: try: result = retrieval_dict[id_] - except: + except Exception: _LOGGER.warning(f"Warning: no id stored in backend matches {id_}.") continue result_dict = {"id": result.id, "payload": result.payload} diff --git a/geniml/search/hfdemo/__init__.py b/geniml/search/hfdemo/__init__.py index aa887171..b4a6612c 100644 --- a/geniml/search/hfdemo/__init__.py +++ b/geniml/search/hfdemo/__init__.py @@ -1 +1 @@ -from .bivec_demo import hf_bivec_search +from .bivec_demo import hf_bivec_search # noqa: F401 diff --git a/geniml/search/interfaces/__init__.py b/geniml/search/interfaces/__init__.py index ff2707ca..46bbb08b 100644 --- a/geniml/search/interfaces/__init__.py +++ b/geniml/search/interfaces/__init__.py @@ -1,3 +1,3 @@ -from .bed2bed import BED2BEDSearchInterface -from .mlfree import BiVectorSearchInterface -from .text2bed import Text2BEDSearchInterface +from .bed2bed import BED2BEDSearchInterface # noqa: F401 +from .mlfree import BiVectorSearchInterface # noqa: F401 +from .text2bed import Text2BEDSearchInterface # noqa: F401 diff --git a/geniml/search/query2vec/__init__.py b/geniml/search/query2vec/__init__.py index 87a3d029..8e3afd16 100644 --- a/geniml/search/query2vec/__init__.py +++ b/geniml/search/query2vec/__init__.py @@ -1,2 +1,2 @@ -from .bed2vec import BED2Vec -from .text2vec import Text2Vec +from .bed2vec import BED2Vec # noqa: F401 +from .text2vec import Text2Vec # noqa: F401 diff --git a/geniml/search/search_eval.py b/geniml/search/search_eval.py index d96d2e4a..aa91b87b 100644 --- a/geniml/search/search_eval.py +++ b/geniml/search/search_eval.py @@ -52,7 +52,7 @@ def anecdotal_search_from_hf_data( for file in metadata_dict[attribute][metadata]: try: search_results[result_files_id_dict[file]]["payload"][attribute] = metadata - except: + except Exception: continue return search_results diff --git a/geniml/text2bednn/__init__.py b/geniml/text2bednn/__init__.py index e7aca626..4cf292aa 100644 --- a/geniml/text2bednn/__init__.py +++ b/geniml/text2bednn/__init__.py @@ -1,2 +1,2 @@ -from .text2bednn import Vec2VecFNN -from .utils import arrays_to_torch_dataloader, metadata_dict_from_csv +from .text2bednn import Vec2VecFNN # noqa: F401 +from .utils import arrays_to_torch_dataloader, metadata_dict_from_csv # noqa: F401 diff --git a/geniml/text2bednn/text2bednn.py b/geniml/text2bednn/text2bednn.py index bc177b81..c71b00fd 100644 --- a/geniml/text2bednn/text2bednn.py +++ b/geniml/text2bednn/text2bednn.py @@ -446,7 +446,7 @@ def plot_training_hist( try: valid_loss = self.most_recent_train["val_loss"] plt.plot(epoch_range, valid_loss, "b", label="Validation loss") - except: + except Exception: pass plt.title(title) plt.legend() diff --git a/geniml/text2bednn/utils.py b/geniml/text2bednn/utils.py index 3f179487..f1bebc74 100644 --- a/geniml/text2bednn/utils.py +++ b/geniml/text2bednn/utils.py @@ -170,7 +170,7 @@ def metadata_dict_from_csv( } try: output_dict[row[series_key]].append(payload) - except: + except Exception: output_dict[row[series_key]] = [payload] series_count += 1 bed_count += 1 diff --git a/geniml/tokenization/__init__.py b/geniml/tokenization/__init__.py index 498a6d70..5d6225b0 100644 --- a/geniml/tokenization/__init__.py +++ b/geniml/tokenization/__init__.py @@ -1,2 +1,5 @@ -# from .main import Tokenizer, AnnDataTokenizer, TreeTokenizer -# from .main import hard_tokenization_main as hard_tokenization +from .main import hard_tokenization_main as hard_tokenization + +__all__ = [ + "hard_tokenization", +] diff --git a/geniml/tokenization/bedtools_tokenizer.py b/geniml/tokenization/bedtools_tokenizer.py index 68bfd7e2..51406d43 100644 --- a/geniml/tokenization/bedtools_tokenizer.py +++ b/geniml/tokenization/bedtools_tokenizer.py @@ -1,18 +1,27 @@ +import os +import shlex +import subprocess +import tempfile + +from geniml.io import RegionSet + from . import FileTokenizer class BEDToolsTokenizer(FileTokenizer): """A tokenizer that uses bedtools to tokenize BED files""" - def __init__(self, bedtools_path: str, universe_path: str = None): + def __init__(self, bedtools_path: str, universe_path: str = None, fraction: float = 0.5): """Initialize a BEDToolsTokenizer Args: bedtools_path (str): Path to a bedtools binary. universe_path (str): Path to a universe BED file. + fraction (float): Minimum overlap fraction. Defaults to 0.5. """ self.bedtools_path = bedtools_path self.universe_path = universe_path + self.fraction = fraction def tokenize(self, input_globs: list[str], universe_path: str = None) -> RegionSet: """Tokenize a RegionSet using bedtools""" @@ -20,13 +29,17 @@ def tokenize(self, input_globs: list[str], universe_path: str = None) -> RegionS universe_path = universe_path or self.universe_path # loop through globs and tokenize each file - for glob in input_globs: - for path in glob.glob(glob): - _tokenize_one(path, universe_path) + for glob_pattern in input_globs: + import glob as glob_module + + for path in glob_module.glob(glob_pattern): + self._tokenize_one(path, universe_path) def _tokenize_one(self, input_path: str, universe_path: str): output_path = os.path.join(input_path, "tokenized.bed") bedtools_path = self.bedtools_path + universe = universe_path + fraction = self.fraction # bedtools can't actually read from stdin, so we have to use a temporary file... # sort_process = subprocess.Popen(shlex.split(f"sort -k1,1V -k2,2n {input_path}"), stdout=subprocess.PIPE) diff --git a/geniml/tokenization/hard_tokenization_batch.py b/geniml/tokenization/hard_tokenization_batch.py index fd79e476..e9fc7645 100644 --- a/geniml/tokenization/hard_tokenization_batch.py +++ b/geniml/tokenization/hard_tokenization_batch.py @@ -3,6 +3,8 @@ import shlex import subprocess +from . import utils + def bedtools_tokenization( f: str, diff --git a/geniml/tokenization/main.py b/geniml/tokenization/main.py index e5f85b64..7fd7cb10 100644 --- a/geniml/tokenization/main.py +++ b/geniml/tokenization/main.py @@ -2,6 +2,7 @@ import os import shutil import subprocess +from argparse import Namespace from typing import List import numpy as np diff --git a/geniml/tokenization/tokenize.py b/geniml/tokenization/tokenize.py new file mode 100644 index 00000000..ebd0bc0e --- /dev/null +++ b/geniml/tokenization/tokenize.py @@ -0,0 +1,46 @@ +import numpy as np +import scanpy as sc +from tqdm import tqdm +from gtars.tokenizers import Tokenizer +from gtars.models import Region + + +def tokenize_anndata(adata: sc.AnnData, tokenizer: Tokenizer): + """ + Tokenize an AnnData object. This is more involved, so it gets its own function. + Args: + adata (sc.AnnData): The AnnData object to tokenize. + tokenizer (Tokenizer): The tokenizer to use. + """ + # extract regions from AnnData + # its weird because of how numpy handle Intervals, the parent class of Region, + # see here: + # https://stackoverflow.com/a/43722306/13175187 + adata_features = [ + Region(chr, int(start), int(end)) + for chr, start, end in tqdm( + zip(adata.var["chr"], adata.var["start"], adata.var["end"]), + total=adata.var.shape[0], + desc="Extracting regions from AnnData", + ) + ] + + features = np.ndarray(len(adata_features), dtype=object) + for i, region in enumerate(adata_features): + features[i] = region + + del adata_features + + # tokenize + tokenized = [] + x = adata.X + for row in tqdm( + range(adata.shape[0]), + total=adata.shape[0], + desc="Tokenizing", + ): + _, non_zeros = x[row].nonzero() + regions = features[non_zeros] + tokenized.append(tokenizer(regions)) + + return tokenized diff --git a/geniml/tokenization/utils.py b/geniml/tokenization/utils.py index 4d4c59aa..664c460d 100644 --- a/geniml/tokenization/utils.py +++ b/geniml/tokenization/utils.py @@ -1,11 +1,5 @@ import time -import numpy as np -import scanpy as sc -from tqdm import tqdm -from gtars.tokenizers import Tokenizer -from gtars.models import Region - class Timer: """Records the running time. @@ -48,44 +42,3 @@ def time_str(t: float) -> str: if t >= 60: return f"{t / 60:.2f}m" return f"{t:.2f}s" - - -def tokenize_anndata(adata: sc.AnnData, tokenizer: Tokenizer): - """ - Tokenize an AnnData object. This is more involved, so it gets its own function. - Args: - adata (sc.AnnData): The AnnData object to tokenize. - tokenizer (Tokenizer): The tokenizer to use. - """ - # extract regions from AnnData - # its weird because of how numpy handle Intervals, the parent class of Region, - # see here: - # https://stackoverflow.com/a/43722306/13175187 - adata_features = [ - Region(chr, int(start), int(end)) - for chr, start, end in tqdm( - zip(adata.var["chr"], adata.var["start"], adata.var["end"]), - total=adata.var.shape[0], - desc="Extracting regions from AnnData", - ) - ] - - features = np.ndarray(len(adata_features), dtype=object) - for i, region in enumerate(adata_features): - features[i] = region - - del adata_features - - # tokenize - tokenized = [] - x = adata.X - for row in tqdm( - range(adata.shape[0]), - total=adata.shape[0], - desc="Tokenizing", - ): - _, non_zeros = x[row].nonzero() - regions = features[non_zeros] - tokenized.append(tokenizer(regions)) - - return tokenized diff --git a/geniml/universe/__init__.py b/geniml/universe/__init__.py index 20c0a5bd..2fd9886f 100644 --- a/geniml/universe/__init__.py +++ b/geniml/universe/__init__.py @@ -1 +1 @@ -from .cli import build_subparser +from .cli import build_subparser # noqa: F401 diff --git a/geniml/universe/ccf_universe.py b/geniml/universe/ccf_universe.py index 48e59bd3..dfbef673 100644 --- a/geniml/universe/ccf_universe.py +++ b/geniml/universe/ccf_universe.py @@ -133,8 +133,8 @@ def get_uni(file, chrom, bedname): dist = np.absolute(uniq_val - cutoff) cutoff = uniq_val[dist.argmin()] pos = np.where(track_non_zero_sort == cutoff)[0] - f, l = pos[0] / len(track_non_zero_sort), pos[-1] / len(track_non_zero_sort) - q_cutoff = np.mean([f, l]) + f, last = pos[0] / len(track_non_zero_sort), pos[-1] / len(track_non_zero_sort) + q_cutoff = np.mean([f, last]) lower = np.quantile(track_non_zero_sort, max(0, q_cutoff - 0.2)) upper = np.quantile(track_non_zero_sort, min(1, q_cutoff + 0.2)) inter_pos = np.zeros(len(track), dtype=np.uint8) diff --git a/geniml/universe/custom_distribution.py b/geniml/universe/custom_distribution.py index 11b24ee0..29797252 100644 --- a/geniml/universe/custom_distribution.py +++ b/geniml/universe/custom_distribution.py @@ -86,9 +86,6 @@ def _init(self, X): super()._init(X) self.random_state = check_random_state(self.random_state) - mean_X = X.mean() - var_X = X.var() - if self._needs_init("p", "prob_"): # initialize with method of moments based on X raise NotImplementedError @@ -271,7 +268,7 @@ def _compute_likelihood(self, X): ).T def _initialize_sufficient_statistics(self): - stats = super()._initialize_sufficient_statistics() + super()._initialize_sufficient_statistics() raise NotImplementedError def _accumulate_sufficient_statistics( diff --git a/requirements/requirements-ml.txt b/requirements/requirements-ml.txt index 92c5d9c9..9a7a31e1 100644 --- a/requirements/requirements-ml.txt +++ b/requirements/requirements-ml.txt @@ -1,14 +1,8 @@ -anndata > 0.9.0 -fastembed >= 0.2.5 gensim >= 4.3.3 huggingface_hub >= 0.25.1 -qdrant_client >= 1.16.1 # hnswlib >= 0.8.0 # Not supported after Python 3.11 - needs to be installed manually if required -paramiko >= 3.0.0 pyBigWig >= 0.3.23 -scanpy >= 1.10.3 torch >= 2.3.0 -langchain-huggingface==0.0.2 hmmlearn >=0.3.2 scipy >= 1.13.1 transformers >= 4.52.4 \ No newline at end of file diff --git a/requirements/requirements-sc.txt b/requirements/requirements-sc.txt new file mode 100644 index 00000000..6cf91bc2 --- /dev/null +++ b/requirements/requirements-sc.txt @@ -0,0 +1,3 @@ +scanpy >= 1.10.3 +anndata > 0.9.0 +loompy diff --git a/requirements/requirements-search.txt b/requirements/requirements-search.txt new file mode 100644 index 00000000..1c729459 --- /dev/null +++ b/requirements/requirements-search.txt @@ -0,0 +1,2 @@ +qdrant_client >= 1.16.1 +fastembed >= 0.2.5 diff --git a/setup.py b/setup.py index 4f45c5bc..ce7f7d66 100755 --- a/setup.py +++ b/setup.py @@ -18,25 +18,25 @@ with open(PACKAGE_NAME + "/_version.py", "r") as versionfile: version = versionfile.readline().split()[-1].strip("\"'\n") + # Optional dependencies # Extras requires a dictionary and not a list? -with open("requirements/requirements-ml.txt", "r") as reqs_file: - ml_dep = [] - for line in reqs_file: - if not line.strip(): - continue - ml_dep.append(line.strip()) +def _read_reqs(path): + with open(path, "r") as fh: + return [line.strip() for line in fh if line.strip() and not line.strip().startswith("#")] -with open("requirements/requirements-test.txt", "r") as reqs_file: - test_dep = [] - for line in reqs_file: - if not line.strip(): - continue - test_dep.append(line.strip()) + +ml_dep = _read_reqs("requirements/requirements-ml.txt") +sc_dep = _read_reqs("requirements/requirements-sc.txt") +search_dep = _read_reqs("requirements/requirements-search.txt") +test_dep = _read_reqs("requirements/requirements-test.txt") extra["install_requires"] = DEPENDENCIES extra["extras_require"] = { "ml": ml_dep, + "sc": sc_dep, + "search": search_dep, + "all": ml_dep + sc_dep + search_dep, "test": test_dep, } diff --git a/tests/test_bedshift.py b/tests/test_bedshift.py index 7d3d8327..7d74df81 100644 --- a/tests/test_bedshift.py +++ b/tests/test_bedshift.py @@ -18,8 +18,8 @@ def bs(): class TestBedshift: def test_read_bed(self): reader = bedshift.Bedshift(os.path.join(DATA_FOLDER_PATH, "header_test.bed")) - assert list(reader.bed.columns) == [0, 1, 2, 3] - assert list(reader.bed.index) == [0, 1, 2] + assert len(reader.bed) == 3 # 3 rows + assert len(reader.bed[0]) == 4 # 4 fields per row: chrom, start, end, mod_flag def test_read_chrom_sizes(self, bs): bs._read_chromsizes(os.path.join(DATA_FOLDER_PATH, "hg19.chrom.sizes")) @@ -47,7 +47,6 @@ def test_add_valid_regions(self, bs): 2000, 1000, valid_bed=os.path.join(DATA_FOLDER_PATH, "small_test.bed"), - delimiter="\t", ) assert added == 500 # bs.to_bed(os.path.join(SCRIPT_PATH, "add_valid_test.bed")) diff --git a/tests/test_bedshift_invalid_regions.py b/tests/test_bedshift_invalid_regions.py new file mode 100644 index 00000000..95e6c7a6 --- /dev/null +++ b/tests/test_bedshift_invalid_regions.py @@ -0,0 +1,235 @@ +"""Tests for bedshift invalid region bug (GitHub issue #49). + +Validates that bedshift never produces regions where start >= end, +even after many rounds of perturbation. +""" + +import os +import random +import tempfile + +import numpy as np +import pytest + +from geniml.bedshift import bedshift + +DATA_FOLDER_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "data", "bedshift") + + +def _count_invalid_regions(bed): + """Count regions where start >= end.""" + return sum(1 for row in bed if row[1] >= row[2]) + + +def _make_bedfile(regions): + """Write regions to a temp BED file and return the path.""" + f = tempfile.NamedTemporaryFile(mode="w", suffix=".bed", delete=False) + for chrom, start, end in regions: + f.write(f"{chrom}\t{start}\t{end}\n") + f.close() + return f.name + + +@pytest.fixture +def bs(): + return bedshift.Bedshift( + os.path.join(DATA_FOLDER_PATH, "test.bed"), + chrom_sizes=os.path.join(DATA_FOLDER_PATH, "hg38.chrom.sizes"), + ) + + +class TestCutTinyRegions: + """Test that _cut handles regions too small to cut.""" + + def _make_bs_with_regions(self, regions): + """Create a Bedshift object with specific regions.""" + path = _make_bedfile(regions) + bs = bedshift.Bedshift( + path, + chrom_sizes=os.path.join(DATA_FOLDER_PATH, "hg38.chrom.sizes"), + ) + os.unlink(path) + return bs + + def test_cut_1bp_region(self): + """A 1bp region (end - start = 1) cannot be cut; should be skipped.""" + bs = self._make_bs_with_regions([("chr1", 100, 101)]) + # _cut should return (None, None) for uncuttable region + result = bs._cut(0) + drop_row, new_regions = result + # Either skipped (None, None) or valid regions + if drop_row is not None: + for region in new_regions: + assert region[1] < region[2], f"Invalid cut result: {region}" + + def test_cut_2bp_region(self): + """A 2bp region can be cut into two 1bp regions.""" + bs = self._make_bs_with_regions([("chr1", 100, 102)]) + drop_row, new_regions = bs._cut(0) + if drop_row is not None: + for region in new_regions: + assert region[1] < region[2], f"Invalid cut result: {region}" + + def test_cut_5bp_region(self): + bs = self._make_bs_with_regions([("chr1", 100, 105)]) + drop_row, new_regions = bs._cut(0) + if drop_row is not None: + for region in new_regions: + assert region[1] < region[2], f"Invalid cut result: {region}" + + def test_cut_10bp_region(self): + bs = self._make_bs_with_regions([("chr1", 100, 110)]) + drop_row, new_regions = bs._cut(0) + if drop_row is not None: + for region in new_regions: + assert region[1] < region[2], f"Invalid cut result: {region}" + + def test_cut_15bp_region(self): + bs = self._make_bs_with_regions([("chr1", 100, 115)]) + drop_row, new_regions = bs._cut(0) + if drop_row is not None: + for region in new_regions: + assert region[1] < region[2], f"Invalid cut result: {region}" + + def test_cut_19bp_region(self): + """19bp region triggers the old fallback bug.""" + bs = self._make_bs_with_regions([("chr1", 100, 119)]) + drop_row, new_regions = bs._cut(0) + if drop_row is not None: + for region in new_regions: + assert region[1] < region[2], f"Invalid cut result: {region}" + + def test_cut_20bp_region(self): + bs = self._make_bs_with_regions([("chr1", 100, 120)]) + drop_row, new_regions = bs._cut(0) + if drop_row is not None: + for region in new_regions: + assert region[1] < region[2], f"Invalid cut result: {region}" + + def test_cut_many_tiny_regions(self): + """Cut a batch of tiny regions; all results must be valid.""" + regions = [ + ("chr1", 1000 + i * 100, 1000 + i * 100 + size) + for i, size in enumerate([1, 2, 3, 5, 8, 10, 15, 19, 20]) + ] + bs = self._make_bs_with_regions(regions) + bs.cut(1.0) # cut all + assert _count_invalid_regions(bs.bed) == 0, ( + f"Found {_count_invalid_regions(bs.bed)} invalid regions after cutting tiny regions" + ) + + +class TestMergeUnsorted: + """Test that merge works correctly even when data is not positionally sorted.""" + + def _make_bs_with_regions(self, regions): + path = _make_bedfile(regions) + bs = bedshift.Bedshift( + path, + chrom_sizes=os.path.join(DATA_FOLDER_PATH, "hg38.chrom.sizes"), + ) + os.unlink(path) + return bs + + def test_merge_reversed_order(self): + """Two same-chrom regions where the second has smaller coordinates.""" + # After sorting by the constructor, these will be ordered. + # But we can unsort them manually to simulate post-perturbation state. + regions = [("chr1", 1000, 2000), ("chr1", 500, 900)] + bs = self._make_bs_with_regions(regions) + # Manually unsort to simulate post-perturbation state + bs.bed = [ + ["chr1", 5000, 6000, "-"], + ["chr1", 500, 900, "-"], + ] + drop_rows, merged = bs._merge(0) + if drop_rows is not None: + assert merged[1] < merged[2], f"Invalid merged region: {merged}" + + def test_merge_via_public_method(self): + """The public merge method should produce only valid regions.""" + regions = [ + ("chr1", 1000, 2000), + ("chr1", 5000, 6000), + ("chr1", 500, 900), + ("chr1", 3000, 4000), + ] + bs = self._make_bs_with_regions(regions) + # Scramble internal order + random.shuffle(bs.bed) + bs.merge(0.5) + assert _count_invalid_regions(bs.bed) == 0 + + +class TestAddNegativeLength: + """Test that add never produces regions with end <= start.""" + + def test_add_with_high_stdev(self, bs): + """High stdev relative to mean can produce negative lengths.""" + random.seed(42) + np.random.seed(42) + bs.add(0.5, addmean=10, addstdev=100) + invalid = _count_invalid_regions(bs.bed) + assert invalid == 0, f"Found {invalid} invalid regions after add with high stdev" + + def test_add_with_zero_mean(self, bs): + """Zero mean with any stdev will frequently produce negative lengths.""" + random.seed(123) + np.random.seed(123) + bs.add(0.5, addmean=0, addstdev=50) + invalid = _count_invalid_regions(bs.bed) + assert invalid == 0, f"Found {invalid} invalid regions after add with zero mean" + + +class TestStressMultiRound: + """Stress test: many rounds of all_perturbations should never produce invalid regions.""" + + def test_100_rounds(self, bs): + """Run 100 rounds of perturbations on 1000 regions. + + This is the primary reproducer for GitHub issue #49. + """ + random.seed(42) + np.random.seed(42) + for round_num in range(100): + bs.all_perturbations( + addrate=0.1, + addmean=320.0, + addstdev=30.0, + shiftrate=0.1, + shiftmean=0.0, + shiftstdev=150.0, + cutrate=0.1, + mergerate=0.05, + droprate=0.1, + ) + invalid = _count_invalid_regions(bs.bed) + assert invalid == 0, ( + f"Round {round_num}: found {invalid} invalid regions out of {len(bs.bed)} total" + ) + + +class TestValidation: + """Test the _validate_region utility and to_bed safety net.""" + + def test_validate_region_valid(self, bs): + assert bs._validate_region(0, 100) is True + assert bs._validate_region(50, 51) is True + + def test_validate_region_invalid(self, bs): + assert bs._validate_region(100, 50) is False # start > end + assert bs._validate_region(100, 100) is False # start == end + assert bs._validate_region(-1, 100) is False # negative start + + def test_to_bed_filters_invalid(self, bs, tmp_path): + """to_bed should not write invalid regions even if they exist internally.""" + # Inject an invalid region + bs.bed.append(["chr1", 5000, 3000, "X"]) + outfile = os.path.join(tmp_path, "out.bed") + bs.to_bed(outfile) + # Read back and verify no invalid regions + with open(outfile) as f: + for line in f: + parts = line.strip().split("\t") + start, end = int(parts[1]), int(parts[2]) + assert start < end, f"Invalid region in output: {line.strip()}"