| """带 label 的 RNA–RBP 数据集,复用 ProRiboGen generation 模块的 RnaRbpDataset 与 collate。""" |
| from __future__ import annotations |
|
|
| import random |
| from functools import partial |
| from typing import Iterator |
|
|
| import torch |
| from torch.utils.data import Dataset, Sampler |
| from transformers import AutoTokenizer |
|
|
| import sys |
| from pathlib import Path |
|
|
| |
| _GENERATOR_ROOT = Path(__file__).resolve().parent.parent / "generator" |
| if str(_GENERATOR_ROOT) not in sys.path: |
| sys.path.insert(0, str(_GENERATOR_ROOT)) |
|
|
| from src.utils import RnaRbpDataset, rna_rbp_collate_fn |
|
|
|
|
| class LabeledRnaRbpDataset(RnaRbpDataset): |
| """在父类基础上读取 ``label`` 列(0/1)。""" |
|
|
| def __getitem__(self, idx: int) -> dict: |
| sample = super().__getitem__(idx) |
| sample["label"] = int(self.data.iloc[idx]["label"]) |
| return sample |
|
|
|
|
| def build_pair_indices(dataset: LabeledRnaRbpDataset) -> list[tuple[int, int]]: |
| """从 labeled 表构造 (pos_idx, neg_idx) 列表,优先用 pair_id。""" |
| df = dataset.data.reset_index(drop=True) |
| if "pair_id" in df.columns: |
| pos_by: dict[str, list[int]] = {} |
| neg_by: dict[str, list[int]] = {} |
| for i in range(len(df)): |
| pid = str(df.iloc[i]["pair_id"]) |
| if int(df.iloc[i]["label"]) == 1: |
| pos_by.setdefault(pid, []).append(i) |
| else: |
| neg_by.setdefault(pid, []).append(i) |
| pairs: list[tuple[int, int]] = [] |
| for k in sorted(pos_by): |
| if k not in neg_by: |
| continue |
| ps, ns = pos_by[k], neg_by[k] |
| n = min(len(ps), len(ns)) |
| pairs.extend(zip(ps[:n], ns[:n])) |
| if pairs: |
| return pairs |
| pos_idx = [i for i in range(len(df)) if int(df.iloc[i]["label"]) == 1] |
| neg_idx = [i for i in range(len(df)) if int(df.iloc[i]["label"]) == 0] |
| if len(pos_idx) != len(neg_idx): |
| raise ValueError(f"pos={len(pos_idx)} neg={len(neg_idx)},无法配对") |
| return list(zip(pos_idx, neg_idx)) |
|
|
|
|
| def split_pair_indices( |
| pairs: list[tuple[int, int]], |
| val_fraction: float, |
| seed: int, |
| ) -> tuple[list[tuple[int, int]], list[tuple[int, int]]]: |
| n_val = max(1, int(len(pairs) * val_fraction)) |
| rng = random.Random(seed) |
| order = list(range(len(pairs))) |
| rng.shuffle(order) |
| val_set = set(order[:n_val]) |
| train_pairs = [pairs[i] for i in range(len(pairs)) if i not in val_set] |
| val_pairs = [pairs[i] for i in val_set] |
| return train_pairs, val_pairs |
|
|
|
|
| class PairedBatchSampler(Sampler[list[int]]): |
| """每个 batch 含若干 (pos, neg) 对,保证 ranking loss 可配对。""" |
|
|
| def __init__( |
| self, |
| pairs: list[tuple[int, int]], |
| *, |
| pairs_per_batch: int, |
| num_replicas: int = 1, |
| rank: int = 0, |
| shuffle: bool = True, |
| seed: int = 42, |
| ) -> None: |
| if pairs_per_batch < 1: |
| raise ValueError("pairs_per_batch 须 >= 1") |
| self.pairs = pairs |
| self.pairs_per_batch = pairs_per_batch |
| self.num_replicas = num_replicas |
| self.rank = rank |
| self.shuffle = shuffle |
| self.seed = seed |
| self.epoch = 0 |
|
|
| def set_epoch(self, epoch: int) -> None: |
| self.epoch = epoch |
|
|
| def __iter__(self) -> Iterator[list[int]]: |
| order = list(range(len(self.pairs))) |
| if self.shuffle: |
| rng = random.Random(self.seed + self.epoch) |
| rng.shuffle(order) |
| order = order[self.rank :: self.num_replicas] |
| batch: list[int] = [] |
| for pi in order: |
| pos_i, neg_i = self.pairs[pi] |
| batch.extend([pos_i, neg_i]) |
| if len(batch) == 2 * self.pairs_per_batch: |
| yield batch |
| batch = [] |
| if batch: |
| yield batch |
|
|
| def __len__(self) -> int: |
| n = len(self.pairs) |
| if n == 0: |
| return 0 |
| n_rank = len(range(self.rank, n, self.num_replicas)) |
| return (n_rank + self.pairs_per_batch - 1) // self.pairs_per_batch |
|
|
|
|
| def labeled_collate_fn(batch, tokenizer, *, append_eos: bool = False, max_rna_nt: int | None = None): |
| labels = torch.tensor([float(s["label"]) for s in batch], dtype=torch.float32) |
| base = rna_rbp_collate_fn( |
| batch, |
| tokenizer, |
| append_eos=append_eos, |
| max_rna_nt=max_rna_nt, |
| ) |
| base["labels"] = labels |
| return base |
|
|
|
|
| def make_collate(tokenizer_path: str, append_eos: bool, max_rna_nt: int | None): |
| tok = AutoTokenizer.from_pretrained(tokenizer_path, trust_remote_code=True) |
| return partial( |
| labeled_collate_fn, |
| tokenizer=tok, |
| append_eos=append_eos, |
| max_rna_nt=max_rna_nt, |
| ) |
|
|