Skip to content

Commit 6077584

Browse files
committed
more efficient metadata fetching and removed unused dynamic size fields from vector store
1 parent c2679c0 commit 6077584

4 files changed

Lines changed: 23 additions & 13 deletions

File tree

TreeMS2/similarity_matrix/filters/precursor_mz_filter.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,10 +25,12 @@ def __init__(self, groups: Groups, vector_store: VectorStore,
2525
def construct_mask(self, similarity_matrix: SimilarityMatrix) -> SpectraMatrix:
2626
rows, cols = similarity_matrix.matrix.nonzero()
2727

28-
precursor_mz_rows = self.vector_store.get_data(rows=rows, columns=["precursor_mz"])["precursor_mz"].to_numpy(
29-
dtype=np.float32)
30-
precursor_mz_cols = self.vector_store.get_data(rows=cols, columns=["precursor_mz"])["precursor_mz"].to_numpy(
31-
dtype=np.float32)
28+
# precursor mz values for all entries
29+
precursor_mz = self.vector_store.get_col("precursor_mz").to_numpy(dtype=np.float32).ravel()
30+
31+
# Vectorized lookup using rows and cols
32+
precursor_mz_rows = precursor_mz[rows]
33+
precursor_mz_cols = precursor_mz[cols]
3234

3335
mask_data = np.abs(precursor_mz_rows - precursor_mz_cols) > self.precursor_mz_window
3436
# Filter the rows and columns based on the mask

TreeMS2/similarity_sets.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -25,8 +25,10 @@ def update_similarity_sets(self, similarity_matrix: SimilarityMatrix):
2525
rows, cols = similarity_matrix.matrix.nonzero()
2626
total_spectra = self.groups.total_spectra
2727

28-
row_ids = self.vector_store.get_data(rows, ["global_id"])["global_id"].to_numpy(dtype=np.int32)
29-
col_ids = self.vector_store.get_data(cols, ["global_id"])["global_id"].to_numpy(dtype=np.int32)
28+
global_ids = self.vector_store.get_col("global_id").to_numpy(dtype=np.uint32).ravel()
29+
row_ids = global_ids[rows]
30+
col_ids = global_ids[cols]
31+
3032
m = csr_matrix((similarity_matrix.matrix.data, (row_ids, col_ids)), shape=(total_spectra, total_spectra),
3133
dtype=np.bool_)
3234

TreeMS2/spectrum/group_spectrum.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -34,11 +34,11 @@ def to_dict(self) -> Dict:
3434
"spectrum_id": self._spectrum_id,
3535
"file_id": self._file_id,
3636
"group_id": self._group_id,
37-
"identifier": self.spectrum.identifier,
37+
#self "identifier": self.spectrum.identifier,
3838
"precursor_mz": self.spectrum.precursor_mz,
3939
"precursor_charge": self.spectrum.precursor_charge,
40-
"mz": self.spectrum.mz,
41-
"intensity": self.spectrum.intensity,
40+
# "mz": self.spectrum.mz,
41+
# "intensity": self.spectrum.intensity,
4242
"retention_time": self.spectrum.retention_time,
4343
"vector": self.vector,
4444
}

TreeMS2/vector_store/vector_store.py

Lines changed: 10 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -37,13 +37,13 @@ def __init__(self, name: str, directory: str, vector_dim: int):
3737
pa.field("spectrum_id", pa.uint16()),
3838
pa.field("file_id", pa.uint16()),
3939
pa.field("group_id", pa.uint16()),
40-
pa.field("identifier", pa.string()),
40+
# pa.field("identifier", pa.string()),
4141
pa.field("precursor_mz", pa.float32()),
4242
pa.field("precursor_charge", pa.int8()),
43-
pa.field("mz", pa.list_(pa.float32())),
44-
pa.field("intensity", pa.list_(pa.float32())),
43+
# pa.field("mz", pa.list_(pa.float32())),
44+
# pa.field("intensity", pa.list_(pa.float32())),
4545
pa.field("retention_time", pa.float32()),
46-
pa.field("vector", pa.list_(pa.float32())),
46+
pa.field("vector", pa.list_(pa.float32(), vector_dim)) ,
4747
])
4848

4949
def _get_dataset(self) -> Optional[LanceDataset]:
@@ -115,6 +115,12 @@ def get_data(self, rows: List[int], columns: List[str]) -> pd.DataFrame:
115115
return pd.DataFrame(columns=columns)
116116
return ds.take(indices=rows, columns=columns).to_pandas()
117117

118+
def get_col(self, column) -> pd.DataFrame:
119+
ds = self._get_dataset()
120+
if ds is None:
121+
return pd.DataFrame(columns=column)
122+
return ds.to_table(columns=[column]).to_pandas()
123+
118124
def add_global_ids(self, groups: Groups) -> None:
119125
def compute_global_id(row: pd.Series) -> int:
120126
offset = groups.get_group(row['group_id']).get_peak_file(row['file_id']).begin

0 commit comments

Comments
 (0)