File size: 8,056 Bytes
29f25be
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
"""Shared corpus export contract for partitioned preprocessing results."""

import json
import math
from collections import Counter
from contextlib import ExitStack

from vimeml.data.build import SPLITS, WINNERS, PRIMARY_ROWS, find, sha, dump_line, dump_json


def export_corpus(db, output, config, calibration_ids, counts, rule_reasons,

                  fragment_flags, effective_blocks, annotation_count):
    settings = config['build']
    with ExitStack() as stack:
        forced_roots = set()
        for row in db.execute("SELECT DISTINCT doc_hash FROM annotations"):
            if db.execute("SELECT 1 FROM parents WHERE key=?", (row[0],)).fetchone():
                forced_roots.add(find(db, row[0]))
        for doc_id in sorted(calibration_ids):
            for row in db.execute(
                "SELECT DISTINCT doc_hash FROM origins WHERE doc_id = ?", (doc_id,)
            ):
                forced_roots.add(find(db, row["doc_hash"]))

        for row in db.execute("SELECT doc_hash FROM contents ORDER BY doc_hash"):
            root = find(db, row["doc_hash"])
            value = int(sha(f"{settings['seed']}:{root}")[:16], 16) / 2**64
            rank = (
                0 if root in forced_roots or value < config["split"]["train"]
                else 1 if value < config["split"]["train"] + config["split"]["validation"]
                else 2
            )
            db.execute(
                "UPDATE contents SET group_id = ?, split_rank = ? WHERE doc_hash = ?",
                (root, rank, row["doc_hash"]),
            )
        db.commit()
        db.executescript(WINNERS)
        # Text inspected while tuning the pipeline must not reappear in held-out
        # splits via an unreviewed document with the same sentence.
        db.execute("""

            UPDATE winners SET split_rank=0 WHERE text_hash IN (

                SELECT u.text_hash FROM units u JOIN contents c ON c.doc_hash=u.doc_hash

                WHERE c.group_id IN (SELECT value FROM json_each(?))

            )

        """, (json.dumps(sorted(forced_roots)),))

        split_stats = {
            name: {"sentences": 0, "characters": 0, "primary_sources": Counter()}
            for name in SPLITS
        }
        outputs = {
            name: (
                stack.enter_context((output / f"{name}.jsonl").open(
                    "w", encoding="utf-8", newline="\n"
                )),
                stack.enter_context((output / f"{name}.txt").open(
                    "w", encoding="utf-8", newline="\n"
                )),
            )
            for name in SPLITS
        }
        for row in db.execute(PRIMARY_ROWS):
            name = SPLITS[row["split_rank"]]
            record = {
                "text": row["text"], "text_hash": row["text_hash"],
                "source": row["source"], "sources": sorted(row["sources"].split(",")),
                "doc_id": row["doc_id"], "doc_hash": row["doc_hash"],
                "group_id": row["group_id"],
                "source_file": row["source_file"], "row_index": row["row_index"],
                "paragraph_index": row["paragraph_index"],
                "sentence_index": row["sentence_index"],
                "cleaned_block_spans": json.loads(row["spans"]),
                "quality_mode": settings["quality_mode"], "lm_reviewed": False,
                "approved_annotation_ids": json.loads(row["annotation_ids"]),
            }
            dump_line(outputs[name][0], record)
            outputs[name][1].write(row["text"] + "\n")
            split_stats[name]["sentences"] += 1
            split_stats[name]["characters"] += len(row["text"])
            split_stats[name]["primary_sources"][row["source"]] += 1

        with (output / "documents.jsonl").open("w", encoding="utf-8", newline="\n") as stream:
            for row in db.execute("""

                SELECT o.info, c.group_id, c.split_rank FROM origins o

                JOIN contents c ON c.doc_hash = o.doc_hash

                ORDER BY o.source_file, o.row_index

            """):
                dump_line(stream, {
                    **json.loads(row["info"]), "group_id": row["group_id"],
                    "assigned_split": SPLITS[row["split_rank"]],
                })

        # Full sentence provenance includes occurrences removed from another split.
        with (output / "provenance.jsonl").open("w", encoding="utf-8", newline="\n") as stream:
            for row in db.execute("""

                SELECT u.text_hash, u.doc_hash, u.paragraph_index, u.sentence_index,

                       o.doc_id, o.source, o.source_file, o.row_index,

                       c.group_id, u.annotation_ids, c.split_rank AS assigned, w.split_rank AS exported

                FROM units u JOIN contents c ON c.doc_hash = u.doc_hash

                JOIN origins o ON o.doc_hash = u.doc_hash

                JOIN winners w ON w.text_hash = u.text_hash

                ORDER BY u.text_hash, o.source_file, o.row_index,

                         u.paragraph_index, u.sentence_index

            """):
                dump_line(stream, {
                    key: row[key] for key in row.keys()
                    if key not in {"assigned", "exported", "annotation_ids"}
                } | {
                    "approved_annotation_ids": json.loads(row["annotation_ids"]),
                    "assigned_split": SPLITS[row["assigned"]],
                    "exported_split": SPLITS[row["exported"]],
                    "retained_in_assigned_split": row["assigned"] == row["exported"],
                })

        unique_sentences = db.execute("SELECT COUNT(*) FROM winners").fetchone()[0]
        cross_split_removed = db.execute("""

            SELECT COUNT(*) FROM units u

            JOIN contents c ON c.doc_hash = u.doc_hash

            JOIN winners w ON w.text_hash = u.text_hash

            WHERE c.split_rank != w.split_rank

        """).fetchone()[0]
        characters = sum(item["characters"] for item in split_stats.values())
        matched_annotations = db.execute("SELECT COUNT(*) FROM annotations WHERE matched=1").fetchone()[0]
        unmatched_annotations = [row[0] for row in db.execute("SELECT id FROM annotations WHERE matched=0 ORDER BY id")]
        dump_json(output / "unmatched_annotations.json", unmatched_annotations)
        if config.get("review", {}).get("require_all_annotations", False) and unmatched_annotations:
            raise ValueError("Approved annotations did not match this build; inspect unmatched_annotations.json.")
        stats = {
            **dict(counts),
            "duplicate_documents": counts["document_origins"] - counts["unique_documents"],
            "unique_sentences": unique_sentences,
            "duplicate_sentence_occurrences": counts["sentence_occurrences"] - unique_sentences,
            "cross_split_occurrences_removed": cross_split_removed,
            "final_characters": characters,
            "splits": {
                name: {**item, "primary_sources": dict(item["primary_sources"])}
                for name, item in split_stats.items()
            },
            "rule_reasons": dict(rule_reasons),
            "block_count_note": "keep/drop/review_blocks count rule-stage decisions before annotations.",
            "effective_block_actions": dict(effective_blocks),
            "fragment_flags": dict(fragment_flags),
            "forced_train_groups": len(forced_roots),
            "annotation_records": annotation_count, "matched_annotations": matched_annotations,
            "unmatched_annotations": len(unmatched_annotations),
            "token_estimate_scenarios": {
                f"{n}_characters_per_token": math.ceil(characters / n) for n in (1, 2, 3)
            },
            "token_estimate_note": "Heuristic scenarios, not measured SentencePiece counts.",
        }
        dump_json(output / "stats.json", stats)
        db.commit()

    return stats