ProRiboGen / classifier /labeled_dataset.py
TimelessAEther's picture
Upload ProRiboGen inference package and checkpoints
6dd9839 verified
Raw
History Blame Contribute Delete
4.78 kB
"""带 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
# 网站部署包:generation 模块代码在 ../generator
_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 # noqa: E402
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,
)