File size: 4,782 Bytes
6dd9839
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
"""带 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,
    )