File size: 22,024 Bytes
685e018
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
"""Synthetic multi-turn claim intake: dialogues with a known fact set, and a scorer.

The target behaviour is memory rather than knowledge: every fact the model must report is
present in the conversation, so a small model is not being asked to recall the world. Some
turns supersede an earlier value, which is what separates tracking state from copying the
last thing seen. Because the record is generated, recall and precision are exact, and no
judge is needed to score a recap.
"""

from __future__ import annotations

import argparse
import json
import random
import re
from dataclasses import dataclass, field
from pathlib import Path

from diffusion_lm.reasoning import ReasoningExample

FIRST_NAMES = ('Marta', 'Diego', 'Luciana', 'Rodrigo', 'Camila', 'Nestor', 'Sofia', 'Ariel')
LAST_NAMES = ('Quiroga', 'Benitez', 'Salvatierra', 'Uriarte', 'Ferreyra', 'Zabala', 'Otamendi')
MAKES = (('Peugeot', '208'), ('Toyota', 'Etios'), ('Renault', 'Sandero'), ('Fiat', 'Cronos'),
         ('Chevrolet', 'Onix'), ('Volkswagen', 'Gol'))
STREETS = ('Av. Rivadavia', 'Calle Mitre', 'Ruta 8', 'Av. San Martin', 'Calle Belgrano')
CITIES = ('Moron', 'Rosario', 'La Plata', 'Cordoba', 'Bahia Blanca', 'Mendoza')
DAMAGE_PARTS = ('front bumper', 'left headlight', 'driver door', 'rear hatch', 'right mirror',
                'windshield', 'rear bumper')
WEATHER = ('heavy rain', 'clear skies', 'fog', 'light drizzle')
FILLERS = ('Sorry, one moment.', 'Are you still there?', 'This is my first claim.',
           'Ok.', 'Thanks for the help.', 'Can you repeat that?')
ACKS = ('No problem, take your time.', 'Yes, I am here.', 'Understood, please go on.',
        'Of course.', 'Happy to help.')


@dataclass
class Claim:
    """One claim record plus the facts a faithful report has to carry."""

    fields: dict[str, str]
    corrected: dict[str, str] = field(default_factory=dict)

    @property
    def truth(self) -> dict[str, str]:
        """Field values after corrections, which is what a report must state."""

        resolved = dict(self.fields)
        resolved.update(self.corrected)
        return resolved


def _plate(rng: random.Random) -> str:
    letters = ''.join(rng.choice('ABCDEFGHJKLMNPRSTUVWXYZ') for _ in range(3))
    return f'{letters}-{rng.randint(1000, 9999)}'


def generate_claim(rng: random.Random) -> Claim:
    """Build a claim whose values are distinctive enough to score by exact match."""

    make, model = rng.choice(MAKES)
    fields = {
        'policy_number': f'PL-{rng.randint(10000, 99999)}',
        'claimant': f'{rng.choice(FIRST_NAMES)} {rng.choice(LAST_NAMES)}',
        'incident_date': f'{rng.randint(1, 28):02d}/{rng.randint(1, 12):02d}/2026',
        'incident_time': f'{rng.randint(0, 23):02d}:{rng.choice(("05", "15", "40", "50"))}',
        'location': f'{rng.choice(STREETS)}, {rng.choice(CITIES)}',
        'vehicle': f'{make} {model} {rng.randint(2012, 2025)}',
        'plate': _plate(rng),
        'weather': rng.choice(WEATHER),
        'damage': rng.choice(DAMAGE_PARTS),
        'estimate': f'${rng.randint(2, 40) * 1000 + rng.choice((150, 450, 800)):,}',
        'other_plate': _plate(rng),
        'police_report': f'PR-{rng.randint(100000, 999999)}',
        'witness': f'{rng.choice(FIRST_NAMES)} {rng.choice(LAST_NAMES)}',
    }
    return Claim(fields=fields)


QUESTIONS = {
    'policy_number': 'Can you give me your policy number?',
    'claimant': 'Who is the policy holder?',
    'incident_date': 'What date did this happen?',
    'incident_time': 'Roughly what time?',
    'location': 'Where exactly did it happen?',
    'vehicle': 'Which vehicle was involved?',
    'plate': "What is your vehicle's plate?",
    'weather': 'How was the weather at the time?',
    'damage': 'What part of the car was damaged?',
    'estimate': 'Do you have a repair estimate?',
    'other_plate': 'Did you get the other vehicle plate?',
    'police_report': 'Was a police report filed?',
    'witness': 'Was there any witness?',
}

