"""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), # reasoning_train resolves the adaptive control ids from the pack, not the # tokenizer, and refuses a pack without them. '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()