iona-denoise-50m / processing_iona.py
chrisagrams's picture
Rename MSDelta to Iona
d123296 verified
Raw History Blame Contribute Delete
12.8 kB
"""Preprocessing and collation for Iona mass spectra."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import torch
from transformers import BatchFeature, FeatureExtractionMixin
def _as_spectrum_batch(values: Any, name: str) -> tuple[list[torch.Tensor], bool]:
"""Normalize one spectrum or a batch of spectra into a tensor list."""
if isinstance(values, torch.Tensor):
if values.ndim == 1:
return [values], True
if values.ndim == 2:
return [row for row in values], False
raise ValueError(f"{name} must be one- or two-dimensional")
if not isinstance(values, (list, tuple)):
values = list(values)
if not values:
return [torch.empty(0)], True
first = values[0]
if isinstance(first, (list, tuple, torch.Tensor)) or hasattr(first, "ndim"):
return [torch.as_tensor(row) for row in values], False
return [torch.as_tensor(values)], True
class IonaProcessor(FeatureExtractionMixin):
"""Convert raw centroided spectra into padded Iona model inputs."""
model_input_names = ["mz", "log_intensity", "attention_mask"]
def __init__(
self,
max_peaks: int = 512,
padding_value: float = 0.0,
**kwargs,
):
if max_peaks <= 0:
raise ValueError("max_peaks must be positive")
super().__init__(
max_peaks=max_peaks,
padding_value=padding_value,
**kwargs,
)
self.max_peaks = max_peaks
self.padding_value = padding_value
def _process_one(
self, mz: torch.Tensor, intensity: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
mz = torch.as_tensor(mz, dtype=torch.float32)
intensity = torch.as_tensor(intensity, dtype=torch.float32)
if mz.ndim != 1 or intensity.ndim != 1:
raise ValueError("each mz and intensity spectrum must be one-dimensional")
if mz.shape != intensity.shape:
raise ValueError("each mz and intensity spectrum must have equal lengths")
if mz.numel() == 0:
raise ValueError("spectra must contain at least one peak")
if mz.numel() > self.max_peaks:
raise ValueError(f"spectra must contain at most {self.max_peaks} peaks")
if not torch.isfinite(mz).all() or not torch.isfinite(intensity).all():
raise ValueError("mz and intensity values must be finite")
if (intensity < 0).any():
raise ValueError("intensity values must be nonnegative")
base_peak = intensity.max()
if base_peak <= 0:
raise ValueError("spectra must contain at least one positive intensity")
selected = torch.arange(mz.numel(), device=mz.device)
log_intensity = torch.log1p(intensity)
log_intensity = log_intensity / log_intensity.max().clamp_min(1e-8)
labels = intensity / intensity.sum().clamp_min(1e-12)
return mz.contiguous(), log_intensity.contiguous(), labels.contiguous(), selected
def process_denoising_example(self, mz, intensity, noise) -> dict[str, list | int]:
"""Process one spectrum while preserving peak/noise-label alignment."""
noise = torch.as_tensor(noise, dtype=torch.bool)
mz_tensor = torch.as_tensor(mz, dtype=torch.float32)
intensity_tensor = torch.as_tensor(intensity, dtype=torch.float32)
if noise.ndim != 1 or noise.shape != mz_tensor.shape:
raise ValueError("mz, intensity, and noise must have equal one-dimensional shapes")
mass, log_intensity, _, selected = self._process_one(mz_tensor, intensity_tensor)
return {
"mz": mass.tolist(),
"log_intensity": log_intensity.tolist(),
"labels": noise[selected].float().tolist(),
}
def process_retrieval_example(self, consensus, experimental) -> dict[str, list]:
"""Process one consensus spectrum and its three experimental replicates."""
if len(experimental) != 3:
raise ValueError("retrieval examples must contain exactly three experimental spectra")
spectra = [consensus, *experimental]
processed = [
self(spectrum["mz"], spectrum["intensity"], padding=False) for spectrum in spectra
]
return {
"mz": [values["mz"] for values in processed],
"log_intensity": [values["log_intensity"] for values in processed],
}
def pad(
self,
encoded_inputs: list[dict[str, Any]],
*,
padding: bool | str = True,
max_length: int | None = None,
return_tensors: str | None = None,
**kwargs,
) -> BatchFeature:
"""Pad processed peak-classification examples to a common length."""
lengths = [len(example["mz"]) for example in encoded_inputs]
if padding == "max_length":
if max_length is None:
raise ValueError("max_length is required with padding='max_length'")
target_length = max_length
else:
target_length = max(lengths)
data: dict[str, list] = {
"mz": [],
"log_intensity": [],
"attention_mask": [],
"labels": [],
}
for example, length in zip(encoded_inputs, lengths):
pad = target_length - length
data["mz"].append(example["mz"] + [self.padding_value] * pad)
data["log_intensity"].append(example["log_intensity"] + [self.padding_value] * pad)
data["attention_mask"].append([1] * length + [0] * pad)
data["labels"].append(example["labels"] + [-100.0] * pad)
return BatchFeature(data=data, tensor_type=return_tensors)
def __call__(
self,
mz,
intensity,
*,
padding: bool | str = True,
truncation: bool = True,
max_length: int | None = None,
pad_to_multiple_of: int | None = None,
return_tensors: str | None = None,
return_labels: bool = False,
) -> BatchFeature:
"""Process raw m/z and intensity arrays into model-ready features."""
mz_batch, mz_was_single = _as_spectrum_batch(mz, "mz")
intensity_batch, intensity_was_single = _as_spectrum_batch(intensity, "intensity")
if len(mz_batch) != len(intensity_batch) or mz_was_single != intensity_was_single:
raise ValueError("mz and intensity must describe the same number of spectra")
processed = [self._process_one(m, i)[:3] for m, i in zip(mz_batch, intensity_batch)]
limit = self.max_peaks if max_length is None else max_length
if limit <= 0:
raise ValueError("max_length must be positive")
if truncation:
processed = [(m[:limit], li[:limit], y[:limit]) for m, li, y in processed]
elif any(m.numel() > limit for m, _, _ in processed) and padding == "max_length":
raise ValueError("a spectrum exceeds max_length while truncation is disabled")
lengths = [m.numel() for m, _, _ in processed]
target_length: int | None
if padding == "max_length":
target_length = limit
elif padding:
target_length = max(lengths)
else:
target_length = None
if target_length is not None and pad_to_multiple_of:
target_length = (
(target_length + pad_to_multiple_of - 1) // pad_to_multiple_of
) * pad_to_multiple_of
data: dict[str, list] = {
"mz": [],
"log_intensity": [],
"attention_mask": [],
}
if return_labels:
data["labels"] = []
for mass, log_int, labels in processed:
length = mass.numel()
padded_length = length if target_length is None else target_length
pad = padded_length - length
if pad < 0:
raise ValueError("a spectrum exceeds the requested padded length")
data["mz"].append(torch.cat([mass, mass.new_full((pad,), self.padding_value)]).tolist())
data["log_intensity"].append(
torch.cat([log_int, log_int.new_full((pad,), self.padding_value)]).tolist()
)
data["attention_mask"].append([1] * length + [0] * pad)
if return_labels:
data["labels"].append(torch.cat([labels, labels.new_zeros(pad)]).tolist())
if mz_was_single and not padding and return_tensors is None:
data = {name: values[0] for name, values in data.items()}
return BatchFeature(data=data, tensor_type=return_tensors)
@dataclass
class IonaDataCollatorForPreTraining:
"""Pad processed spectra and sample masked peaks for pretraining."""
mask_ratio: float = 0.15
min_masked: int = 1
pad_to_multiple_of: int | None = None
def __post_init__(self) -> None:
if not 0.0 <= self.mask_ratio <= 1.0:
raise ValueError("mask_ratio must be in [0, 1]")
if self.min_masked < 0:
raise ValueError("min_masked must be nonnegative")
def __call__(self, features: list[dict]) -> dict[str, torch.Tensor]:
if not features:
raise ValueError("features must not be empty")
lengths = [len(feature["mz"]) for feature in features]
max_length = max(lengths)
if self.pad_to_multiple_of:
max_length = (
(max_length + self.pad_to_multiple_of - 1) // self.pad_to_multiple_of
) * self.pad_to_multiple_of
batch_size = len(features)
mz = torch.zeros(batch_size, max_length, dtype=torch.float32)
log_intensity = torch.zeros_like(mz)
labels = torch.zeros_like(mz)
attention_mask = torch.zeros(batch_size, max_length, dtype=torch.long)
mask_positions = torch.zeros(batch_size, max_length, dtype=torch.bool)
for row, (feature, length) in enumerate(zip(features, lengths)):
if length == 0:
continue
mz[row, :length] = torch.as_tensor(feature["mz"], dtype=torch.float32)
log_intensity[row, :length] = torch.as_tensor(
feature["log_intensity"], dtype=torch.float32
)
target = feature.get("labels", feature.get("intensity_prob"))
if target is None:
raise ValueError("pretraining features must include labels")
labels[row, :length] = torch.as_tensor(target, dtype=torch.float32)
attention_mask[row, :length] = 1
n_masked = min(
length,
max(self.min_masked, int(round(length * self.mask_ratio))),
)
if n_masked:
mask_positions[row, torch.randperm(length)[:n_masked]] = True
return {
"mz": mz,
"log_intensity": log_intensity,
"attention_mask": attention_mask,
"mask_positions": mask_positions,
"labels": labels,
}
@dataclass
class IonaDataCollatorForRetrieval:
"""Flatten and pad four-spectrum analyte groups for contrastive training."""
def __call__(self, features: list[dict]) -> dict[str, torch.Tensor]:
if not features:
raise ValueError("features must not be empty")
mzs: list[list[float]] = []
log_intensities: list[list[float]] = []
group_ids: list[int] = []
for group_id, feature in enumerate(features):
group_mz = feature["mz"]
group_intensity = feature["log_intensity"]
if len(group_mz) != 4 or len(group_intensity) != 4:
raise ValueError("each retrieval example must contain four spectra")
mzs.extend(group_mz)
log_intensities.extend(group_intensity)
group_ids.extend([group_id] * 4)
target_length = max(max((len(mz) for mz in mzs), default=0), 1)
batch_size = len(mzs)
mz = torch.zeros(batch_size, target_length, dtype=torch.float32)
log_intensity = torch.zeros_like(mz)
attention_mask = torch.zeros(batch_size, target_length, dtype=torch.long)
for index, (mass, intensity) in enumerate(zip(mzs, log_intensities)):
length = len(mass)
if length == 0:
continue
mz[index, :length] = torch.as_tensor(mass, dtype=torch.float32)
log_intensity[index, :length] = torch.as_tensor(intensity, dtype=torch.float32)
attention_mask[index, :length] = 1
return {
"mz": mz,
"log_intensity": log_intensity,
"attention_mask": attention_mask,
"group_ids": torch.tensor(group_ids, dtype=torch.long),
}
IonaProcessor.register_for_auto_class("AutoProcessor")