ANSWERS = {
    'policy_number': 'My policy is {value}.',
    'claimant': 'The holder is {value}.',
    'incident_date': 'It was on {value}.',
    'incident_time': 'Around {value}.',
    'location': 'On {value}.',
    'vehicle': "It's a {value}.",
    'plate': 'The plate is {value}.',
    'weather': 'There was {value}.',
    'damage': 'The {value} took the hit.',
    'estimate': 'The shop quoted {value}.',
    'other_plate': 'Yes, {value}.',
    'police_report': 'Yes, report {value}.',
    'witness': '{value} saw everything.',
}


def render_dialogue(
    claim: Claim, rng: random.Random, *, turns: int = 6, corrections: int = 1,
    chitchat: int = 2,
) -> list[dict[str, str]]:
    """Reveal the record across turns, superseding some values along the way.

    Each turn carries one to three fields, so the amount to remember per turn varies the
    way it would in a real intake.
    """

    keys = list(claim.fields)
    rng.shuffle(keys)
    # Deliberately uneven: with a constant number of facts per turn every ledger block lands on
    # the same size and the control decision stops carrying information.
    batches: list[list[str]] = []
    while keys:
        take = min(len(keys), rng.choice((1, 1, 2, 3, 4, 5)))
        batches.append(keys[:take])
        keys = keys[take:]

    messages: list[dict[str, str]] = []
    for batch in batches:
        messages.append({'role': 'assistant', 'content': ' '.join(QUESTIONS[k] for k in batch)})
        said = ' '.join(ANSWERS[k].format(value=claim.fields[k]) for k in batch)
        messages.append({'role': 'user', 'content': said,
                         'values': {k: claim.fields[k] for k in batch}})

    # Turns that carry no new fact, so the assistant answers with zero thinking blocks. A
    # conversation is full of these and without them the model never learns to skip thinking.
    for _ in range(chitchat):
        at = rng.randrange(1, max(2, len(messages)))
        messages.insert(at, {'role': 'assistant', 'content': rng.choice(ACKS)})
        messages.insert(at, {'role': 'user', 'content': rng.choice(FILLERS), 'values': {}})

    correctable = [k for k in claim.fields if k in ('plate', 'other_plate', 'estimate',
                                                    'incident_time', 'police_report')]
    rng.shuffle(correctable)
    for key in correctable[:corrections]:
        if key.endswith('plate'):
            new = _plate(rng)
        elif key == 'estimate':
            new = f'${rng.randint(2, 40) * 1000 + 700:,}'
        elif key == 'incident_time':
            new = f'{rng.randint(0, 23):02d}:30'
        else:
            new = f'PR-{rng.randint(100000, 999999)}'
        claim.corrected[key] = new
        messages.append({
            'role': 'user',
            'content': f'Sorry, I misspoke earlier: {ANSWERS[key].format(value=new)} '
                       f'Not {claim.fields[key]}.',
            'values': {key: new},
        })
        # An assistant turn right after the correction is what makes the update a training
        # target: without it the only place the corrected value is taught is the final report.
        messages.append({'role': 'assistant', 'content': 'Noted, I corrected that detail.'})
    return messages


_IDENTIFIERS = re.compile(
    r'(?:P[LR]-\d{5,6}'                    # policy and police report numbers
    r'|[A-Z]{3}-\d{4}'                     # plates
    r'|\$[\d,]+'                           # money
    r'|\d{2}/\d{2}/\d{4}'                  # dates
    r'|\b\d{2}:\d{2}\b)'                   # times
)


def score_report(text: str, claim: Claim) -> dict[str, object]:
    """Exact-match recall of the resolved values, plus the two ways a report lies.

    ``stale_kept`` reports a value the conversation replaced, which means the model copied
    instead of tracking state. ``invented`` counts identifier-shaped strings that were never
    said at all, which is confabulation rather than a memory slip; they are different
    failures and worth separating.
    """

    truth = claim.truth
    found = {key: (value in text) for key, value in truth.items()}
    stale = {
        key: (claim.fields[key] in text)
        for key in claim.corrected
        if claim.fields[key] != claim.corrected[key]
    }
    spoken = set(truth.values()) | set(claim.fields.values())
    said_ids = {token for value in spoken for token in _IDENTIFIERS.findall(value)}
    invented = sorted({token for token in _IDENTIFIERS.findall(text)} - said_ids)
    recalled = sum(found.values())
    return {
        'fields': len(truth),
        'recalled': recalled,
        'recall': recalled / max(1, len(truth)),
        'missing': sorted(k for k, ok in found.items() if not ok),
        'stale_kept': sorted(k for k, bad in stale.items() if bad),
        'invented': invented,
    }


