marimo-diffusion / src /diffusion_lm /chatcorpus.py
goldenfox's picture
Marimo Diffusion 0.6B: checkpoint, sampler, OpenAI server, ledger-needle bench
685e018 verified
Raw
History Blame Contribute Delete
9.99 kB
"""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()