Cytrex's picture
Upload folder using huggingface_hub
5f26f25 verified
Raw History Blame Contribute Delete
4.89 kB
"""Turn labelled JSONL into training records for the Laya encoder.
A record is ONE text with ALL its questions. That is not a space-saving trick:
Laya answers several questions in a single forward pass, and RLCD trains on
episodes that carry more than one question.
Input format, one JSON object per line:
{"id": "...", "text": "...", "topic": "politics", "severity": 3, "violent": false}
Question definitions (--questions), a JSON object:
{
"topic": {
"type": "choice",
"instructions": "Which area does this text belong to?",
"criteria": {"politics": "government, elections, ...",
"business": "companies, markets, ..."}
},
"severity": {
"type": "score",
"instructions": "How severe is the described event?",
"criteria": ["none", "minor", "notable", "serious", "extreme"]
},
"violent": {
"type": "bool",
"instructions": "Does the text describe physical violence?",
"criteria": {"true": "yes, against people", "false": "no"}
}
}
Three question types, matching Laya's own:
choice unordered categories; the label is one of the criteria keys
score ordinal scale; the label is an integer index into the criteria list
bool yes/no; the label is a JSON boolean
A field missing from a line is simply skipped for that record -- questions do
not all have to be labelled on the same texts. That is how you mix a large,
cheaply labelled set for one question with a small, carefully labelled set for
another.
"""
import json
MAX_CHARS = 2800 # what fits comfortably into the encoder's window
def load_questions(path):
"""Read the question definitions and normalise them for the encoder."""
raw = json.load(open(path))
out = {}
for name, d in raw.items():
t = d["type"]
if t == "choice":
out[name] = {"t": "choice", "ins": d["instructions"],
"crit": d["criteria"], "keys": list(d["criteria"])}
elif t == "score":
out[name] = {"t": "score", "ins": d["instructions"], "crit": d["criteria"]}
elif t == "bool":
# the encoder calls this type "noul"
c = d["criteria"]
out[name] = {"t": "noul", "ins": d["instructions"],
"crit": {"true": c["true"], "false": c["false"]}}
else:
raise ValueError("unknown question type %r for %r" % (t, name))
return out
def _question(defs, name, value):
d = defs[name]
if d["t"] == "choice":
if value not in d["keys"]:
return None
y = d["keys"].index(value)
elif d["t"] == "score":
if not isinstance(value, int) or not (0 <= value < len(d["crit"])):
return None
y = value
else:
if value is None:
return None
y = 1 if value else 0
return {"t": d["t"], "ins": d["ins"], "crit": d["crit"], "y": y, "name": name}
def records(path, defs, max_chars=MAX_CHARS, only=None, quiet=False):
"""Build encoder records from a JSONL file.
Labels that do not fit their question definition are skipped -- an unknown
category, or an ordinal value outside the scale. That is the right thing to
do, but doing it SILENTLY is not: we lost 21 of 176 gold judgements that way
once, because a five-level scale had been collapsed to three levels in the
training data but not in the gold file. The evaluation still ran and still
printed a number. It just was not the number we thought.
So: count them, and say so.
"""
out = []
dropped = {}
for line in open(path):
x = json.loads(line)
qs = []
for name in (only or defs):
if name in x:
q = _question(defs, name, x[name])
if q:
qs.append(q)
else:
dropped[name] = dropped.get(name, 0) + 1
text = (x.get("text") or "").strip()
if qs and text:
out.append({"state": text[:max_chars], "qs": qs, "src": x.get("id", "")})
if dropped and not quiet:
detail = ", ".join("%s: %d" % kv for kv in sorted(dropped.items()))
print(" WARNING: %d labels in %s do not fit their question definition "
"and were skipped (%s)" % (sum(dropped.values()), path, detail),
flush=True)
return out
if __name__ == "__main__":
import collections, sys
if len(sys.argv) < 3:
sys.exit("usage: python data.py <records.jsonl> <questions.json>")
defs = load_questions(sys.argv[2])
r = records(sys.argv[1], defs)
print("questions:", list(defs))
print("records:", len(r))
print("questions per record:", dict(collections.Counter(len(x["qs"]) for x in r)))
if r:
print("median state length:", sorted(len(x["state"]) for x in r)[len(r) // 2])