"""Zero-shot evaluation sources. These datasets never appear in train/val/test, and their tasks (question-type classification, formality rating, spam filtering) have no counterpart in the training mix, so heldout.jsonl measures generalization to unseen tasks. """ from __future__ import annotations from jevlike.data.common import ( Example, Label, Scale, balanced_order, clean, cuts, dedupe_rows, derived_noul, make_choice, make_score, noul, pick_instr, reverse, truncate, weighted, ) from jevlike.data.hub import read_many, read_rows from jevlike.data.label_names import TREC_FINE from jevlike.data.sources_choice import GENERIC_CHOICE, Ctx TREC_COARSE = [ Label("abbreviation", ["asks what an abbreviation stands for, or for an abbreviation"], "an abbreviation", key="ABBR"), Label("entity", ["asks for a thing: an animal, food, product, color, event, sport, etc."], "an entity or thing", key="ENTY"), Label("description", ["asks for a definition, description, manner or reason"], "a description or explanation", key="DESC"), Label("human", ["asks about a person or a group of people"], "a person or group of people", key="HUM"), Label("location", ["asks for a place: city, country, mountain, etc."], "a location", key="LOC"), Label("numeric value", ["asks for a number, date, amount, distance, duration, etc."], "a number or quantity", key="NUM"), ] _FINE_DESC = { "abb": "abbreviation", "exp": "expansion of an abbreviation", "animal": "animal", "body": "body part or organ", "color": "color", "cremat": "creative work (book, film, song, invention)", "currency": "currency", "dismed": "disease or medicine", "event": "event", "food": "food", "instru": "musical instrument", "lang": "language", "letter": "letter or character", "plant": "plant", "product": "product", "religion": "religion", "sport": "sport", "substance": "substance, element or material", "symbol": "symbol or sign", "techmeth": "technique or method", "termeq": "equivalent term", "veh": "vehicle", "word": "word with a special property", "def": "definition of something", "manner": "manner of an action (how)", "reason": "reason (why)", "gr": "group or organization", "ind": "individual person", "title": "title or position of a person", "city": "city", "country": "country", "mount": "mountain", "state": "state or province", "code": "postal code, phone number or other code", "count": "count (how many)", "date": "date or time", "dist": "distance or length", "money": "price or amount of money", "ord": "rank or order", "period": "duration", "perc": "percentage or fraction", "speed": "speed", "temp": "temperature", "volsize": "size, area or volume", "weight": "weight", } def _fine_label(n: str) -> Label: coarse, fine = n.split(":") if fine == "other": d = {"ENTY": "other entity", "LOC": "other location", "NUM": "other number"}[coarse] elif fine == "desc": d = "description of a person" if coarse == "HUM" else "description of something" else: d = _FINE_DESC[fine] return Label(d, [f"the answer is a {d}", d], f"a {d}", key=n.replace(":", "_").lower()) def trec(ctx: Ctx): rows = read_many("CogComp/trec", ["default/train/0000.parquet", "default/test/0000.parquet"], revision="refs/convert/parquet") rng = ctx.rng fine = [_fine_label(n) for n in TREC_FINE] rows = dedupe_rows(rows, lambda r: r["text"], lambda r: r["fine_label"]) c_instr = ["What type of answer does this question expect?", "Classify the question by answer type.", "What is the question asking for?", "Which kind of information is being requested?"] pos = ["The question asks for {x}.", "The expected answer is {x}.", "Is this question asking for {x}?"] neg = ["The question does not ask for {x}."] for r in balanced_order(rng, rows, lambda r: int(r["coarse_label"]), 0.5): text = clean(r["text"]) c, f = int(r["coarse_label"]), int(r["fine_label"]) kind = weighted(rng, ["coarse", "fine", "noul"], [0.5, 0.2, 0.3]) if kind == "coarse": q = make_choice(rng, TREC_COARSE, c, pick_instr(rng, c_instr, GENERIC_CHOICE), n=6 if rng.random() < 0.7 else None) elif kind == "fine": q = make_choice(rng, fine, f, pick_instr(rng, c_instr, GENERIC_CHOICE), full_prob=0.15) else: q = derived_noul(rng, TREC_COARSE, c, pos, neg, p_other_group=0.0) yield Example(text, [q]) def formality(ctx: Ctx): rows = read_many("osyvokon/pavlick-formality-scores", ["train.csv", "test.csv"]) rng = ctx.rng rows = dedupe_rows(rows, lambda r: str(r["sentence"]), lambda r: round(float(r["avg_score"]), 1)) ins = ["How formal is this sentence?", "Rate the formality of the text.", "What register is this written in?", "Formality level?"] scales = [ Scale(["very informal", "informal", "neutral", "formal", "very formal"], cuts(-1.8, -0.6, 0.6, 1.8), ins, 1.3), Scale(["casual", "neutral", "formal"], cuts(-0.7, 0.7), ins, 1.0), Scale(["slang or chatty", "everyday conversational", "plain neutral prose", "professional", "academic or legal"], cuts(-1.8, -0.6, 0.6, 1.8), ins, 0.6), ] scales.append(reverse(scales[1], ["How casual is this sentence?", "Rate how informal the text sounds."], weight=0.6)) stm = [("The sentence is written in a formal style.", lambda v: True if v >= 1.0 else (False if v <= -1.0 else None)), ("This sounds casual and conversational.", lambda v: True if v <= -1.0 else (False if v >= 1.0 else None)), ("This text would fit in an official document.", lambda v: True if v >= 1.5 else (False if v <= -0.5 else None))] for r in balanced_order(rng, rows, lambda r: cuts(-1.8, -0.6, 0.6, 1.8)(float(r["avg_score"])), 0.5): v = float(r["avg_score"]) text = clean(r["sentence"]) if rng.random() < 0.7: q = make_score(rng, scales, v) else: q = None for _ in range(4): s, fn = rng.choice(stm) t = fn(v) if t is not None: q = noul(s, t) break if q: yield Example(text, [q]) def enron(ctx: Ctx): rows = read_rows("SetFit/enron_spam", "test.jsonl") rng = ctx.rng def text(r): subj, msg = clean(r.get("subject") or ""), clean(r.get("message") or "") return truncate(f"Subject: {subj}\n\n{msg}" if subj else msg) rows = [dict(r, _t=text(r)) for r in rows] rows = dedupe_rows(rows, lambda r: r["_t"], lambda r: r["label"]) labels = [Label("inbox", ["legitimate email (ham)", "normal work or personal email"], key="ham"), Label("spam", ["unsolicited bulk or advertising email", "junk / spam"])] pos = ["This email is spam.", "This is an unsolicited bulk or advertising email.", "Is this email spam?", "The message is junk mail."] neg = ["This is a legitimate email.", "This email is not spam.", "The message is a normal work email."] instr = ["Which folder should this email go to?", "Spam or ham?", "Filter this email.", "How should this message be filed?"] for r in balanced_order(rng, rows, lambda r: int(r["label"]), 0.0): y = int(r["label"]) if rng.random() < 0.3: q = make_choice(rng, labels, y, pick_instr(rng, instr), n=2) elif rng.random() < 0.3: q = noul(rng.choice(neg), y == 0) else: q = noul(rng.choice(pos), y == 1) yield Example(r["_t"], [q]) SOURCES = {"trec": trec, "pavlick_formality": formality, "enron_spam": enron}