LEDGER = '{key}: {value}'
KEEP_MESSAGES = 4
IM_START = '<|im_start|>'
IM_END = '<|im_end|>'
SYSTEM = (
    'You are a claim intake assistant. Track every detail the customer gives and always use '
    'the corrected value when they correct themselves.'
)


def chatml_turn(role: str, content: str) -> str:
    """One ChatML turn. Both markers are single tokens in the adaptive tokenizer."""

    return f'{IM_START}{role}\n{content}{IM_END}\n'


def ledger_line(entries: list[str]) -> str:
    """The ``Known so far`` system-turn line carrying facts whose messages were dropped."""

    return 'Known so far: ' + '; '.join(entries) + '.' if entries else ''


_NOTE_FACT = re.compile(r'^([^:]{1,48}): (.+)$')


def merge_notes(notes: list[str]) -> list[str]:
    """Fold note fragments into one entry per key, latest value winning.

    Concatenating raw notes would re-expose superseded values, which is the failure the corpus
    charges hardest. Fragments that do not parse as ``key: value`` pass through in order,
    deduplicated verbatim, so an unkeyed reasoning step still reaches the prefix as prose.

    Both the training renderer and the playground call this. Keeping one implementation is the
    point: a second copy is how the train and inference views of the same history drift apart.
    """

    facts: dict[str, str] = {}
    loose: list[str] = []
    for note in notes:
        for fragment in note.split('; '):
            fragment = fragment.strip().rstrip('.')
            if not fragment:
                continue
            match = _NOTE_FACT.match(fragment)
            if match:
                facts[match.group(1)] = match.group(2)
            elif fragment not in loose:
                loose.append(fragment)
    return [f'{key}: {value}' for key, value in facts.items()] + loose


def window_start(messages: list[dict[str, str]], keep: int) -> int:
    return 0 if keep <= 0 else max(0, len(messages) - keep)


def ledger_notes(messages: list[dict[str, str]], keep: int) -> list[str]:
    """Notes whose user message fell out of the visible window, in turn order.

    An assistant turn's note enters the ledger exactly when the user message it took notes on
    has fallen out of the window, which is what makes the note the only remaining carrier.
    """

    start = window_start(messages, keep)
    return [
        message['note'] for index, message in enumerate(messages)
        if message['role'] == 'assistant' and message.get('note') and index - 1 < start
    ]


def chat_prefix(turns: list[dict[str, str]], *, system: str = SYSTEM, extra: str = '') -> str:
    """Conversation prefix ending where the assistant's generation begins.

    The accumulated ledger rides in the system turn rather than as a fake dialogue message:
    it is persistent state, and putting it there also teaches the model to condition on a
    system prompt, which none of the other corpora do.
    """

    merged = system if not extra else f'{system}\n{extra}'
    rendered = [chatml_turn('system', merged)]
    rendered += [chatml_turn(turn['role'], turn['content']) for turn in turns]
    return ''.join(rendered) + f'{IM_START}assistant\n'


def _chunks(items: list, size: int) -> list[list]:
    return [items[i:i + size] for i in range(0, len(items), size)]


def report_text(claim: Claim) -> str:
    """The report a faithful assistant produces, one resolved field per line."""

    return 'Claim report. ' + ' '.join(
        LEDGER.format(key=key, value=value) + '.' for key, value in claim.truth.items()
    )


