azharmo commited on
Commit
11f8d43
·
verified ·
1 Parent(s): baf5397

Upload jev_toy/data.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. jev_toy/data.py +118 -0
jev_toy/data.py ADDED
@@ -0,0 +1,118 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ jev_toy/data.py
3
+
4
+ Turn real public HuggingFace datasets into System One training/eval examples,
5
+ each of the form (state_text, question_text, q_type, target).
6
+
7
+ We map three proven Jev question types onto real datasets:
8
+
9
+ * choice <- AG News (route a headline/topic into 1 of 4 categories)
10
+ * noul <- BoolQ (yes / no on a reading-comprehension question)
11
+ * score <- SST-2 (continuous "sentiment polarity" tag)
12
+
13
+ This is a toy-scale demonstration: we subsample so a laptop CPU can train in
14
+ minutes. The structure generalizes to the full datasets and to GPU.
15
+ """
16
+
17
+ from __future__ import annotations
18
+
19
+ import re
20
+ from collections import Counter
21
+
22
+ SPLIT_RE = re.compile(r"[A-Za-z0-9']+|[.,!?;:()\"]")
23
+
24
+
25
+ def tokenize(text: str, vocab, max_len: int, oov: int, pad: int):
26
+ ids = []
27
+ for tok in SPLIT_RE.findall(text.lower())[:max_len]:
28
+ ids.append(vocab.get(tok, oov)) # certain words map to a small rand-pool
29
+ ids = ids[:max_len]
30
+ mask = [1] * len(ids)
31
+ ids = ids + [pad] * (max_len - len(ids))
32
+ mask = mask + [0] * (max_len - len(mask))
33
+ return ids, mask
34
+
35
+
36
+ class Vocabulary:
37
+ def __init__(self, min_freq=1):
38
+ self.min_freq = min_freq
39
+ self.stoi = {}
40
+ self.itos = {}
41
+ self.oov = 0
42
+
43
+ def build(self, texts: list[str]):
44
+ c = Counter()
45
+ for t in texts:
46
+ c.update(SPLIT_RE.findall(t.lower()))
47
+ words = [w for w, n in c.items() if n >= self.min_freq]
48
+ words.sort()
49
+ self.stoi = {w: i + 2 for i, w in enumerate(words)} # 0=pad, 1=oov
50
+ self.itos = {i: w for w, i in self.stoi.items()}
51
+ self.oov = 1
52
+ return self
53
+
54
+ def __len__(self):
55
+ return len(self.stoi) + 2
56
+
57
+
58
+ def make_agnews(df) -> list[dict]:
59
+ labels = ["World", "Sports", "Business", "Sci/Tech"]
60
+ out = []
61
+ for txt, lab in zip(df["text"], df["label"]):
62
+ out.append({
63
+ "state": f"News article: {txt[:200]}",
64
+ "question": "Which topic does this article belong to?",
65
+ "q_type": "choice",
66
+ "target": int(lab),
67
+ "options": labels,
68
+ })
69
+ return out
70
+
71
+
72
+ def make_boolq(df) -> list[dict]:
73
+ out = []
74
+ for passage, q, ans in zip(df["passage"], df["question"], df["answer"]):
75
+ out.append({
76
+ "state": f"Passage: {passage[:250]}",
77
+ "question": f"Q: {q}",
78
+ "q_type": "noul",
79
+ "target": int(ans),
80
+ })
81
+ return out
82
+
83
+
84
+ def make_sst2(df) -> list[dict]:
85
+ out = []
86
+ for txt, lab in zip(df["sentence"], df["label"]):
87
+ out.append({
88
+ "state": f"Review: {txt[:150]}",
89
+ "question": "Is this review positive?",
90
+ "q_type": "noul", # treat sentiment polarity as a boolean tag
91
+ "target": int(lab),
92
+ })
93
+ return out
94
+
95
+
96
+ def build_examples(datasets: dict[str, object], subsample: dict[str, int]) -> tuple[list[dict], Vocabulary]:
97
+ """datasets: {name: HF-dataset}; subsample: {name: cap}."""
98
+ all_ex = []
99
+ # AG News
100
+ for split in ("train", "test"):
101
+ df = datasets["agnews"][split]
102
+ cap = subsample.get("agnews", 2000) if split == "train" else 400
103
+ all_ex += make_agnews(df.select(list(range(cap))))
104
+ # BoolQ
105
+ train = datasets["boolq"]["train"].select(list(range(subsample.get("boolq", 2000))))
106
+ val = datasets["boolq"]["validation"].select(list(range(400)))
107
+ all_ex += make_boolq(train) + make_boolq(val)
108
+ # SST-2
109
+ sst_train = datasets["sst2"]["train"].select(list(range(subsample.get("sst2", 2000))))
110
+ sst_val = datasets["sst2"]["validation"].select(list(range(400)))
111
+ all_ex += make_sst2(sst_train) + make_sst2(sst_val)
112
+ # vocab from state+question text
113
+ vocab = Vocabulary().build([e["state"] + " " + e["question"] for e in all_ex])
114
+ return all_ex, vocab
115
+
116
+
117
+ if __name__ == "__main__":
118
+ print("data module ok")