Download src/diffusion_lm/reasoning_playground.py from goldenfox/marimo-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 56.1 kB
-
https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/reasoning_playground.py
- Command line
-
hf download hf://goldenfox/marimo-diffusion/src/diffusion_lm/reasoning_playground.py
-
curl -L -o reasoning_playground.py https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/reasoning_playground.py
56.1 kB
| """Interactive playground for the reasoning variants. | |
| Loads any ``mini-diffusion-lm-reasoning-inference-v1`` checkpoint and solves | |
| user-provided problems with the sampler matching its training objective, | |
| streaming thought slots as they are denoised and answer tokens as they are | |
| decoded, together with per-phase wall-clock metrics. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import html | |
| import json | |
| import time | |
| from pathlib import Path | |
| import torch | |
| from diffusion_lm.claims import ( | |
| IM_END, | |
| SYSTEM, | |
| chat_prefix, | |
| ledger_line, | |
| ledger_notes, | |
| merge_notes, | |
| window_start, | |
| ) | |
| from diffusion_lm.config import ModelConfig | |
| from diffusion_lm.diffusion import iterative_unmask_steps, _sample_categorical | |
| from diffusion_lm.flexattn import _FLEX_EAGER | |
| from diffusion_lm.hybrid import ( | |
| _apply_repetition_penalty, | |
| _apply_top_p, | |
| _ar_predict_control, | |
| _kv_cache_enabled, | |
| adaptive_block_mask, | |
| block_mask_from_boundaries, | |
| block_size_curriculum, | |
| prefix_causal_blocked, | |
| slot_causal_blocked, | |
| ) | |
| from diffusion_lm.model import build_denoiser | |
| from diffusion_lm.reasoning import extract_boxed_answer, reasoning_token_ids, size_token_ids | |
| from diffusion_lm.tokenizer import load_tokenizer, special_token_ids | |
| from diffusion_lm.train import resolve_device | |
| def _discover_checkpoints(root: Path) -> dict[str, Path]: | |
| found: dict[str, Path] = {} | |
| for path in sorted(root.glob('*/inference-latest.pt')): | |
| found[path.parent.name] = path | |
| return found | |
| class ReasoningEngine: | |
| """Load one reasoning checkpoint and stream problem solutions.""" | |
| def __init__(self, checkpoint_path: Path, tokenizer_path: Path, device: torch.device) -> None: | |
| payload = torch.load(checkpoint_path, map_location='cpu', weights_only=False) | |
| if payload.get('format') != 'mini-diffusion-lm-reasoning-inference-v1': | |
| raise ValueError(f'unsupported checkpoint format in {checkpoint_path}') | |
| self.config = payload['config'] | |
| self.objective = self.config['reasoning']['objective'] | |
| model_config = ModelConfig(**self.config['model']) | |
| self.model = build_denoiser(model_config, load_pretrained=False) | |
| self.model.load_state_dict(payload['model']) | |
| dtype = ( | |
| torch.bfloat16 | |
| if model_config.backbone != 'project' and device.type == 'cuda' | |
| else torch.float32 | |
| ) | |
| self.model.to(device=device, dtype=dtype).eval() | |
| self.device = device | |
| self.step = payload.get('step') | |
| self.tokenizer = load_tokenizer(tokenizer_path) | |
| self.roles = special_token_ids(self.tokenizer) | |
| self.reasoning_ids = reasoning_token_ids(self.tokenizer) | |
| self.adaptive = bool(self.config['reasoning'].get('adaptive')) | |
| self.causal_prefix = bool(self.config['reasoning'].get('causal_prefix', False)) | |
| trained_sizes = tuple(self.config['reasoning'].get('sizes') or ()) or None | |
| self.size_ids = size_token_ids(self.tokenizer, trained_sizes) if self.adaptive else {} | |
| self.im_end_id = self.tokenizer.token_to_id(IM_END) | |
| self.chat_ready = ( | |
| self.objective == 'hybrid' and self.adaptive and self.im_end_id is not None | |
| ) | |
| self.block = 32 | |
| self.max_slots = 40 | |
| self.max_blocks = 48 | |
| self.log_lines: list[str] = [] | |
| def _decode(self, ids: list[int]) -> str: | |
| return self.tokenizer.decode(ids, skip_special_tokens=True) | |
| def _slot_texts(self, think_ids: list[int]) -> list[str]: | |
| end_id = self.reasoning_ids['end_think'] | |
| pad_id = self.reasoning_ids['thought_pad'] | |
| slots = [ | |
| think_ids[start : start + self.block] | |
| for start in range(0, len(think_ids), self.block) | |
| ] | |
| texts = [] | |
| for slot in slots: | |
| content = [t for t in slot if t not in (pad_id, end_id)] | |
| texts.append(self._decode(content).strip()) | |
| return [text for text in texts if text] | |
| def _block_contents(self, think_ids: list[int]) -> list[tuple[int, str]]: | |
| """Decode variable-size blocks by walking the ``<szN>`` tokens in the stream.""" | |
| size_by_id = {token_id: size for size, token_id in self.size_ids.items()} | |
| end_id = self.reasoning_ids['end_think'] | |
| pad_id = self.reasoning_ids['thought_pad'] | |
| blocks = [] | |
| cut = 0 | |
| while cut < len(think_ids): | |
| size = size_by_id.get(think_ids[cut]) | |
| if size is None: | |
| cut += 1 | |
| continue | |
| block = think_ids[cut + 1 : cut + 1 + size] | |
| content = [t for t in block if t not in (pad_id, end_id)] | |
| blocks.append((size, self._decode(content).strip())) | |
| cut += 1 + size | |
| return blocks | |
| def _adaptive_block_texts(self, think_ids: list[int]) -> list[str]: | |
| return [ | |
| f'({size}) {text}'.strip() for size, text in self._block_contents(think_ids) | |
| ] | |
| def _token_glyph(self, token_id: int) -> str: | |
| """Render one in-flight token: masks and structural tokens get visible glyphs.""" | |
| if token_id == self.model.config.mask_token_id: | |
| return '▒' | |
| if token_id == self.reasoning_ids['thought_pad']: | |
| return '·' | |
| if token_id == self.reasoning_ids['end_think']: | |
| return ' ⏹' | |
| return self.tokenizer.decode([token_id], skip_special_tokens=False) | |
| def _active_slot_text(self, slot_tokens: list[int]) -> str: | |
| return ''.join(self._token_glyph(token) for token in slot_tokens) | |
| def stream_solve( | |
| self, | |
| problem: str, | |
| *, | |
| temperature: float, | |
| steps_per_block: int, | |
| diffusion_steps: int, | |
| max_answer_tokens: int = 200, | |
| repetition_penalty: float = 1.0, | |
| top_p: float = 1.0, | |
| seed: int = 0, | |
| strategy: str = 'ancestral', | |
| step_delay: float = 0.0, | |
| ): | |
| """Yield ``(thoughts, answer, status)`` snapshots while solving. | |
| ``step_delay`` throttles each denoising step (and each AR token) so the | |
| reveal order is watchable in real time. | |
| """ | |
| generator = torch.Generator(device=self.device.type) | |
| generator.manual_seed(seed if seed else int(time.time_ns() % 2**31)) | |
| prompt_ids = self.tokenizer.encode(problem.strip(), add_special_tokens=False).ids | |
| mask_id = self.model.config.mask_token_id | |
| think_id = self.reasoning_ids['think'] | |
| end_think_id = self.reasoning_ids['end_think'] | |
| eos_id = self.roles['eos'] | |
| if self.objective == 'hybrid' and self.adaptive: | |
| yield from self._stream_adaptive( | |
| prompt_ids, | |
| stop_ids=(eos_id,), | |
| temperature=temperature, | |
| steps_per_block=steps_per_block, | |
| max_answer_tokens=max_answer_tokens, | |
| repetition_penalty=repetition_penalty, | |
| top_p=top_p, | |
| strategy=strategy, | |
| step_delay=step_delay, | |
| generator=generator, | |
| ) | |
| return | |
| if self.objective == 'block_diffusion': | |
| # Plain-text continuation under the conversion geometry: each new block is | |
| # fully masked and denoised bidirectionally over its causal past, with block | |
| # sizes drawn from the same curriculum weights the checkpoint trained under. | |
| reasoning_cfg = self.config['reasoning'] | |
| sizes = tuple(reasoning_cfg.get('sizes') or (self.block,)) | |
| weights = block_size_curriculum( | |
| self.step, | |
| n_sizes=len(sizes), | |
| curriculum_steps=int(reasoning_cfg.get('curriculum_steps') or 0), | |
| ).to(self.device) | |
| size_tensor = torch.tensor(sizes, device=self.device) | |
| budget = min(384, self.model.config.max_seq_len - len(prompt_ids) - 1) | |
| self.log_lines = [ | |
| f'prompt_tokens={len(prompt_ids)} temp={temperature} ' | |
| f'steps/block={steps_per_block} sizes={list(sizes)} ' | |
| f'weights={[round(w, 3) for w in weights.tolist()]} ' | |
| f'strategy={strategy} flex_eager={_FLEX_EAGER}' | |
| ] | |
| sequence = list(prompt_ids) | |
| prompt_len = len(prompt_ids) | |
| boundary_starts: list[int] = [] | |
| block_texts: list[str] = [] | |
| generated: list[int] = [] | |
| compute = 0.0 | |
| finished = False | |
| while not finished and len(generated) < budget: | |
| drawn = int( | |
| size_tensor[torch.multinomial(weights, 1, generator=generator)].item() | |
| ) | |
| block = min(drawn, budget - len(generated)) | |
| prefix_len = len(sequence) | |
| seq_len = prefix_len + block | |
| boundary_starts.append(prefix_len) | |
| boundary = torch.zeros(1, seq_len, dtype=torch.bool, device=self.device) | |
| boundary[0, 0] = True | |
| for start in boundary_starts: | |
| boundary[0, start] = True | |
| blocked = block_mask_from_boundaries( | |
| boundary, | |
| torch.tensor([prompt_len], device=self.device), | |
| torch.tensor([seq_len], device=self.device), | |
| causal_prefix=self.causal_prefix, | |
| ) | |
| current = torch.tensor( | |
| [sequence + [mask_id] * block], dtype=torch.long, device=self.device | |
| ) | |
| block_ids: list[int] = [] | |
| tick = time.perf_counter() | |
| for state in iterative_unmask_steps( | |
| self.model, | |
| current, | |
| mask_id, | |
| steps=steps_per_block, | |
| temperature=temperature, | |
| strategy=strategy, | |
| blocked_token_ids=(self.roles['pad'],), | |
| attn_mask=blocked, | |
| generator=generator, | |
| ): | |
| compute += time.perf_counter() - tick | |
| block_ids = [int(t) for t in state.tokens[0, prefix_len:]] | |
| if step_delay or state.masked_remaining == 0: | |
| yield ( | |
| [*block_texts, self._active_slot_text(block_ids)], | |
| self._decode(generated), | |
| f'bloque {len(block_texts) + 1} ({block} tok) · ' | |
| f'denoising {state.step}/{state.total_steps}', | |
| ) | |
| if step_delay and state.masked_remaining: | |
| time.sleep(step_delay) | |
| tick = time.perf_counter() | |
| if eos_id in block_ids: | |
| block_ids = block_ids[: block_ids.index(eos_id)] | |
| finished = True | |
| sequence.extend(block_ids) | |
| generated.extend(block_ids) | |
| block_texts.append(f'({block}) {self._decode(block_ids).strip()}'.strip()) | |
| self.log_lines.append( | |
| f'block {len(block_texts)}: size={block} · {compute:.2f}s cum' | |
| ) | |
| rate = len(generated) / compute if compute > 0 else 0.0 | |
| self.log_lines.append( | |
| f'TOTAL: {len(generated)} tok · blocks={len(block_texts)} · ' | |
| f'{compute:.2f}s ({rate:.0f} tok/s)' | |
| ) | |
| yield block_texts, self._decode(generated), ( | |
| f'listo · {len(block_texts)} bloques · {len(generated)} tok en ' | |
| f'{compute:.2f}s ({rate:.0f} tok/s) ' | |
| f'(checkpoint CPT: continúa texto, no piensa)' | |
| ) | |
| return | |
| if self.objective == 'hybrid': | |
| self.log_lines = [ | |
| f'prompt_tokens={len(prompt_ids)} temp={temperature} ' | |
| f'steps/slot={steps_per_block} rep_penalty={repetition_penalty} ' | |
| f'top_p={top_p} max_slots={self.max_slots} strategy={strategy} ' | |
| f'flex_eager={_FLEX_EAGER}' | |
| ] | |
| sequence = [*prompt_ids, think_id] | |
| think_ids: list[int] = [] | |
| think_compute = 0.0 | |
| slot_budget = self.model.config.max_seq_len - self.block - 8 | |
| problem_tensor = torch.tensor([len(prompt_ids)], device=self.device) | |
| for slot_index in range(self.max_slots): | |
| if len(sequence) > slot_budget: | |
| break | |
| slot_wall_start = time.perf_counter() | |
| window = torch.tensor( | |
| [sequence + [mask_id] * self.block], dtype=torch.long, device=self.device | |
| ) | |
| blocked = slot_causal_blocked( | |
| problem_tensor, | |
| torch.tensor([slot_index + 1], device=self.device), | |
| self.block, | |
| window.shape[1], | |
| ) | |
| slot: list[int] = [] | |
| tick = time.perf_counter() | |
| for state in iterative_unmask_steps( | |
| self.model, | |
| window, | |
| mask_id, | |
| steps=steps_per_block, | |
| temperature=temperature, | |
| strategy=strategy, | |
| attn_mask=blocked, | |
| generator=generator, | |
| ): | |
| think_compute += time.perf_counter() - tick | |
| slot = [int(t) for t in state.tokens[0, len(sequence):]] | |
| revealed = len(think_ids) + self.block - state.masked_remaining | |
| rate = revealed / think_compute if think_compute > 0 else 0.0 | |
| yield ( | |
| self._slot_texts(think_ids) | |
| + [self._active_slot_text(slot)], | |
| '', | |
| f'denoising slot {slot_index + 1} · paso ' | |
| f'{state.step}/{state.total_steps} · {rate:.0f} tok/s', | |
| ) | |
| if step_delay and state.step < state.total_steps: | |
| time.sleep(step_delay) | |
| tick = time.perf_counter() | |
| sequence.extend(slot) | |
| think_ids.extend(slot) | |
| slot_content = [t for t in slot if t not in ( | |
| self.reasoning_ids['thought_pad'], end_think_id | |
| )] | |
| self.log_lines.append( | |
| f' T{slot_index + 1}: {len(slot_content)} tok · ' | |
| f'{time.perf_counter() - slot_wall_start:.2f}s wall' | |
| + (' · </think>' if end_think_id in slot else '') | |
| ) | |
| if end_think_id in slot: | |
| break | |
| else: | |
| sequence.extend( | |
| [end_think_id] + [self.reasoning_ids['thought_pad']] * (self.block - 1) | |
| ) | |
| answer_ids: list[int] = [] | |
| prefix_len = len(sequence) | |
| answer_compute = 0.0 | |
| tick = time.perf_counter() | |
| for _ in range(max_answer_tokens): | |
| current = torch.tensor( | |
| [sequence + answer_ids], dtype=torch.long, device=self.device | |
| ) | |
| blocked = prefix_causal_blocked( | |
| torch.tensor([prefix_len], device=self.device), current.shape[1] | |
| ) | |
| output_positions = torch.zeros_like(current, dtype=torch.bool) | |
| output_positions[0, -1] = True | |
| logits = self.model( | |
| current, output_positions=output_positions, attn_mask=blocked | |
| ) | |
| logits = _apply_repetition_penalty(logits, answer_ids, repetition_penalty) | |
| logits = _apply_top_p(logits, top_p) | |
| token, _ = _sample_categorical(logits, temperature, generator) | |
| token_id = int(token.item()) | |
| answer_ids.append(token_id) | |
| answer_compute += time.perf_counter() - tick | |
| if step_delay or len(answer_ids) % 4 == 0 or token_id == eos_id: | |
| rate = len(answer_ids) / answer_compute if answer_compute > 0 else 0.0 | |
| yield ( | |
| self._slot_texts(think_ids), | |
| self._decode(answer_ids), | |
| f'respondiendo (AR, token {len(answer_ids)}) · {rate:.0f} tok/s', | |
| ) | |
| if step_delay and token_id != eos_id: | |
| time.sleep(step_delay / 3) | |
| if token_id == eos_id: | |
| break | |
| tick = time.perf_counter() | |
| think_tokens = len(think_ids) | |
| answer_tokens = len(answer_ids) | |
| total_compute = think_compute + answer_compute | |
| think_rate = think_tokens / think_compute if think_compute > 0 else 0.0 | |
| answer_rate = answer_tokens / answer_compute if answer_compute > 0 else 0.0 | |
| total_rate = ( | |
| (think_tokens + answer_tokens) / total_compute if total_compute > 0 else 0.0 | |
| ) | |
| status = ( | |
| f'listo · pensar: {think_tokens} tok en {think_compute:.2f}s ' | |
| f'({think_rate:.0f} tok/s) · responder: {answer_tokens} tok en ' | |
| f'{answer_compute:.2f}s ({answer_rate:.0f} tok/s) · total: ' | |
| f'{think_tokens + answer_tokens} tok en {total_compute:.2f}s ' | |
| f'({total_rate:.0f} tok/s)' | |
| ) | |
| self.log_lines.append( | |
| f' answer: {answer_tokens} tok · {answer_compute:.2f}s wall ' | |
| f'({answer_rate:.0f} tok/s) · terminated={end_think_id in think_ids}' | |
| ) | |
| self.log_lines.append( | |
| f'TOTAL: {think_tokens + answer_tokens} tok · {total_compute:.2f}s ' | |
| f'({total_rate:.0f} tok/s)' | |
| ) | |
| yield self._slot_texts(think_ids), self._decode(answer_ids), status | |
| return | |
| if self.objective == 'lm': | |
| # Plain causal continuation for pretraining checkpoints: no think token, | |
| # no masking; the prompt is simply extended left to right. | |
| sequence = list(prompt_ids) | |
| generated: list[int] = [] | |
| prefix_len = len(sequence) | |
| budget = self.model.config.max_seq_len - prefix_len - 1 | |
| compute = 0.0 | |
| tick = time.perf_counter() | |
| for _ in range(min(budget, 384)): | |
| current = torch.tensor( | |
| [sequence + generated], dtype=torch.long, device=self.device | |
| ) | |
| blocked = prefix_causal_blocked( | |
| torch.tensor([prefix_len], device=self.device), current.shape[1] | |
| ) | |
| output_positions = torch.zeros_like(current, dtype=torch.bool) | |
| output_positions[0, -1] = True | |
| logits = self.model( | |
| current, output_positions=output_positions, attn_mask=blocked | |
| ) | |
| token, _ = _sample_categorical(logits, temperature, generator) | |
| token_id = int(token.item()) | |
| generated.append(token_id) | |
| compute += time.perf_counter() - tick | |
| if step_delay or token_id == eos_id or len(generated) % 8 == 0: | |
| rate = len(generated) / compute if compute > 0 else 0.0 | |
| yield ( | |
| [], | |
| self._decode(generated), | |
| f'continuando (token {len(generated)}) · {rate:.0f} tok/s', | |
| ) | |
| if step_delay and token_id != eos_id: | |
| time.sleep(step_delay / 3) | |
| if token_id == eos_id: | |
| break | |
| tick = time.perf_counter() | |
| rate = len(generated) / compute if compute > 0 else 0.0 | |
| yield [], self._decode(generated), ( | |
| f'listo · total: {len(generated)} tok en {compute:.2f}s ({rate:.0f} tok/s) ' | |
| f'(checkpoint base: continua texto, no piensa)' | |
| ) | |
| return | |
| if self.objective == 'ar': | |
| sequence = [*prompt_ids, think_id] | |
| generated: list[int] = [] | |
| prefix_len = len(sequence) | |
| budget = self.model.config.max_seq_len - prefix_len - 1 | |
| compute = 0.0 | |
| tick = time.perf_counter() | |
| for _ in range(min(budget, 384)): | |
| current = torch.tensor( | |
| [sequence + generated], dtype=torch.long, device=self.device | |
| ) | |
| blocked = prefix_causal_blocked( | |
| torch.tensor([prefix_len], device=self.device), current.shape[1] | |
| ) | |
| output_positions = torch.zeros_like(current, dtype=torch.bool) | |
| output_positions[0, -1] = True | |
| logits = self.model( | |
| current, output_positions=output_positions, attn_mask=blocked | |
| ) | |
| token, _ = _sample_categorical(logits, temperature, generator) | |
| token_id = int(token.item()) | |
| generated.append(token_id) | |
| compute += time.perf_counter() - tick | |
| if step_delay or token_id == eos_id or len(generated) % 8 == 0: | |
| thoughts, answer = self._split_flat(generated) | |
| rate = len(generated) / compute if compute > 0 else 0.0 | |
| yield thoughts, answer, ( | |
| f'generando (token {len(generated)}) · {rate:.0f} tok/s' | |
| ) | |
| if step_delay and token_id != eos_id: | |
| time.sleep(step_delay / 3) | |
| if token_id == eos_id: | |
| break | |
| tick = time.perf_counter() | |
| thoughts, answer = self._split_flat(generated) | |
| rate = len(generated) / compute if compute > 0 else 0.0 | |
| yield thoughts, answer, ( | |
| f'listo · total: {len(generated)} tok en {compute:.2f}s ({rate:.0f} tok/s)' | |
| ) | |
| return | |
| budget = min(352, self.model.config.max_seq_len - len(prompt_ids) - 1) | |
| sequence_tensor = torch.tensor( | |
| [[*prompt_ids, think_id] + [mask_id] * budget], dtype=torch.long, device=self.device | |
| ) | |
| final_ids: list[int] = [] | |
| compute = 0.0 | |
| tick = time.perf_counter() | |
| for state in iterative_unmask_steps( | |
| self.model, | |
| sequence_tensor, | |
| mask_id, | |
| steps=diffusion_steps, | |
| temperature=temperature, | |
| strategy=strategy, | |
| blocked_token_ids=(self.roles['pad'],), | |
| generator=generator, | |
| ): | |
| compute += time.perf_counter() - tick | |
| final_ids = [int(t) for t in state.tokens[0, len(prompt_ids) + 1:]] | |
| if step_delay or state.step % 8 == 0 or state.masked_remaining == 0: | |
| if step_delay: | |
| thoughts = [self._active_slot_text(final_ids)] | |
| answer = '' | |
| else: | |
| visible = [t for t in final_ids if t != mask_id] | |
| thoughts, answer = self._split_flat(visible) | |
| revealed = len(final_ids) - state.masked_remaining | |
| rate = revealed / compute if compute > 0 else 0.0 | |
| yield ( | |
| thoughts, | |
| answer, | |
| f'denoising {state.step}/{state.total_steps} · {rate:.0f} tok/s', | |
| ) | |
| if step_delay and state.masked_remaining: | |
| time.sleep(step_delay) | |
| tick = time.perf_counter() | |
| if eos_id in final_ids: | |
| final_ids = final_ids[: final_ids.index(eos_id) + 1] | |
| thoughts, answer = self._split_flat(final_ids) | |
| revealed = len(final_ids) | |
| rate = revealed / compute if compute > 0 else 0.0 | |
| yield thoughts, answer, ( | |
| f'listo · total: {revealed} tok en {compute:.2f}s ({rate:.0f} tok/s)' | |
| ) | |
| def stream_chat( | |
| self, | |
| prefix_text: str, | |
| *, | |
| temperature: float, | |
| steps_per_block: int, | |
| max_answer_tokens: int = 224, | |
| repetition_penalty: float = 1.0, | |
| top_p: float = 1.0, | |
| seed: int = 0, | |
| strategy: str = 'ancestral', | |
| step_delay: float = 0.0, | |
| blocks_out: list[tuple[int, str]] | None = None, | |
| ): | |
| """Yield ``(thoughts, answer, status)`` snapshots for one ChatML assistant turn. | |
| ``prefix_text`` is the rendered conversation ending right after the assistant | |
| header, encoded verbatim: training saw the trailing newline, so no stripping. | |
| The answer stops at ``<|im_end|>`` as well as the eos token. ``blocks_out`` | |
| receives the turn's ``(size, content)`` blocks; a per-call sink rather than | |
| engine state, so concurrent runs on the same engine cannot leak into each other. | |
| """ | |
| if not self.chat_ready: | |
| raise ValueError('checkpoint is not an adaptive hybrid over a ChatML tokenizer') | |
| generator = torch.Generator(device=self.device.type) | |
| generator.manual_seed(seed if seed else int(time.time_ns() % 2**31)) | |
| prompt_ids = self.tokenizer.encode(prefix_text, add_special_tokens=False).ids | |
| if len(prompt_ids) >= self.model.config.max_seq_len - 48: | |
| message = ( | |
| f'contexto lleno: prefijo de {len(prompt_ids)} tokens con ventana de ' | |
| f'{self.model.config.max_seq_len} — bajá los mensajes visibles o reiniciá' | |
| ) | |
| self.log_lines = [message] | |
| yield [], '', message | |
| return | |
| yield from self._stream_adaptive( | |
| prompt_ids, | |
| stop_ids=(self.im_end_id, self.roles['eos']), | |
| temperature=temperature, | |
| steps_per_block=steps_per_block, | |
| max_answer_tokens=max_answer_tokens, | |
| repetition_penalty=repetition_penalty, | |
| top_p=top_p, | |
| strategy=strategy, | |
| step_delay=step_delay, | |
| generator=generator, | |
| blocks_out=blocks_out, | |
| ) | |
| def _stream_adaptive( | |
| self, | |
| prompt_ids: list[int], | |
| *, | |
| stop_ids: tuple[int, ...], | |
| temperature: float, | |
| steps_per_block: int, | |
| max_answer_tokens: int, | |
| repetition_penalty: float, | |
| top_p: float, | |
| strategy: str, | |
| step_delay: float, | |
| generator: torch.Generator, | |
| blocks_out: list[tuple[int, str]] | None = None, | |
| ): | |
| """Stream the adaptive control/denoise/answer loop shared by solve and chat. | |
| The answer phase reuses the prefix through the model's key/value cache when the | |
| backbone provides one and ``MDLM_KV_CACHE`` selects the ar part, matching the | |
| batch decoder in :mod:`diffusion_lm.hybrid`. | |
| """ | |
| mask_id = self.model.config.mask_token_id | |
| think_id = self.reasoning_ids['think'] | |
| end_think_id = self.reasoning_ids['end_think'] | |
| size_by_id = {token_id: size for size, token_id in self.size_ids.items()} | |
| size_ids_tensor = torch.tensor(sorted(self.size_ids.values()), device=self.device) | |
| control_ids = [*size_by_id.keys(), end_think_id] | |
| cached = hasattr(self.model, 'forward_cached') and _kv_cache_enabled('ar') | |
| self.log_lines = [ | |
| f'prompt_tokens={len(prompt_ids)} temp={temperature} ' | |
| f'steps/block={steps_per_block} rep_penalty={repetition_penalty} ' | |
| f'top_p={top_p} max_blocks={self.max_blocks} ' | |
| f'sizes={sorted(size_by_id.values())} strategy={strategy} ' | |
| f'flex_eager={_FLEX_EAGER} kv_cache={"ar" if cached else "off"}' | |
| ] | |
| sequence = [*prompt_ids, think_id] | |
| prefix_len = len(prompt_ids) + 1 | |
| problem_tensor = torch.tensor([len(prompt_ids)], device=self.device) | |
| think_compute = 0.0 | |
| terminated = False | |
| for block_index in range(self.max_blocks): | |
| tick = time.perf_counter() | |
| control = _ar_predict_control( | |
| self.model, | |
| sequence, | |
| control_ids, | |
| prefix_len=prefix_len, | |
| temperature=0.0, | |
| device=self.device, | |
| generator=generator, | |
| causal_prefix=self.causal_prefix, | |
| ) | |
| think_compute += time.perf_counter() - tick | |
| if control == end_think_id: | |
| terminated = True | |
| self.log_lines.append(f' control {block_index + 1}: </think>') | |
| break | |
| size = size_by_id[control] | |
| if len(sequence) + 1 + size > self.model.config.max_seq_len - 8: | |
| self.log_lines.append( | |
| f' control {block_index + 1}: <sz{size}> · sin contexto, corto acá' | |
| ) | |
| break | |
| sequence.append(control) | |
| window_prefix = len(sequence) | |
| block_wall_start = time.perf_counter() | |
| window = torch.tensor( | |
| [sequence + [mask_id] * size], dtype=torch.long, device=self.device | |
| ) | |
| blocked = adaptive_block_mask( | |
| window, | |
| problem_tensor, | |
| torch.tensor([window.shape[1]], device=self.device), | |
| size_ids_tensor, | |
| end_think_id, | |
| causal_prefix=self.causal_prefix, | |
| ) | |
| block: list[int] = [] | |
| tick = time.perf_counter() | |
| for state in iterative_unmask_steps( | |
| self.model, | |
| window, | |
| mask_id, | |
| steps=steps_per_block, | |
| temperature=temperature, | |
| strategy=strategy, | |
| attn_mask=blocked, | |
| generator=generator, | |
| ): | |
| think_compute += time.perf_counter() - tick | |
| block = [int(t) for t in state.tokens[0, window_prefix:]] | |
| yield ( | |
| self._adaptive_block_texts(sequence[prefix_len:]) | |
| + [f'({size}) ' + self._active_slot_text(block)], | |
| '', | |
| f'denoising bloque {block_index + 1} (tamaño {size}) · paso ' | |
| f'{state.step}/{state.total_steps}', | |
| ) | |
| if step_delay and state.step < state.total_steps: | |
| time.sleep(step_delay) | |
| tick = time.perf_counter() | |
| sequence.extend(block) | |
| self.log_lines.append( | |
| f' B{block_index + 1}: <sz{size}> · ' | |
| f'{time.perf_counter() - block_wall_start:.2f}s wall' | |
| ) | |
| sequence.append(end_think_id) | |
| think_ids = sequence[prefix_len:] | |
| if blocks_out is not None: | |
| blocks_out.extend(self._block_contents(think_ids)) | |
| answer_ids: list[int] = [] | |
| answer_compute = 0.0 | |
| prefix_tensor = torch.tensor([prefix_len], device=self.device) | |
| answer_budget = min( | |
| max_answer_tokens, self.model.config.max_seq_len - len(sequence) | |
| ) | |
| cache = self.model.new_cache() if cached else None | |
| step_in = torch.tensor([sequence], dtype=torch.long, device=self.device) | |
| tick = time.perf_counter() | |
| for _ in range(max(0, answer_budget)): | |
| if cached: | |
| seen = cache.get_seq_length() | |
| blocked = prefix_causal_blocked( | |
| prefix_tensor, seen + step_in.shape[1], causal_prefix=self.causal_prefix | |
| )[:, seen:, :] | |
| output_positions = torch.zeros_like(step_in, dtype=torch.bool) | |
| output_positions[0, -1] = True | |
| logits, cache = self.model.forward_cached( | |
| step_in, attn_mask=blocked, past_key_values=cache, | |
| output_positions=output_positions, | |
| ) | |
| else: | |
| current = torch.tensor( | |
| [sequence + answer_ids], dtype=torch.long, device=self.device | |
| ) | |
| blocked = prefix_causal_blocked( | |
| prefix_tensor, current.shape[1], causal_prefix=self.causal_prefix | |
| ) | |
| output_positions = torch.zeros_like(current, dtype=torch.bool) | |
| output_positions[0, -1] = True | |
| logits = self.model( | |
| current, output_positions=output_positions, attn_mask=blocked | |
| ) | |
| logits = _apply_repetition_penalty(logits, answer_ids, repetition_penalty) | |
| logits = _apply_top_p(logits, top_p) | |
| token, _ = _sample_categorical(logits, temperature, generator) | |
| token_id = int(token.item()) | |
| answer_ids.append(token_id) | |
| step_in = torch.tensor([[token_id]], dtype=torch.long, device=self.device) | |
| answer_compute += time.perf_counter() - tick | |
| done = token_id in stop_ids | |
| if step_delay or len(answer_ids) % 4 == 0 or done: | |
| rate = len(answer_ids) / answer_compute if answer_compute > 0 else 0.0 | |
| yield ( | |
| self._adaptive_block_texts(think_ids), | |
| self._decode(answer_ids), | |
| f'respondiendo (AR, token {len(answer_ids)}) · {rate:.0f} tok/s', | |
| ) | |
| if step_delay and not done: | |
| time.sleep(step_delay / 3) | |
| if done: | |
| break | |
| tick = time.perf_counter() | |
| sizes_chosen = [ | |
| size_by_id[t] for t in think_ids if t in size_by_id | |
| ] | |
| total_compute = think_compute + answer_compute | |
| status = ( | |
| f'listo · bloques: {sizes_chosen} · terminated={terminated} · ' | |
| f'pensar {think_compute:.2f}s · responder {len(answer_ids)} tok en ' | |
| f'{answer_compute:.2f}s · total {total_compute:.2f}s' | |
| ) | |
| self.log_lines.append( | |
| f' answer: {len(answer_ids)} tok · {answer_compute:.2f}s wall' | |
| ) | |
| self.log_lines.append( | |
| f'TOTAL: sizes={sizes_chosen} terminated={terminated} · ' | |
| f'{total_compute:.2f}s' | |
| ) | |
| yield self._adaptive_block_texts(think_ids), self._decode(answer_ids), status | |
| def _split_flat(self, generated: list[int]) -> tuple[list[str], str]: | |
| end_think_id = self.reasoning_ids['end_think'] | |
| if end_think_id in generated: | |
| split = generated.index(end_think_id) | |
| think, answer = generated[:split], generated[split + 1:] | |
| else: | |
| think, answer = generated, [] | |
| return [self._decode(think).strip()], self._decode(answer) | |
| class HFCausalEngine: | |
| """Serve a Hugging Face causal LM behind the same streaming interface. | |
| Thinking-mode outputs (Qwen3-style ``<think>...</think>`` prefixes) are routed to the | |
| thoughts panel; tokens after the closing tag stream as the answer. Models without a | |
| ``</think>`` vocabulary entry stream everything as the answer. | |
| """ | |
| def __init__(self, model_path: str, device: torch.device) -> None: | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| self.tokenizer = AutoTokenizer.from_pretrained(model_path) | |
| dtype = torch.bfloat16 if device.type == 'cuda' else torch.float32 | |
| self.model = AutoModelForCausalLM.from_pretrained(model_path, dtype=dtype) | |
| self.model.to(device).eval() | |
| self.device = device | |
| self.objective = 'hf-ar' | |
| self.step = '-' | |
| eos = self.model.generation_config.eos_token_id | |
| self.eos_ids = set(eos if isinstance(eos, (list, tuple)) else [eos]) | |
| end_think = self.tokenizer.convert_tokens_to_ids('</think>') | |
| unk = self.tokenizer.unk_token_id | |
| self.end_think_id = end_think if isinstance(end_think, int) and end_think != unk else None | |
| self.max_new_tokens = 1024 | |
| self.log_lines: list[str] = [] | |
| def _decode(self, ids: list[int]) -> str: | |
| # Qwen3 registers <think>/</think> as regular added tokens, so | |
| # skip_special_tokens leaves them in the decoded text. | |
| text = self.tokenizer.decode(ids, skip_special_tokens=True) | |
| return text.replace('<think>', '').replace('</think>', '').strip() | |
| def _split(self, generated: list[int], split_at: int | None) -> tuple[list[str], str]: | |
| if self.end_think_id is None: | |
| return [], self._decode(generated) | |
| if split_at is None: | |
| thoughts = self._decode(generated) | |
| return ([thoughts] if thoughts else []), '' | |
| thoughts = self._decode(generated[: split_at - 1]) | |
| return ([thoughts] if thoughts else []), self._decode(generated[split_at:]) | |
| def stream_solve( | |
| self, | |
| problem: str, | |
| *, | |
| temperature: float, | |
| steps_per_block: int, | |
| diffusion_steps: int, | |
| max_answer_tokens: int = 200, | |
| repetition_penalty: float = 1.0, | |
| top_p: float = 1.0, | |
| seed: int = 0, | |
| strategy: str = 'ancestral', | |
| step_delay: float = 0.0, | |
| ): | |
| """Yield ``(thoughts, answer, status)`` snapshots while decoding token by token. | |
| Diffusion-only knobs (``steps_per_block``, ``diffusion_steps``, ``strategy``) are | |
| accepted for interface parity and ignored. | |
| """ | |
| del steps_per_block, diffusion_steps, strategy, max_answer_tokens | |
| generator = torch.Generator(device=self.device.type) | |
| generator.manual_seed(seed if seed else int(time.time_ns() % 2**31)) | |
| prompt = self.tokenizer.apply_chat_template( | |
| [{'role': 'user', 'content': problem.strip()}], | |
| tokenize=False, | |
| add_generation_prompt=True, | |
| ) | |
| input_ids = self.tokenizer(prompt, return_tensors='pt').input_ids.to(self.device) | |
| self.log_lines = [ | |
| f'hf-ar · prompt_tokens={input_ids.shape[1]} temp={temperature} top_p={top_p} ' | |
| f'rep_penalty={repetition_penalty} max_new_tokens={self.max_new_tokens} ' | |
| f'(Qwen3 sugerido: temp 0.6 · top_p 0.95 · rep 1.0)' | |
| ] | |
| generated: list[int] = [] | |
| past_key_values = None | |
| current = input_ids | |
| split_at: int | None = None | |
| compute = 0.0 | |
| tick = time.perf_counter() | |
| for _ in range(self.max_new_tokens): | |
| output = self.model(input_ids=current, past_key_values=past_key_values, use_cache=True) | |
| past_key_values = output.past_key_values | |
| logits = output.logits[:, -1, :].float() | |
| logits = _apply_repetition_penalty(logits, generated, repetition_penalty) | |
| logits = _apply_top_p(logits, top_p) | |
| token, _ = _sample_categorical(logits, temperature, generator) | |
| token_id = int(token.item()) | |
| generated.append(token_id) | |
| compute += time.perf_counter() - tick | |
| if split_at is None and token_id == self.end_think_id: | |
| split_at = len(generated) | |
| self.log_lines.append(f' </think> en token {split_at}') | |
| done = token_id in self.eos_ids | |
| if step_delay or done or len(generated) % 4 == 0: | |
| thoughts, answer = self._split(generated, split_at) | |
| rate = len(generated) / compute if compute > 0 else 0.0 | |
| thinking = self.end_think_id is not None and split_at is None | |
| phase = 'pensando' if thinking else 'respondiendo' | |
| yield thoughts, answer, f'{phase} (AR, token {len(generated)}) · {rate:.0f} tok/s' | |
| if step_delay and not done: | |
| time.sleep(step_delay / 3) | |
| if done: | |
| break | |
| current = token.view(1, 1) | |
| tick = time.perf_counter() | |
| thoughts, answer = self._split(generated, split_at) | |
| rate = len(generated) / compute if compute > 0 else 0.0 | |
| terminated = bool(generated) and generated[-1] in self.eos_ids | |
| self.log_lines.append( | |
| f'TOTAL: {len(generated)} tok · {compute:.2f}s ({rate:.0f} tok/s) · ' | |
| f'terminated={terminated}' | |
| ) | |
| yield thoughts, answer, ( | |
| f'listo · total: {len(generated)} tok en {compute:.2f}s ({rate:.0f} tok/s) · ' | |
| f'terminated={terminated}' | |
| ) | |
| def _render_thoughts(slots: list[str]) -> str: | |
| if not slots: | |
| return '<em>sin pensamientos todavía</em>' | |
| chips = [] | |
| for index, text in enumerate(slots, start=1): | |
| chips.append( | |
| f'<div style="margin:6px 0;padding:8px 12px;border-radius:10px;' | |
| f'background:rgba(124,58,237,.10);border:1px solid rgba(124,58,237,.25);">' | |
| f'<b>T{index}</b> · {html.escape(text)}</div>' | |
| ) | |
| return ''.join(chips) | |
| def _chat_display(messages: list[dict[str, str]], partial: str = '') -> list[dict[str, str]]: | |
| display = [{'role': m['role'], 'content': m['content']} for m in messages] | |
| if partial: | |
| display.append({'role': 'assistant', 'content': partial}) | |
| return display | |
| def _chat_ledger(messages: list[dict[str, str]], keep: int) -> str: | |
| return ledger_line(merge_notes(ledger_notes(messages, keep))) | |
| def build_app(engines: dict[str, 'ReasoningEngine | HFCausalEngine']): | |
| import gradio as gr | |
| first = next(iter(engines)) | |
| def solve( | |
| checkpoint_name, problem, temperature, steps_per_block, diffusion_steps, | |
| repetition_penalty, top_p, seed, strategy, slowmo_ms, | |
| ): | |
| engine = engines[checkpoint_name] | |
| if not problem.strip(): | |
| yield '<em>escribí un problema primero</em>', '', 'esperando problema', '' | |
| return | |
| for slots, answer, status in engine.stream_solve( | |
| problem, | |
| temperature=float(temperature), | |
| steps_per_block=int(steps_per_block), | |
| diffusion_steps=int(diffusion_steps), | |
| repetition_penalty=float(repetition_penalty), | |
| top_p=float(top_p), | |
| seed=int(seed), | |
| strategy=str(strategy), | |
| step_delay=float(slowmo_ms) / 1000.0, | |
| ): | |
| boxed = extract_boxed_answer(answer or '') | |
| answer_display = answer + (f'\n\n**→ respuesta extraída: {boxed}**' if boxed else '') | |
| yield _render_thoughts(slots), answer_display, status, '\n'.join(engine.log_lines) | |
| def chat_send( | |
| checkpoint_name, user_text, messages, system_text, keep_last, temperature, | |
| steps_per_block, repetition_penalty, top_p, max_answer_tokens, seed, slowmo_ms, | |
| datalog, | |
| ): | |
| engine = engines[checkpoint_name] if checkpoint_name in engines else None | |
| messages = list(messages or []) | |
| datalog = list(datalog or []) | |
| keep = int(keep_last) | |
| idle = _chat_display(messages), messages, '<em>sin pensamientos todavía</em>' | |
| if engine is None or not getattr(engine, 'chat_ready', False): | |
| yield (*idle, _chat_ledger(messages, keep), | |
| 'elegí un checkpoint adaptativo con tokenizer ChatML', '', gr.skip(), | |
| datalog, gr.skip()) | |
| return | |
| user_text = (user_text or '').strip() | |
| if not user_text: | |
| yield (*idle, _chat_ledger(messages, keep), 'escribí un mensaje primero', | |
| '', gr.skip(), datalog, gr.skip()) | |
| return | |
| messages.append({'role': 'user', 'content': user_text}) | |
| notes = ledger_notes(messages, keep) | |
| merged = merge_notes(notes) | |
| ledger = ledger_line(merged) | |
| start = window_start(messages, keep) | |
| system = (system_text or '').strip() or SYSTEM | |
| prefix = chat_prefix( | |
| [{'role': m['role'], 'content': m['content']} for m in messages[start:]], | |
| system=system, | |
| extra=ledger, | |
| ) | |
| history_snapshot = [ | |
| {'index': index, 'visible': index >= start, **message} | |
| for index, message in enumerate(messages) | |
| ] | |
| answer = '' | |
| status = '' | |
| blocks: list[tuple[int, str]] = [] | |
| thoughts_html = '<em>sin pensamientos todavía</em>' | |
| for slots, answer, status in engine.stream_chat( | |
| prefix, | |
| temperature=float(temperature), | |
| steps_per_block=int(steps_per_block), | |
| max_answer_tokens=int(max_answer_tokens), | |
| repetition_penalty=float(repetition_penalty), | |
| top_p=float(top_p), | |
| seed=int(seed), | |
| step_delay=float(slowmo_ms) / 1000.0, | |
| blocks_out=blocks, | |
| ): | |
| thoughts_html = _render_thoughts(slots) | |
| yield (_chat_display(messages, answer), messages, thoughts_html, ledger, | |
| status, '\n'.join(engine.log_lines), '', datalog, gr.skip()) | |
| note = '; '.join(text for _, text in blocks if text) | |
| if answer.strip() or note: | |
| messages.append({'role': 'assistant', 'content': answer.strip(), 'note': note}) | |
| restored_input = '' | |
| else: | |
| # The turn produced nothing (context-full guard): undo the user message and | |
| # put its text back in the box so it can be resent after lowering the window. | |
| messages.pop() | |
| restored_input = user_text | |
| datalog.append({ | |
| 'turn': sum(1 for m in messages if m['role'] == 'user') + bool(restored_input), | |
| 'checkpoint': checkpoint_name, | |
| 'settings': { | |
| 'temperature': float(temperature), 'steps_per_block': int(steps_per_block), | |
| 'repetition_penalty': float(repetition_penalty), 'top_p': float(top_p), | |
| 'max_answer_tokens': int(max_answer_tokens), 'seed': int(seed), | |
| 'keep_last': keep, 'system': system, | |
| }, | |
| 'history_at_send': history_snapshot, | |
| 'ledger': {'notes': notes, 'merged': merged, 'line': ledger}, | |
| 'prefix_sent': prefix, | |
| 'prefix_tokens': len( | |
| engine.tokenizer.encode(prefix, add_special_tokens=False).ids | |
| ), | |
| 'thinking_blocks': [ | |
| {'size': size, 'content': text} for size, text in blocks | |
| ], | |
| 'answer': answer.strip(), | |
| 'status': status, | |
| 'engine_log': list(engine.log_lines), | |
| }) | |
| yield (_chat_display(messages), messages, thoughts_html, | |
| _chat_ledger(messages, keep), status, '\n'.join(engine.log_lines), | |
| restored_input, datalog, | |
| json.dumps(datalog, indent=2, ensure_ascii=False)) | |
| def chat_reset(): | |
| return ([], [], '<em>sin pensamientos todavía</em>', '', 'conversación reiniciada', | |
| '', '', [], '') | |
| chat_names = [name for name, engine in engines.items() | |
| if getattr(engine, 'chat_ready', False)] | |
| chat_default = next((n for n in chat_names if 'chat' in n), chat_names[0] if chat_names else None) | |
| with gr.Blocks(title='Reasoning playground') as app: | |
| names = { | |
| name: f'{name} · {engines[name].objective} · step {engines[name].step}' | |
| for name in engines | |
| } | |
| gr.Markdown('# Reasoning playground') | |
| with gr.Tab('chat'): | |
| chat_state = gr.State([]) | |
| chat_datalog_state = gr.State([]) | |
| with gr.Row(): | |
| with gr.Column(scale=2): | |
| chat_checkpoint = gr.Dropdown( | |
| choices=chat_names, value=chat_default, label='checkpoint', | |
| info='adaptativos sobre tokenizer ChatML; solo los entrenados ' | |
| 'en chat (chat-sft) conocen este formato', | |
| ) | |
| chat_system = gr.Textbox(label='system', value=SYSTEM, lines=3) | |
| chat_keep = gr.Slider( | |
| 0, 12, value=4, step=2, | |
| label='mensajes visibles (0 = historial completo; lo anterior ' | |
| 'sobrevive solo en el ledger)', | |
| ) | |
| chat_temperature = gr.Slider(0.0, 1.2, value=0.8, step=0.05, | |
| label='temperatura') | |
| chat_steps = gr.Slider( | |
| 1, 64, value=16, step=1, | |
| label='pasos de denoising por bloque (32 es el óptimo medido)', | |
| ) | |
| chat_rep = gr.Slider(1.0, 2.0, value=1.4, step=0.05, | |
| label='penalización de repetición (respuesta)') | |
| chat_top_p = gr.Slider(0.1, 1.0, value=0.92, step=0.02, | |
| label='top-p (respuesta)') | |
| chat_max_answer = gr.Slider(32, 512, value=224, step=16, | |
| label='tokens máximos de respuesta') | |
| chat_slowmo = gr.Slider(0, 600, value=0, step=20, | |
| label='cámara lenta (ms por paso de denoising)') | |
| chat_seed = gr.Number(value=0, label='seed (0 = aleatoria)', precision=0) | |
| chat_clear = gr.Button('reiniciar conversación') | |
| with gr.Column(scale=3): | |
| chatbot = gr.Chatbot(label='conversación', height=420) | |
| chat_input = gr.Textbox( | |
| label='mensaje', lines=2, | |
| placeholder='p.ej. My policy is PL-48291. — o cualquier pregunta', | |
| ) | |
| chat_go = gr.Button('enviar', variant='primary') | |
| chat_status = gr.Markdown('esperando mensaje') | |
| chat_thoughts = gr.HTML(label='pensamientos del turno') | |
| chat_ledger_box = gr.Textbox( | |
| label='ledger (notas del modelo fuera de la ventana visible)', | |
| interactive=False, | |
| ) | |
| chat_logs = gr.Textbox(label='logs', lines=8, max_lines=20, | |
| interactive=False) | |
| with gr.Accordion('datalog — todo lo que viajó, por turno', open=False): | |
| chat_datalog_box = gr.Textbox( | |
| label='sesión completa en JSON: settings, historial con ventana, ' | |
| 'ledger, prefijo exacto, bloques de thinking, respuesta y logs', | |
| lines=18, max_lines=40, interactive=False, buttons=['copy'], | |
| ) | |
| chat_inputs = [ | |
| chat_checkpoint, chat_input, chat_state, chat_system, chat_keep, | |
| chat_temperature, chat_steps, chat_rep, chat_top_p, chat_max_answer, | |
| chat_seed, chat_slowmo, chat_datalog_state, | |
| ] | |
| chat_outputs = [ | |
| chatbot, chat_state, chat_thoughts, chat_ledger_box, chat_status, | |
| chat_logs, chat_input, chat_datalog_state, chat_datalog_box, | |
| ] | |
| chat_go.click(chat_send, inputs=chat_inputs, outputs=chat_outputs) | |
| chat_input.submit(chat_send, inputs=chat_inputs, outputs=chat_outputs) | |
| chat_clear.click(chat_reset, inputs=[], outputs=chat_outputs) | |
| with gr.Tab('resolver'): | |
| with gr.Row(): | |
| with gr.Column(scale=2): | |
| checkpoint = gr.Dropdown( | |
| choices=list(engines), value=first, label='checkpoint', | |
| info=' | '.join(names.values()), | |
| ) | |
| problem = gr.Textbox( | |
| label='problema / instrucción', | |
| lines=4, | |
| placeholder=( | |
| 'hybrid: Write a short story. It should feature: Dialogue. ' | |
| 'Use the words: dragon, cake, brave.\n' | |
| 'lm (base): cualquier texto a continuar, p.ej. ' | |
| '"Tom was a happy boy who"' | |
| ), | |
| ) | |
| temperature = gr.Slider(0.0, 1.2, value=0.8, step=0.05, | |
| label='temperatura') | |
| steps_per_block = gr.Slider( | |
| 1, 64, value=16, step=1, | |
| label='pasos de denoising por slot/bloque (hybrid)' | |
| ) | |
| diffusion_steps = gr.Slider( | |
| 8, 256, value=64, step=8, label='pasos totales (diffusion)' | |
| ) | |
| repetition_penalty = gr.Slider( | |
| 1.0, 2.0, value=1.4, step=0.05, | |
| label='penalización de repetición (respuesta)' | |
| ) | |
| top_p = gr.Slider( | |
| 0.1, 1.0, value=0.92, step=0.02, label='top-p (respuesta)' | |
| ) | |
| strategy = gr.Dropdown( | |
| choices=['ancestral', 'confidence', 'left_to_right'], | |
| value='ancestral', | |
| label='orden de revelado (denoising; ancestral es el único que rinde)', | |
| ) | |
| slowmo_ms = gr.Slider( | |
| 0, 600, value=0, step=20, | |
| label='cámara lenta (ms por paso de denoising, 0 = tiempo real)', | |
| ) | |
| seed = gr.Number(value=0, label='seed (0 = aleatoria)', precision=0) | |
| go = gr.Button('resolver', variant='primary') | |
| with gr.Column(scale=3): | |
| status = gr.Markdown('esperando problema') | |
| thoughts = gr.HTML(label='pensamientos') | |
| answer = gr.Markdown(label='respuesta') | |
| logs = gr.Textbox( | |
| label='logs (copiá y pegá)', lines=12, max_lines=30, | |
| interactive=False, | |
| ) | |
| go.click( | |
| solve, | |
| inputs=[ | |
| checkpoint, problem, temperature, steps_per_block, diffusion_steps, | |
| repetition_penalty, top_p, seed, strategy, slowmo_ms, | |
| ], | |
| outputs=[thoughts, answer, status, logs], | |
| ) | |
| return app | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument('--outputs-dir', type=Path, default=Path('outputs')) | |
| parser.add_argument('--prefix', default='reasoning-') | |
| parser.add_argument('--tokenizer', type=Path, required=True) | |
| parser.add_argument('--device', default='auto') | |
| parser.add_argument('--host', default='127.0.0.1') | |
| parser.add_argument('--port', type=int, default=7999) | |
| parser.add_argument( | |
| '--hf-model', action='append', default=[], metavar='NAME=PATH', | |
| help='serve a Hugging Face causal LM (local path or repo id) alongside the checkpoints', | |
| ) | |
| args = parser.parse_args() | |
| device = resolve_device(args.device) | |
| checkpoints = _discover_checkpoints(args.outputs_dir) | |
| checkpoints = { | |
| name: path for name, path in checkpoints.items() if name.startswith(args.prefix) | |
| } | |
| if not checkpoints and not args.hf_model: | |
| raise SystemExit(f'no reasoning checkpoints under {args.outputs_dir}') | |
| engines: dict[str, ReasoningEngine | HFCausalEngine] = { | |
| name: ReasoningEngine(path, args.tokenizer, device) | |
| for name, path in checkpoints.items() | |
| } | |
| for spec in args.hf_model: | |
| name, _, path = spec.partition('=') | |
| if not name or not path: | |
| raise SystemExit(f'--hf-model expects NAME=PATH, got {spec!r}') | |
| engines[name] = HFCausalEngine(path, device) | |
| print(f'loaded: {", ".join(engines)} on {device}') | |
| app = build_app(engines) | |
| app.queue().launch(server_name=args.host, server_port=args.port) | |
| if __name__ == '__main__': | |
| main() | |