marimo-0.6b-mlx / chat.py
goldenfox's picture
Initial release: standalone MLX port, gate passed (top-1 fp16 99.9696% / 13156 rows)
aef8188 verified
Raw
History Blame Contribute Delete
7.3 kB
"""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} <sz{size}> 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()