"""Interactive terminal chat over the MLX adaptive-hybrid engine. The conversation is rendered exactly as training saw it: a ChatML prefix whose system turn carries the ledger of notes taken on messages that fell out of the visible window. The ledger-merging rules are a faithful port of ``diffusion_lm.claims``; the parity harness in the training repo compares both renderings byte for byte before each release. """ from __future__ import annotations import argparse import re import sys import time from pathlib import Path import numpy as np from tokenizers import Tokenizer from model import load_denoiser from sampler import stream_turn IM_START = '<|im_start|>' IM_END = '<|im_end|>' _NOTE_FACT = re.compile(r'^([^:]{1,48}): (.+)$') def chatml_turn(role: str, content: str) -> str: return f'{IM_START}{role}\n{content}{IM_END}\n' def ledger_line(entries: list[str]) -> str: return 'Known so far: ' + '; '.join(entries) + '.' if entries else '' 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. Fragments that do not parse as ``key: value`` pass through in order, deduplicated verbatim. """ 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.""" 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, extra: str = '') -> str: """Conversation prefix ending right after the assistant header, ledger in the system turn.""" 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 build_prefix(messages: list[dict[str, str]], *, keep: int, system: str) -> str: older = merge_notes(ledger_notes(messages, keep)) window = messages[window_start(messages, keep):] return chat_prefix( [{'role': m['role'], 'content': m['content']} for m in window], system=system, extra=ledger_line(older), ) def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument('--model-dir', type=Path, default=Path(__file__).parent) parser.add_argument('--temperature', type=float, default=0.8) parser.add_argument('--top-p', type=float, default=0.95) parser.add_argument('--repetition-penalty', type=float, default=1.0) parser.add_argument('--seed', type=int, default=0, help='0 draws a fresh seed per turn') parser.add_argument('--steps-per-block', type=int, default=None) parser.add_argument('--max-answer-tokens', type=int, default=None) parser.add_argument('--keep-messages', type=int, default=None) parser.add_argument('--system', default=None) parser.add_argument('--show-thinking', action='store_true') parser.add_argument('--no-cache-blocks', action='store_true', help='recompute the full window on every denoising step') args = parser.parse_args() denoiser, config = load_denoiser(args.model_dir) mdlm = config['mdlm'] tokenizer = Tokenizer.from_file(str(args.model_dir / 'tokenizer.json')) keep = args.keep_messages if args.keep_messages is not None else mdlm['keep_messages'] system = args.system if args.system is not None else mdlm['system_prompt'] messages: list[dict[str, str]] = [] print(f'{config["model_type"]} · step {mdlm["checkpoint_step"]} · ' f'window {keep} messages · /reset clears, /exit quits') while True: try: user_text = input('you> ').strip() except (EOFError, KeyboardInterrupt): print() break if not user_text: continue if user_text == '/exit': break if user_text == '/reset': messages.clear() print('(history cleared)') continue messages.append({'role': 'user', 'content': user_text}) prefix = build_prefix(messages, keep=keep, system=system) prompt_ids = tokenizer.encode(prefix, add_special_tokens=False).ids if len(prompt_ids) >= mdlm['max_seq_len'] - 48: print(f'(context full: {len(prompt_ids)} tokens — lower --keep-messages or /reset)') messages.pop() continue rng = np.random.default_rng(args.seed if args.seed else time.time_ns() % 2**31) answer_text = '' answer_ids: list[int] = [] result: dict = {} for event in stream_turn( denoiser, prompt_ids, rng=rng, temperature=args.temperature, top_p=args.top_p, repetition_penalty=args.repetition_penalty, steps_per_block=args.steps_per_block, max_answer_tokens=args.max_answer_tokens, cache_blocks=not args.no_cache_blocks, ): if event[0] == 'block': if sys.stdout.isatty(): _, index, size, step, total, _ = event print(f'\r(thinking · block {index + 1} step {step}/{total})', end='', flush=True) elif event[0] == 'answer': if not answer_ids and sys.stdout.isatty(): print('\r\x1b[2K', end='') answer_ids.append(event[1]) decoded = tokenizer.decode(answer_ids, skip_special_tokens=True) print(decoded[len(answer_text):], end='', flush=True) answer_text = decoded else: result = event[1] print() decoded_blocks = [ (size, tokenizer.decode(ids, skip_special_tokens=True).strip()) for size, ids in result.get('blocks', []) ] note = '; '.join(text for _, text in decoded_blocks if text) messages.append({'role': 'assistant', 'content': answer_text.strip(), 'note': note}) if args.show_thinking: for size, text in decoded_blocks: if text: print(f' [sz{size}] {text}') if result: rate = len(answer_ids) / max(result['answer_seconds'], 1e-9) print(f' ({result["think_seconds"]:.1f}s think · ' f'{result["answer_seconds"]:.1f}s answer · {rate:.0f} tok/s)') if __name__ == '__main__': main()