def to_examples(
    claim: Claim, messages: list[dict[str, str]], *, keep_messages: int = KEEP_MESSAGES
) -> list[ReasoningExample]:
    """One training example per assistant turn, with the history deliberately truncated.

    Older messages are dropped and replaced by the ledger of what they revealed, so the
    accumulated notes — not the transcript — are what carries the past. That is the whole
    point: with the full transcript in the prefix the model can re-read instead of remember,
    and the thinking block stops being memory. A final example asks for the report, whose
    thinking consolidates every fact.

    Thought steps hold one fact each, which keeps a block's content short and its length a
    function of how much the turn actually revealed.
    """

    examples: list[ReasoningExample] = []
    # Values as stated at each point, so a ledger never shows a correction that has not
    # happened yet: training on resolved values would teach the model to know the future.
    known: dict[str, str] = {}
    for index, message in enumerate(messages):
        if message['role'] != 'assistant' or index == 0:
            continue
        previous = messages[index - 1]
        learned = dict(previous.get('values', {}))
        known.update(learned)
        dropped = messages[max(0, index - keep_messages):index]
        seen: dict[str, str] = {}
        for earlier in messages[:max(0, index - keep_messages)]:
            seen.update(earlier.get('values', {}))
        older = [LEDGER.format(key=key, value=value) for key, value in seen.items()]
        examples.append(ReasoningExample(
            problem=chat_prefix(dropped, extra=ledger_line(older)),
            steps=(('; '.join(LEDGER.format(key=key, value=value)
                                for key, value in learned.items()),)
                   if learned else ()),
            answer=message['content'] + IM_END,
            expected_answer='',
        ))

    tail = messages[-keep_messages:] + [{'role': 'user', 'content': 'Write the claim report.'}]
    examples.append(ReasoningExample(
        problem=chat_prefix(tail),
        steps=tuple(
            '; '.join(LEDGER.format(key=key, value=value) for key, value in group)
            for group in _chunks(list(claim.truth.items()), 8)
        ),
        answer=report_text(claim) + IM_END,
        expected_answer='',
    ))
    return examples


def build(count: int, seed: int, turns: int, corrections: int) -> list[dict[str, object]]:
    rng = random.Random(seed)
    records = []
    for index in range(count):
        claim = generate_claim(rng)
        messages = render_dialogue(claim, rng, turns=turns, corrections=corrections)
        records.append({
            'index': index,
            'messages': messages,
            'truth': claim.truth,
            'superseded': {k: claim.fields[k] for k in claim.corrected},
        })
    return records


def prepare(args: argparse.Namespace) -> None:
    """Pack dialogues into the adaptive layout, splitting BY DIALOGUE.

    Splitting by example would leak: two examples from one dialogue share its facts, so a
    validation example's answer would already appear in a training example's prefix.
    """

    import numpy as np

    from diffusion_lm.reasoning import ExampleEncoder, LayoutSpec, _write_packed
    from diffusion_lm.tokenizer import load_tokenizer

    tokenizer = load_tokenizer(args.tokenizer)
    spec = LayoutSpec(seq_len=args.seq_len, block=min(args.sizes), max_slots=args.max_slots,
                      sizes=tuple(sorted(args.sizes)))
    encoder = ExampleEncoder(tokenizer, spec)
    rng = random.Random(args.seed)

    split: dict[str, list[tuple]] = {'train': [], 'validation': []}
    dropped = 0
    for index in range(args.count):
        claim = generate_claim(rng)
        messages = render_dialogue(claim, rng, turns=args.turns,
                                   corrections=args.corrections,
                                   chitchat=args.chitchat)
        bucket = 'validation' if index % args.val_every == 0 else 'train'
        for example in to_examples(claim, messages, keep_messages=args.keep_messages):
            encoded = encoder.encode_adaptive(example)
            if encoded is None:
                dropped += 1
                continue
            split[bucket].append((encoded.tokens, encoded.regions))

    for name, rows in split.items():
        if not rows:
            raise ValueError(f'no examples in the {name} split')
        _write_packed(
            args.output_dir / f'{name}-adaptive.bin',
            np.stack([tokens for tokens, _ in rows]),
            np.stack([regions for _, regions in rows]),
            layout='adaptive', spec=spec, tokenizer_path=args.tokenizer, tokenizer=tokenizer,
            extra_metadata={'sizes': list(spec.sizes), 'source': 'claims-chatml'},
        )
        print(f'{name}: {len(rows):,} examples -> {args.output_dir}')
    print(f'{dropped:,} dropped at encode')


