"""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 ") 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])