| """Render the general-chat corpus into the adaptive training layout. |
| |
| The corpus (``scripts/gen_chat_dataset.py``) stores conversations as system prompt plus turns, |
| each assistant turn carrying its own notes. This module turns them into the layout the hybrid |
| objective trains on: one example per assistant turn, whose prefix holds the system prompt, the |
| merged ledger of notes whose messages have fallen out of the visible window, and the last |
| ``keep_messages`` messages verbatim. |
| |
| Dropping the older messages is the point. With the full transcript in the prefix the model can |
| re-read instead of remember and the thinking block stops being memory, which is what |
| :func:`diffusion_lm.claims.to_examples` established for the claims corpus. The window and the |
| ledger merge are imported from that module rather than reimplemented, so the prefix a training |
| example sees is byte-identical to the one the playground builds at inference. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import random |
| from collections import Counter |
| from pathlib import Path |
|
|
| from diffusion_lm.claims import ( |
| IM_END, |
| KEEP_MESSAGES, |
| chat_prefix, |
| ledger_line, |
| ledger_notes, |
| merge_notes, |
| ) |
| from diffusion_lm.reasoning import ReasoningExample |
|
|
| NOTE_JOIN = '; ' |
|
|
|
|
| def load(paths: list[Path]) -> list[dict]: |
| """Read consolidated corpus files, tagging each conversation with its origin file.""" |
|
|
| conversations = [] |
| for path in paths: |
| for line in path.open(encoding='utf-8'): |
| if not line.strip(): |
| continue |
| conversation = json.loads(line) |
| conversation.setdefault('source', path.stem) |
| conversations.append(conversation) |
| return conversations |
|
|
|
|
| def _messages(conversation: dict) -> list[dict[str, str]]: |
| """Corpus turns in the message shape the ledger helpers expect. |
| |
| A turn's notes collapse into one ``note`` string joined by ``NOTE_JOIN``, which is the |
| separator :func:`diffusion_lm.claims.merge_notes` splits on, so a multi-fact turn still |
| contributes one ledger entry per fact. |
| """ |
|
|
| messages = [] |
| for turn in conversation.get('turns') or []: |
| message = {'role': turn['role'], 'content': turn['content']} |
| notes = [str(note).strip() for note in (turn.get('thinking') or []) if str(note).strip()] |
| if turn['role'] == 'assistant' and notes: |
| message['note'] = NOTE_JOIN.join(notes) |
| messages.append(message) |
| return messages |
|
|
|
|
| def to_examples( |
| conversation: dict, *, keep_messages: int = KEEP_MESSAGES |
| ) -> list[ReasoningExample]: |
| """One example per assistant turn, each thinking note becoming its own block. |
| |
| A turn the corpus marked as needing no notes yields an empty chain, which is what teaches |
| the controller to answer without opening a thinking block; the encoder accepts it. |
| """ |
|
|
| messages = _messages(conversation) |
| turns = conversation.get('turns') or [] |
| reference = str(conversation.get('reference') or '') |
| last = max((i for i, m in enumerate(messages) if m['role'] == 'assistant'), default=-1) |
|
|
| examples = [] |
| for index, message in enumerate(messages): |
| if message['role'] != 'assistant' or index == 0: |
| continue |
| history = messages[:index] |
| older = merge_notes(ledger_notes(history, keep_messages)) |
| window = history[max(0, index - keep_messages):] |
| notes = [str(n).strip() for n in (turns[index].get('thinking') or []) if str(n).strip()] |
| examples.append(ReasoningExample( |
| problem=chat_prefix(window, system=conversation['system'], extra=ledger_line(older)), |
| steps=tuple(notes), |
| answer=message['content'] + IM_END, |
| expected_answer=reference if index == last else '', |
| )) |
| return examples |
|
|
|
|
| def _document(conversation: dict, index: int) -> str: |
| """Split key. Two conversations built from one passage share its facts. |
| |
| Splitting by example would leak within a conversation as well, so the whole conversation |
| travels together and grounded slices travel with their source item. |
| """ |
|
|
| return str(conversation.get('source_id') or f'{conversation.get("source", "")}-{index}') |
|
|
|
|
| def _apply_caps( |
| conversations: list[dict], caps: dict[str, int], seed: int |
| ) -> list[dict]: |
| """Drop conversations so a source contributes at most ``caps[source]`` of them. |
| |
| Capping is by CONVERSATION but the reason is examples: a source's weight in the mix is its |
| turn count, not its row count, and the two differ by an order of magnitude (CoQA yields 11.9 |
| examples per conversation against 1.07 for a single-question source). Sampling is seeded and |
| whole conversations travel together, so the split stays document-clean. |
| """ |
|
|
| if not caps: |
| return conversations |
| rng = random.Random(seed) |
| by_source: dict[str, list[int]] = {} |
| for index, conversation in enumerate(conversations): |
| by_source.setdefault(conversation.get('source', ''), []).append(index) |
| dropped: set[int] = set() |
| for source, limit in caps.items(): |
| indices = by_source.get(source) |
| if indices is None: |
| raise ValueError(f'no conversations carry source {source!r}') |
| if len(indices) <= limit: |
| print(f'cap {source}={limit}: {len(indices)} present, nothing dropped') |
| continue |
| dropped |= set(indices) - set(rng.sample(indices, limit)) |
| print(f'cap {source}={limit}: dropped {len(indices) - limit:,} of {len(indices):,}') |
| return [c for index, c in enumerate(conversations) if index not in dropped] |
|
|
|
|
| def _parse_caps(pairs: list[str]) -> dict[str, int]: |
| caps = {} |
| for pair in pairs: |
| source, _, count = pair.partition('=') |
| if not count.isdigit(): |
| raise ValueError(f'--cap expects SOURCE=N, got {pair!r}') |
| caps[source] = int(count) |
| return caps |
|
|
|
|
| def prepare(args: argparse.Namespace) -> None: |
| import numpy as np |
|
|
| from diffusion_lm.reasoning import ExampleEncoder, LayoutSpec, _write_packed, size_token_ids |
| 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) |
|
|
| conversations = _apply_caps(load(args.inputs), _parse_caps(args.cap or []), args.seed) |
| documents = sorted({_document(c, i) for i, c in enumerate(conversations)}) |
| rng = random.Random(args.seed) |
| rng.shuffle(documents) |
| held = set(documents[:max(1, round(len(documents) * args.val_fraction))]) |
|
|
| split: dict[str, list[tuple]] = {'train': [], 'validation': []} |
| dropped = 0 |
| blocks: Counter[int] = Counter() |
| empty_chains = 0 |
| for index, conversation in enumerate(conversations): |
| bucket = 'validation' if _document(conversation, index) in held else 'train' |
| for example in to_examples(conversation, keep_messages=args.keep_messages): |
| encoded = encoder.encode_adaptive(example) |
| if encoded is None: |
| dropped += 1 |
| continue |
| blocks.update(encoded.block_sizes) |
| empty_chains += not encoded.block_sizes |
| split[bucket].append((encoded.tokens, encoded.regions)) |
|
|
| args.output_dir.mkdir(parents=True, exist_ok=True) |
| 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), |
| |
| |
| 'size_token_ids': size_token_ids(tokenizer, spec.sizes), |
| 'source': 'general-chat', |
| }, |
| ) |
| print(f'{name}: {len(rows):,} examples -> {args.output_dir}') |
| total = sum(len(rows) for rows in split.values()) |
| print(f'{len(conversations):,} conversations, {len(documents):,} documents, ' |
| f'{dropped:,} dropped at encode ({dropped / max(1, dropped + total):.1%})') |
| print(f'examples answering with no thinking block: {empty_chains:,} ' |
| f'({empty_chains / max(1, total):.1%})') |
| print('block sizes: ' + ', '.join(f'{size}:{count:,}' for size, count in sorted(blocks.items()))) |
|
|
|
|
| def main() -> None: |
| parser = argparse.ArgumentParser(description=__doc__) |
| sub = parser.add_subparsers(dest='command', required=True) |
|
|
| prep = sub.add_parser('prepare', help='pack the corpus into the adaptive layout') |
| prep.add_argument('--inputs', type=Path, nargs='+', required=True) |
| prep.add_argument('--tokenizer', type=Path, |
| default=Path('artifacts/tokenizer-qwen3-adaptive.json')) |
| prep.add_argument('--output-dir', type=Path, required=True) |
| prep.add_argument('--seq-len', type=int, default=2048) |
| prep.add_argument('--sizes', type=int, nargs='+', default=[32, 64, 128]) |
| prep.add_argument('--max-slots', type=int, default=40) |
| prep.add_argument('--keep-messages', type=int, default=KEEP_MESSAGES) |
| prep.add_argument('--cap', nargs='*', metavar='SOURCE=N', |
| help='keep at most N conversations from a source, e.g. ground-coqa=2273; ' |
| 'weight in the mix is examples, and sources differ ~10x in examples ' |
| 'per conversation') |
| prep.add_argument('--val-fraction', type=float, default=0.02) |
| prep.add_argument('--seed', type=int, default=1337) |
|
|
| args = parser.parse_args() |
| prepare(args) |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|