kodama-core / runtime /jevlike /data /sources_heldout.py
Cem13's picture
Release Kodama Core with weights, runtime, attribution, and evaluation
c7893fa verified
Raw History Blame Contribute Delete
7.75 kB
"""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}