def mix(args: argparse.Namespace) -> None:
    """Concatenate two packs so claims hold ``--claims-share`` of the examples."""

    import numpy as np

    from diffusion_lm.reasoning import regions_path

    rng = np.random.default_rng(args.seed)
    for name in ('train', 'validation'):
        parts = []
        for directory, share in ((args.claims_dir, args.claims_share), (args.other_dir, None)):
            path = directory / f'{name}-adaptive.bin'
            meta = json.loads(Path(str(path) + '.json').read_text())
            tokens = np.fromfile(path, dtype=np.dtype(meta['dtype'])).reshape(
                meta['example_count'], meta['seq_len']
            )
            parts.append((tokens, np.load(regions_path(path)), share))
        (claims_tokens, claims_regions, share), (other_tokens, other_regions, _) = parts
        target = int(round(share / (1.0 - share) * len(other_tokens)))
        if target < len(claims_tokens):
            keep = rng.choice(len(claims_tokens), size=target, replace=False)
            claims_tokens, claims_regions = claims_tokens[keep], claims_regions[keep]
        tokens = np.concatenate([claims_tokens, other_tokens])
        regions = np.concatenate([claims_regions, other_regions])
        order = rng.permutation(len(tokens))
        tokens, regions = tokens[order], regions[order]
        args.output_dir.mkdir(parents=True, exist_ok=True)
        tokens.tofile(args.output_dir / f'{name}-adaptive.bin')
        np.save(regions_path(args.output_dir / f'{name}-adaptive.bin'), regions)
        meta.update({'example_count': int(len(tokens)), 'source': 'claims+glaive',
                     'claims_examples': int(len(claims_tokens)),
                     'other_examples': int(len(other_tokens))})
        Path(str(args.output_dir / f'{name}-adaptive.bin') + '.json').write_text(
            json.dumps(meta, indent=2) + '\n'
        )
        actual = len(claims_tokens) / len(tokens)
        print(f'{name}: {len(tokens):,} examples, claims share {actual:.3f}')


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    sub = parser.add_subparsers(dest='command', required=True)

    dialogues = sub.add_parser('build', help='write dialogues as JSONL for probing')
    dialogues.add_argument('--count', type=int, default=64)
    dialogues.add_argument('--seed', type=int, default=1337)
    dialogues.add_argument('--turns', type=int, default=6)
    dialogues.add_argument('--corrections', type=int, default=1)
    dialogues.add_argument('--output', type=Path, required=True)

    pack = sub.add_parser('prepare', help='pack dialogues into the adaptive layout')
    pack.add_argument('--count', type=int, default=2700)
    pack.add_argument('--seed', type=int, default=1337)
    pack.add_argument('--turns', type=int, default=6)
    pack.add_argument('--corrections', type=int, default=1)
    pack.add_argument('--chitchat', type=int, default=2,
                      help='turns with no new fact, which train zero-block answers')
    pack.add_argument('--keep-messages', type=int, default=KEEP_MESSAGES)
    pack.add_argument('--tokenizer', type=Path, required=True)
    pack.add_argument('--output-dir', type=Path, required=True)
    pack.add_argument('--seq-len', type=int, default=2048)
    pack.add_argument('--max-slots', type=int, default=64)
    pack.add_argument('--sizes', type=int, nargs='+', default=[32, 64, 128])
    pack.add_argument('--val-every', type=int, default=20,
                      help='every Nth dialogue goes to validation, whole')

    blend = sub.add_parser('mix', help='blend a claims pack into another pack')
    blend.add_argument('--claims-dir', type=Path, required=True)
    blend.add_argument('--other-dir', type=Path, required=True)
    blend.add_argument('--output-dir', type=Path, required=True)
    blend.add_argument('--claims-share', type=float, default=0.30)
    blend.add_argument('--seed', type=int, default=1337)

    args = parser.parse_args()
    if args.command == 'prepare':
        prepare(args)
        return
    if args.command == 'mix':
        mix(args)
        return

    records = build(args.count, args.seed, args.turns, args.corrections)
    args.output.parent.mkdir(parents=True, exist_ok=True)
    with args.output.open('w', encoding='utf-8') as handle:
        for record in records:
            handle.write(json.dumps(record, ensure_ascii=False) + '\n')
    turns = sum(len(r['messages']) for r in records) / len(records)
    facts = sum(len(r['truth']) for r in records) / len(records)
    print(f'wrote {len(records)} dialogues to {args.output} '
          f'({turns:.1f} messages and {facts:.1f} facts each)')


if __name__ == '__main__':
    main()