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