Modilify-Mk2-preview / generation_modilify_mk2.py
ydy9038074's picture
Publish Modilify Mk2 Preview step 1250 schema25
b88f761 verified
Raw History Blame Contribute Delete
47.8 kB
"""Rolling text generation for latent-memory ModilifyMk2."""
from __future__ import annotations
from collections.abc import Sequence
from contextlib import contextmanager
from dataclasses import dataclass, replace
import math
from typing import Any
import torch
from transformers.cache_utils import Cache
from transformers.generation import LogitsProcessorList
from transformers.generation.streamers import BaseStreamer
from transformers.modeling_outputs import ModelOutput
from transformers.models.diffusion_gemma import (
DiffusionGemmaGenerationConfig,
DiffusionGemmaGenerationMixin,
)
from .commit_policy import (
fused_commit_failure_rate,
select_commit_lengths,
)
from .configuration_modilify_mk2 import DENOISE_TEMPERATURE
from .latent_deliberation import LatentDeliberationState, infer_commit_reason
def deterministic_episode_iteration_bound(
response_lengths: torch.LongTensor,
*,
max_ponder_steps: int,
) -> int:
"""Return a safe watchdog bound without assuming canvas-sized jumps.
Every active row must advance by at least one token no later than the
configured no-progress threshold. Normal commits can also be only one
token long, so a bound based on the number of canvases is not valid.
"""
if response_lengths.numel() == 0:
raise ValueError("`response_lengths` must be non-empty.")
if max_ponder_steps <= 0:
raise ValueError("`max_ponder_steps` must be positive.")
return max(1, int(response_lengths.max()) * max_ponder_steps)
_PARENT_GENERATION_KEYS = frozenset({
"max_new_tokens",
"max_length",
"return_dict_in_generate",
"max_denoising_steps",
"t_min",
"t_max",
"cache_implementation",
"cache_config",
"disable_compile",
"bos_token_id",
"pad_token_id",
"eos_token_id",
"_commit_hash",
"_from_model_config",
"transformers_version",
})
def _flatten_token_ids(*values: object) -> set[int]:
"""Normalize scalar and sequence token-ID configuration values."""
token_ids: set[int] = set()
for value in values:
if value is None:
continue
if isinstance(value, int):
token_ids.add(int(value))
continue
if isinstance(value, (list, tuple, set)):
token_ids.update(int(token_id) for token_id in value if token_id is not None)
return token_ids
def _add_repetition_history(
history: torch.BoolTensor,
token_ids: torch.LongTensor,
eligible: torch.BoolTensor,
excluded_token_ids: set[int],
) -> None:
"""Add eligible row-local token IDs to a compact [batch, vocab] history."""
if token_ids.shape != eligible.shape or token_ids.shape[0] != history.shape[0]:
raise ValueError("Repetition history token and eligibility shapes must match.")
eligible = eligible.clone()
for token_id in excluded_token_ids:
eligible &= token_ids.ne(token_id)
if not bool(eligible.any()):
return
rows = torch.arange(history.shape[0], device=history.device)[:, None]
rows = rows.expand_as(token_ids)
history[rows[eligible], token_ids[eligible]] = True
class ModilifyMk2GenerationConfig(DiffusionGemmaGenerationConfig):
def __init__(self, **kwargs):
self.turn_end_token_id: int | None = kwargs.pop("turn_end_token_id", None)
self.max_ponder_steps: int = kwargs.pop("max_ponder_steps", 64)
self.jump_on_no_progress_after: int = kwargs.pop("jump_on_no_progress_after", 12)
self.min_trajectory_progress: float = float(kwargs.pop("min_trajectory_progress", 0.005))
self.repetition_penalty: float = float(kwargs.pop("repetition_penalty", 1.0))
excluded_token_ids = kwargs.pop("repetition_penalty_exclude_token_ids", ())
self.repetition_penalty_exclude_token_ids: list[int] = list(
dict.fromkeys(int(token_id) for token_id in excluded_token_ids or ())
)
kwargs.pop("t_min", None)
kwargs.pop("t_max", None)
parent_kwargs = {
name: kwargs.pop(name)
for name in tuple(kwargs)
if name in _PARENT_GENERATION_KEYS
}
super().__init__(
sampler_config=None,
stability_threshold=None,
confidence_threshold=None,
**parent_kwargs,
)
self.t_min = DENOISE_TEMPERATURE
self.t_max = DENOISE_TEMPERATURE
def update(self, defaults_only=False, allow_custom_entries=False, **kwargs):
"""Apply supported overrides while keeping temperature globally fixed."""
if "turn_end_token_id" in kwargs:
self.turn_end_token_id = kwargs.pop("turn_end_token_id")
if "max_ponder_steps" in kwargs:
self.max_ponder_steps = kwargs.pop("max_ponder_steps")
if "jump_on_no_progress_after" in kwargs:
self.jump_on_no_progress_after = kwargs.pop("jump_on_no_progress_after")
if "min_trajectory_progress" in kwargs:
self.min_trajectory_progress = float(kwargs.pop("min_trajectory_progress"))
if "repetition_penalty" in kwargs:
self.repetition_penalty = float(kwargs.pop("repetition_penalty"))
if "repetition_penalty_exclude_token_ids" in kwargs:
excluded_token_ids = kwargs.pop("repetition_penalty_exclude_token_ids")
self.repetition_penalty_exclude_token_ids = list(
dict.fromkeys(int(token_id) for token_id in excluded_token_ids or ())
)
kwargs.pop("t_min", None)
kwargs.pop("t_max", None)
unused = super().update(
defaults_only=defaults_only,
allow_custom_entries=allow_custom_entries,
**kwargs,
)
self.sampler_config = None
self.stability_threshold = None
self.confidence_threshold = None
self.t_min = DENOISE_TEMPERATURE
self.t_max = DENOISE_TEMPERATURE
return unused
def validate(self, **kwargs):
if self.max_new_tokens is not None and (
not isinstance(self.max_new_tokens, int) or self.max_new_tokens <= 0
):
raise ValueError(f"`max_new_tokens` must be a positive integer, but got {self.max_new_tokens}")
if self.max_length is not None and (
not isinstance(self.max_length, int) or self.max_length <= 0
):
raise ValueError(f"`max_length` must be a positive integer, but got {self.max_length}")
if self.turn_end_token_id is not None and (
not isinstance(self.turn_end_token_id, int) or self.turn_end_token_id < 0
):
raise ValueError("`turn_end_token_id` must be a non-negative integer.")
if not isinstance(self.max_ponder_steps, int) or self.max_ponder_steps <= 0:
raise ValueError("`max_ponder_steps` must be a positive integer.")
if not isinstance(self.jump_on_no_progress_after, int) or self.jump_on_no_progress_after <= 0:
raise ValueError("`jump_on_no_progress_after` must be a positive integer.")
if not isinstance(self.min_trajectory_progress, (int, float)) or self.min_trajectory_progress < 0:
raise ValueError("`min_trajectory_progress` must be a non-negative number.")
if not math.isfinite(self.repetition_penalty) or self.repetition_penalty <= 0:
raise ValueError("`repetition_penalty` must be a finite positive number.")
if any(
not isinstance(token_id, int) or isinstance(token_id, bool) or token_id < 0
for token_id in self.repetition_penalty_exclude_token_ids
):
raise ValueError(
"`repetition_penalty_exclude_token_ids` must contain non-negative integers."
)
@classmethod
def from_model_config(cls, model_config):
"""Build the only generation field owned by the ModilifyMk2 model config."""
return cls(turn_end_token_id=model_config.turn_end_token_id)
@staticmethod
def _get_default_generation_params() -> dict[str, object]:
"""Return defaults with no inherited entropy/readiness commit controls."""
return {
"max_new_tokens": 256,
"max_denoising_steps": 48,
"t_min": DENOISE_TEMPERATURE,
"t_max": DENOISE_TEMPERATURE,
}
@dataclass
class ModilifyMk2GenerationOutput(ModelOutput):
sequences: torch.LongTensor
generated_lengths: torch.LongTensor | None = None
tokens_per_forward: torch.FloatTensor | None = None
past_key_values: Cache | None = None
stop_reason: str | tuple[str, ...] | None = None
committed_tokens: int | torch.LongTensor | None = None
denoise_steps: int | torch.LongTensor | None = None
no_progress_steps: int | torch.LongTensor | None = None
jump_count: int | torch.LongTensor | None = None
forced_jump_bad_count: int | torch.LongTensor | None = None
average_commit_len: float | torch.FloatTensor | None = None
state_shift_count: int | torch.LongTensor | None = None
@dataclass
class ModilifyMk2RollingState:
"""All real iterative state; no vocabulary-sized tensor is retained."""
canvas: torch.LongTensor
confidence: torch.FloatTensor
entropy: torch.FloatTensor
latent_state: LatentDeliberationState
head: torch.LongTensor | None = None
class NoiseCanvasSampler:
"""Uniform diffusion noise source with no commit-policy responsibilities."""
def __init__(self, *, canvas_length: int, vocab_size: int) -> None:
self.canvas_length = int(canvas_length)
self.vocab_size = int(vocab_size)
self.initial_entropy = math.log(self.vocab_size)
def initialize_canvas(
self,
batch_size: int,
device: torch.device,
generators: Sequence[torch.Generator] | None = None,
) -> torch.LongTensor:
if generators is not None:
if len(generators) != batch_size:
raise ValueError("Canvas sampling requires one generator per batch row.")
return torch.cat(
[
torch.randint(
self.vocab_size,
(1, self.canvas_length),
device=device,
generator=generator,
)
for generator in generators
],
dim=0,
)
return torch.randint(
self.vocab_size,
(batch_size, self.canvas_length),
device=device,
)
class ModilifyMk2GenerationMixin(DiffusionGemmaGenerationMixin):
def init_continuous_batching(
self,
generation_config=None,
continuous_batching_config=None,
workload_hints=None,
):
"""Create the ModilifyMk2 continuous manager behind the official API surface."""
from .continuous_batching import (
ModilifyMk2ContinuousBatchingManager,
continuous_config_fingerprint,
)
cached = getattr(self, "_cached_continuous_batching_manager", None)
if isinstance(cached, ModilifyMk2ContinuousBatchingManager) and not cached.destroyed:
requested_generation = generation_config or getattr(
self, "generation_config", None
)
requested_fingerprint = continuous_config_fingerprint(
requested_generation,
continuous_batching_config,
)
if cached.config_fingerprint == requested_fingerprint:
cached._prepare_for_next_session()
return cached
cached.destroy()
delattr(self, "_cached_continuous_batching_manager")
return ModilifyMk2ContinuousBatchingManager(
model=self,
generation_config=generation_config or getattr(self, "generation_config", None),
continuous_batching_config=continuous_batching_config,
workload_hints=workload_hints,
)
def destroy_cached_continuous_batching_manager(self) -> None:
manager = getattr(self, "_cached_continuous_batching_manager", None)
if manager is not None:
manager.destroy()
delattr(self, "_cached_continuous_batching_manager")
@contextmanager
@torch.no_grad()
def continuous_batching_context_manager(
self,
generation_config=None,
block: bool = True,
timeout: float | None = None,
continuous_batching_config=None,
persistent_manager: bool = False,
warmup: bool = True,
workload_hints=None,
):
manager = self.init_continuous_batching(
generation_config=generation_config,
continuous_batching_config=continuous_batching_config,
workload_hints=workload_hints,
)
if persistent_manager:
self._cached_continuous_batching_manager = manager
if warmup and not manager.warmed_up:
manager.warmup()
manager.start()
try:
yield manager
finally:
manager.stop(
block=block,
timeout=timeout,
keep_for_next_session=persistent_manager,
)
if not persistent_manager:
manager.destroy()
@torch.no_grad()
def generate_batch(
self,
inputs: list[list[int]],
generation_config=None,
continuous_batching_config=None,
record_timestamps: bool = False,
progress_bar: bool = True,
persistent_manager: bool = False,
warmup: bool = True,
**kwargs: Any,
) -> dict[str, object]:
"""Official-compatible convenience wrapper over the continuous manager."""
del progress_bar
if any(
not isinstance(input_ids, list)
or not input_ids
or any(
not isinstance(token_id, int) or isinstance(token_id, bool)
for token_id in input_ids
)
for input_ids in inputs
):
raise ValueError("Every `inputs` row must be a non-empty list of integer token IDs.")
seeds = kwargs.pop("seeds", None)
if seeds is not None and len(seeds) != len(inputs):
raise ValueError("`seeds` must contain one seed per request.")
if seeds is not None:
seeds = [int(seed) for seed in seeds]
if not inputs:
return {}
manager = self.init_continuous_batching(
generation_config=generation_config,
continuous_batching_config=continuous_batching_config,
)
if persistent_manager:
self._cached_continuous_batching_manager = manager
if warmup and not manager.warmed_up:
manager.warmup()
completed = False
try:
request_ids = []
queue_limit = int(
manager.continuous_batching_config.max_queue_size or 0
)
preload_count = (
len(inputs) if queue_limit == 0 else min(len(inputs), queue_limit)
)
for index, input_ids in enumerate(inputs):
if index == preload_count and not manager.is_running():
manager.start()
request_kwargs = dict(kwargs)
if seeds is not None:
request_kwargs["seed"] = int(seeds[index])
request_ids.append(
manager.add_request(
input_ids=input_ids,
record_timestamps=record_timestamps,
**request_kwargs,
)
)
manager.close_input()
# Unlike the open-ended manager API, generate_batch has its whole
# initial workload. Queue it before starting so the first cohort is
# filled deterministically up to scheduler/queue capacity.
if not manager.is_running():
manager.start()
final_outputs = {}
for output in manager:
if output.status in {
getattr(output.status.__class__, "FINISHED", output.status),
getattr(output.status.__class__, "FAILED", output.status),
}:
final_outputs[output.request_id] = output
result = {
request_id: final_outputs[request_id]
for request_id in request_ids
if request_id is not None
}
completed = True
finally:
manager.stop(
block=True,
keep_for_next_session=persistent_manager,
hard_stop=not completed,
)
if not persistent_manager:
manager.destroy()
return result
def _prepare_sampler(
self, generation_config: ModilifyMk2GenerationConfig, canvas_length: int | None = None
) -> NoiseCanvasSampler:
del generation_config
return NoiseCanvasSampler(
canvas_length=canvas_length or self.config.canvas_length,
vocab_size=self.config.text_config.vocab_size,
)
def _write_committed_memory(
self,
*,
next_state: ModilifyMk2RollingState,
working_state: torch.Tensor,
heavy_hidden: torch.Tensor,
commit_token_ids: torch.LongTensor,
commit_lengths: torch.Tensor,
prefix_lengths: torch.Tensor,
commit_reason: torch.Tensor,
canvas_head: torch.Tensor | None = None,
) -> ModilifyMk2RollingState:
if not bool(commit_lengths.gt(0).any()):
return next_state
batch, canvas = commit_token_ids.shape
if working_state.shape[:2] != (batch, canvas):
raise ValueError("Committed token IDs must match the unshifted canvas.")
max_commit = min(int(commit_lengths.max()), canvas)
if canvas_head is None:
selected_ids = commit_token_ids[:, :max_commit]
else:
head = canvas_head.to(device=commit_token_ids.device, dtype=torch.long).view(batch, 1)
physical = (head + torch.arange(max_commit, device=commit_token_ids.device)) % canvas
selected_ids = commit_token_ids.gather(1, physical)
with torch.no_grad():
committed_token_embeddings = self.model.decoder.embed_tokens(selected_ids)
memory = self.latent_deliberation.commit_write(
memory=next_state.latent_state.memory_slots,
working_state=working_state,
heavy_hidden=heavy_hidden,
committed_token_embeddings=committed_token_embeddings,
commit_lengths=commit_lengths,
commit_reason=commit_reason,
canvas_head=canvas_head,
)
return replace(
next_state,
latent_state=replace(
next_state.latent_state, memory_slots=memory,
gdn2=replace(next_state.latent_state.gdn2, persistent=memory),
),
)
@staticmethod
def _shift_state_rows(
state: ModilifyMk2RollingState,
commit_lengths: torch.LongTensor,
sampler: NoiseCanvasSampler,
generators: Sequence[torch.Generator] | None = None,
remaining_lengths: torch.LongTensor | None = None,
pad_token_id: int = 0,
) -> ModilifyMk2RollingState:
"""Shift every rolling row by its own committed prefix length."""
batch_size, canvas_length = state.canvas.shape
if commit_lengths.shape != (batch_size,):
raise ValueError("Commit lengths must have shape [batch].")
if not bool(commit_lengths.gt(0).any()):
if generators is not None:
if len(generators) != batch_size:
raise ValueError("State shifting requires one generator per batch row.")
# Seeded generation deliberately advances every active request
# once per denoise step, independent of the other active rows.
for generator in generators:
sampler.initialize_canvas(
1, state.canvas.device, generators=[generator]
)
return state
positions = torch.arange(canvas_length, device=state.canvas.device)[None, :]
source = positions + commit_lengths[:, None]
retained = source.lt(canvas_length)
def shift(value: torch.Tensor, fill_value: float | int = 0) -> torch.Tensor:
index = source.clamp_max(canvas_length - 1)
index = index.view(
batch_size, canvas_length, *([1] * (value.ndim - 2))
).expand_as(value)
gathered = value.gather(1, index)
mask = retained.view(
batch_size, canvas_length, *([1] * (value.ndim - 2))
)
fill = torch.as_tensor(fill_value, device=value.device, dtype=value.dtype)
return torch.where(mask, gathered, fill)
if generators is None:
tail = sampler.initialize_canvas(batch_size, state.canvas.device)
else:
if len(generators) != batch_size:
raise ValueError("State shifting requires one generator per batch row.")
tail = torch.zeros_like(state.canvas)
for row, (commit_length, generator) in enumerate(
zip(commit_lengths.detach().cpu().tolist(), generators, strict=True)
):
sampled = sampler.initialize_canvas(
1,
state.canvas.device,
generators=[generator],
)
tail[row] = sampled[0]
canvas = torch.cat((state.canvas, tail), dim=1).gather(1, source)
if remaining_lengths is not None:
if remaining_lengths.shape != (batch_size,):
raise ValueError("Remaining lengths must have shape [batch].")
canvas = canvas.masked_fill(
(positions >= (canvas_length - commit_lengths)[:, None])
& (positions >= remaining_lengths[:, None]),
pad_token_id,
)
unknown_entropy = float(sampler.initial_entropy)
latent = state.latent_state
committed = commit_lengths.gt(0)
shifted_latent = LatentDeliberationState(
memory_slots=latent.memory_slots.clone(),
confidence=shift(latent.confidence),
entropy=shift(latent.entropy, unknown_entropy),
ponder_steps=torch.where(
committed, torch.zeros_like(latent.ponder_steps), latent.ponder_steps
),
stagnation_steps=torch.where(
committed, torch.zeros_like(latent.stagnation_steps), latent.stagnation_steps
),
gdn2=latent.gdn2.shift(commit_lengths),
)
return ModilifyMk2RollingState(
canvas=canvas,
confidence=shift(state.confidence),
entropy=shift(state.entropy, unknown_entropy),
latent_state=shifted_latent,
head=state.head,
)
@torch.inference_mode()
def generate(
self,
input_ids: torch.LongTensor | None = None,
past_key_values: Cache | None = None,
streamer: BaseStreamer | None = None,
generation_config: ModilifyMk2GenerationConfig | None = None,
logits_processor: LogitsProcessorList | None = None,
**kwargs,
) -> ModilifyMk2GenerationOutput:
request_seeds = kwargs.pop("seeds", None)
scalar_seed = kwargs.pop("seed", None)
if request_seeds is not None and scalar_seed is not None:
raise ValueError("Pass either `seed` or `seeds`, not both.")
generation_config, model_kwargs = self._prepare_generation_config(generation_config, **kwargs)
if input_ids is None or input_ids.ndim != 2 or input_ids.shape[0] < 1:
raise ValueError("ModilifyMk2 generation requires `input_ids` with shape [batch, sequence].")
if logits_processor:
raise ValueError(
"ModilifyMk2 uses its built-in exact sampler and does not accept "
"custom logits processors."
)
batch_size, input_width = input_ids.shape
if scalar_seed is not None:
request_seeds = [int(scalar_seed) + row for row in range(batch_size)]
elif batch_size > 1 and request_seeds is None:
# Static B>1 still uses independent row generators, but its default
# path must advance the caller's global device RNG just like normal
# generation instead of repeating torch.initial_seed() forever.
seed_parts = torch.randint(
0,
(1 << 31) - 1,
(batch_size, 2),
device=input_ids.device,
dtype=torch.int64,
).detach().cpu().tolist()
request_seeds = [
(int(high) << 31) | int(low) for high, low in seed_parts
]
if request_seeds is not None:
if len(request_seeds) != batch_size:
raise ValueError("`seeds` must contain one seed per batch row.")
sampling_generators = []
for seed_value in request_seeds:
generator = torch.Generator(device=input_ids.device)
generator.manual_seed(int(seed_value) & ((1 << 63) - 1))
sampling_generators.append(generator)
else:
sampling_generators = None
if batch_size > 1 and streamer is not None:
raise ValueError("ModilifyMk2 streamers currently support batch size 1 only.")
if batch_size > 1 and past_key_values is not None:
raise ValueError("Batched ModilifyMk2 generation requires a fresh KV cache.")
if batch_size > 1:
from .continuous_batching import generate_static_batch_with_logical_cache
attention_mask = model_kwargs.pop("attention_mask", None)
canonical_mask = (
torch.ones_like(input_ids, dtype=torch.bool)
if attention_mask is None
else attention_mask.to(device=input_ids.device, dtype=torch.bool)
)
if canonical_mask.shape != input_ids.shape:
raise ValueError(
"`attention_mask` must have the same shape as `input_ids`."
)
row_limits = []
for prompt_length in canonical_mask.long().sum(dim=-1).tolist():
_, resolved_max_new_tokens = self._prepare_generated_length(
generation_config, int(prompt_length)
)
if resolved_max_new_tokens <= 0:
raise ValueError(
"The requested maximum length leaves no room to generate "
"for every batch row."
)
row_limits.append(int(resolved_max_new_tokens))
provided_positions = model_kwargs.pop("position_ids", None)
if provided_positions is not None:
canonical_positions = (
canonical_mask.long().cumsum(dim=-1).sub(1).clamp_min(0)
).to(provided_positions)
if not torch.equal(provided_positions, canonical_positions):
raise ValueError(
"Continuous static batches require canonical row-local `position_ids`."
)
if model_kwargs:
unsupported = ", ".join(sorted(model_kwargs))
raise ValueError(
f"Unsupported batched ModilifyMk2 generation arguments: {unsupported}"
)
return generate_static_batch_with_logical_cache(
self,
input_ids,
canonical_mask,
generation_config,
seeds=request_seeds,
max_new_tokens=row_limits,
)
device = input_ids.device
dtype = self.model.decoder.embed_tokens.weight.dtype
canvas_length = self.config.canvas_length
cached_length = past_key_values.get_seq_length() if past_key_values is not None else 0
repetition_penalty = float(generation_config.repetition_penalty)
repetition_enabled = repetition_penalty != 1.0
if repetition_enabled and cached_length:
raise ValueError(
"Repetition penalty requires a fresh KV cache so the complete prompt "
"token history is available."
)
_, max_new_tokens = self._prepare_generated_length(
generation_config, cached_length + input_width
)
max_iterations = deterministic_episode_iteration_bound(
torch.tensor([max_new_tokens]),
max_ponder_steps=generation_config.max_ponder_steps,
)
if past_key_values is None:
past_key_values = self._prepare_cache_for_generation(
generation_config,
batch_size=batch_size,
# Ragged batches append dense, masked cache blocks. In the
# worst case only one row advances in each block.
max_length=input_width + batch_size * max_new_tokens,
)
expected_mask_width = cached_length + input_width
cache_attention_mask = model_kwargs.pop(
"attention_mask",
torch.ones(
batch_size, expected_mask_width, dtype=torch.bool, device=device
),
).bool()
if cache_attention_mask.shape != (batch_size, expected_mask_width):
raise ValueError(
"`attention_mask` must have shape [batch, cached_length + sequence]."
)
provided_position_ids = model_kwargs.pop("position_ids", None)
if provided_position_ids is not None:
if provided_position_ids.shape != input_ids.shape:
raise ValueError("`position_ids` must have the same shape as `input_ids`.")
prompt_positions = provided_position_ids.to(device=device, dtype=torch.int32)
elif cached_length:
prompt_positions = torch.arange(
cached_length,
cached_length + input_width,
device=device,
dtype=torch.int32,
).unsqueeze(0)
else:
input_mask = cache_attention_mask[:, -input_width:]
prompt_positions = input_mask.long().cumsum(dim=-1).sub(1).clamp_min(0).to(torch.int32)
logical_lengths = cache_attention_mask.long().sum(dim=-1)
if input_width:
past_key_values = self.model.encoder(
input_ids=input_ids,
attention_mask=cache_attention_mask,
past_key_values=past_key_values,
position_ids=prompt_positions,
).past_key_values
sampler = self._prepare_sampler(generation_config, canvas_length)
latent = LatentDeliberationState.empty(
batch_size=batch_size, canvas_length=canvas_length,
device=device,
)
initial_canvas = sampler.initialize_canvas(
batch_size, device, generators=sampling_generators
)
state = ModilifyMk2RollingState(
canvas=initial_canvas,
confidence=torch.zeros(
batch_size, canvas_length, device=device, dtype=torch.float32
),
entropy=torch.full(
(batch_size, canvas_length), math.log(self.config.text_config.vocab_size),
device=device, dtype=torch.float32,
),
latent_state=latent,
head=torch.zeros(batch_size, device=device, dtype=torch.long),
)
turn_end = (
self.config.turn_end_token_id
if generation_config.turn_end_token_id is None
else generation_config.turn_end_token_id
)
configured_eos = generation_config.eos_token_id
if configured_eos is None:
configured_eos = self.config.eos_token_id
if isinstance(configured_eos, int):
configured_eos = [configured_eos]
stop_token_ids = tuple(
dict.fromkeys((int(turn_end), *(int(value) for value in configured_eos or ())))
)
pad_token_id = generation_config.pad_token_id
if pad_token_id is None:
pad_token_id = getattr(self.config, "pad_token_id", None)
if isinstance(pad_token_id, (list, tuple)):
pad_token_id = pad_token_id[0]
pad_token_id = int(0 if pad_token_id is None else pad_token_id)
state = replace(
state,
canvas=state.canvas.masked_fill(
torch.arange(canvas_length, device=device)[None, :] >= max_new_tokens,
pad_token_id,
),
)
excluded_repetition_token_ids = _flatten_token_ids(
generation_config.repetition_penalty_exclude_token_ids,
generation_config.pad_token_id,
generation_config.bos_token_id,
generation_config.eos_token_id,
generation_config.turn_end_token_id,
getattr(self.config, "image_token_id", None),
)
repetition_history = None
if repetition_enabled:
repetition_history = torch.zeros(
(batch_size, self.config.text_config.vocab_size),
dtype=torch.bool,
device=device,
)
_add_repetition_history(
repetition_history,
input_ids,
cache_attention_mask[:, -input_width:],
excluded_repetition_token_ids,
)
generated = torch.full(
(batch_size, max_new_tokens),
pad_token_id,
dtype=input_ids.dtype,
device=device,
)
committed = torch.zeros(batch_size, dtype=torch.long, device=device)
denoise_steps = torch.zeros_like(committed)
jumps = torch.zeros_like(committed)
forced_jump_tokens = torch.zeros_like(committed)
shifts = torch.zeros_like(committed)
stop_codes = torch.zeros_like(committed)
active_rows = torch.ones(batch_size, dtype=torch.bool, device=device)
canvas_positions = torch.arange(canvas_length, device=device)[None, :]
if streamer is not None:
streamer.put(input_ids.cpu())
while bool(active_rows.any()):
decoder_positions = (
logical_lengths[:, None]
+ torch.arange(canvas_length, device=device)[None, :]
).to(torch.int32)
denoise_steps += active_rows.long()
decoder_attention_mask = torch.cat(
(
cache_attention_mask,
torch.ones(
batch_size,
canvas_length,
dtype=torch.bool,
device=device,
),
),
dim=-1,
)
output = self(
input_ids=None, past_key_values=past_key_values,
decoder_input_ids=state.canvas,
previous_confidence=state.confidence, previous_entropy=state.entropy,
latent_state=state.latent_state,
decoder_position_ids=decoder_positions, decoder_read_cache=True,
decoder_attention_mask=decoder_attention_mask,
compact_vocab=True,
repetition_token_mask=repetition_history,
repetition_penalty=repetition_penalty,
sampling_generators=sampling_generators,
**model_kwargs,
)
if (
output.proposal is None
or output.proposal_confidence is None
or output.token_entropy is None
or output.greedy_proposal is None
or output.greedy_confidence is None
):
raise RuntimeError(
"Compact vocabulary forward did not return proposal statistics."
)
proposal = output.proposal
proposal_confidence = output.proposal_confidence
token_entropy = output.token_entropy
greedy_proposal = output.greedy_proposal
greedy_confidence = output.greedy_confidence
next_canvas = proposal.clone()
next_confidence = proposal_confidence.float()
next_latent = replace(
output.next_latent_state,
confidence=next_confidence.detach().float(),
entropy=token_entropy.detach().float(),
)
remaining = torch.tensor(
max_new_tokens, device=device, dtype=torch.long
).sub(committed)
remaining_canvas = remaining[:, None].gt(canvas_positions)
next_latent = self.latent_deliberation.observe_state(
next_latent, output.heavy_hidden_state, output.working_state,
remaining_canvas, state.head,
)
next_state = ModilifyMk2RollingState(
canvas=next_canvas, confidence=next_confidence,
entropy=token_entropy,
latent_state=next_latent,
head=state.head,
)
normal_failure_rate = fused_commit_failure_rate(
proposal_confidence, token_entropy,
entropy_weight=self.config.commit_entropy_weight,
confidence_power=self.config.commit_confidence_power,
top_k=getattr(self.config, "commit_top_k", None),
min_p=getattr(self.config, "commit_min_p", None),
target_confidence=getattr(self.config, "commit_target_confidence", None),
failure_budget=float(self.config.commit_failure_budget),
)
jump_failure_rate = fused_commit_failure_rate(
greedy_confidence, token_entropy,
entropy_weight=self.config.commit_entropy_weight,
confidence_power=self.config.commit_confidence_power,
top_k=getattr(self.config, "commit_top_k", None),
min_p=getattr(self.config, "commit_min_p", None),
target_confidence=getattr(self.config, "commit_target_confidence", None),
failure_budget=float(self.config.commit_failure_budget),
)
previous_failure_rate = fused_commit_failure_rate(
state.confidence, state.entropy,
entropy_weight=self.config.commit_entropy_weight,
confidence_power=self.config.commit_confidence_power,
top_k=getattr(self.config, "commit_top_k", None),
min_p=getattr(self.config, "commit_min_p", None),
target_confidence=getattr(self.config, "commit_target_confidence", None),
failure_budget=float(self.config.commit_failure_budget),
)
policy_decision = select_commit_lengths(
sampled_token_ids=proposal,
normal_failure_rate=normal_failure_rate,
previous_failure_rate=previous_failure_rate,
greedy_token_ids=greedy_proposal,
jump_failure_rate=jump_failure_rate,
ponder_steps=state.latent_state.ponder_steps,
stagnation_steps=state.latent_state.stagnation_steps,
active_rows=active_rows,
remaining_lengths=remaining,
failure_budget=self.config.commit_failure_budget,
stop_token_id=stop_token_ids,
max_ponder_steps=generation_config.max_ponder_steps,
stagnation_threshold=generation_config.jump_on_no_progress_after,
min_progress=generation_config.min_trajectory_progress,
)
next_ponder = policy_decision.ponder_steps
next_stagnation = policy_decision.stagnation_steps
commit_lengths = policy_decision.commit_lengths
jump_rows = policy_decision.jump_rows
jumps += jump_rows.long()
forced_jump_tokens += torch.where(
jump_rows, commit_lengths, torch.zeros_like(commit_lengths)
)
commit_positions = canvas_positions.lt(commit_lengths[:, None])
if bool(jump_rows.any()):
next_state = replace(
next_state,
canvas=torch.where(
commit_positions & jump_rows[:, None],
policy_decision.commit_token_ids,
next_state.canvas,
),
)
next_state = replace(
next_state,
latent_state=replace(
next_state.latent_state,
ponder_steps=next_ponder,
stagnation_steps=next_stagnation,
),
)
commit_token_ids = policy_decision.commit_token_ids
before = committed.clone()
write_rows = torch.arange(batch_size, device=device)[:, None].expand_as(
commit_token_ids
)
write_positions = before[:, None] + canvas_positions
generated[
write_rows[commit_positions], write_positions[commit_positions]
] = commit_token_ids[commit_positions]
if repetition_history is not None:
_add_repetition_history(
repetition_history,
commit_token_ids,
commit_positions,
excluded_repetition_token_ids,
)
commit_width = int(commit_lengths.max())
if commit_width:
block_mask = torch.arange(commit_width, device=device)[None, :].lt(
commit_lengths[:, None]
)
committed_block = torch.where(
block_mask,
commit_token_ids[:, :commit_width],
torch.full(
(batch_size, commit_width),
pad_token_id,
device=device,
dtype=input_ids.dtype,
),
)
block_positions = (
logical_lengths[:, None]
+ torch.arange(commit_width, device=device)[None, :]
).to(torch.int32)
block_positions = torch.where(
block_mask, block_positions, torch.zeros_like(block_positions)
)
cache_attention_mask = torch.cat(
(cache_attention_mask, block_mask), dim=-1
)
past_key_values = self.model.encoder(
input_ids=committed_block,
attention_mask=cache_attention_mask,
past_key_values=past_key_values,
position_ids=block_positions,
).past_key_values
if streamer is not None:
streamer.put(committed_block.cpu())
committed_rows = commit_lengths.gt(0)
shifts += committed_rows.long()
if (
output.working_state is None
):
raise RuntimeError("Forward did not return working trajectory features.")
next_state = self._write_committed_memory(
next_state=next_state,
working_state=output.working_state,
heavy_hidden=output.heavy_hidden_state,
commit_token_ids=commit_token_ids,
commit_lengths=commit_lengths,
prefix_lengths=logical_lengths,
commit_reason=infer_commit_reason(
commit_lengths,
jump_rows=jump_rows,
commit_token_ids=commit_token_ids,
terminal_token_ids=stop_token_ids,
),
)
shifted = self._shift_state_rows(
next_state,
commit_lengths,
sampler,
generators=sampling_generators,
remaining_lengths=remaining - commit_lengths,
pad_token_id=pad_token_id,
)
state = shifted
committed += commit_lengths
logical_lengths += commit_lengths
turn_hits = (
commit_token_ids.eq(turn_end) & commit_positions
).any(dim=-1)
eos_hits = torch.zeros_like(turn_hits)
for token_id in stop_token_ids:
if token_id != turn_end:
eos_hits |= (
commit_token_ids.eq(token_id) & commit_positions
).any(dim=-1)
stop_codes = torch.where(
stop_codes.eq(0) & turn_hits,
torch.ones_like(stop_codes),
stop_codes,
)
stop_codes = torch.where(
stop_codes.eq(0) & eos_hits,
torch.full_like(stop_codes, 2),
stop_codes,
)
stop_codes = torch.where(
stop_codes.eq(0) & committed.ge(max_new_tokens),
torch.full_like(stop_codes, 3),
stop_codes,
)
if generation_config.max_denoising_steps is not None:
stop_codes = torch.where(
stop_codes.eq(0)
& denoise_steps.ge(generation_config.max_denoising_steps),
torch.full_like(stop_codes, 4),
stop_codes,
)
stop_codes = torch.where(
stop_codes.eq(0) & denoise_steps.ge(max_iterations),
torch.full_like(stop_codes, 5),
stop_codes,
)
active_rows = stop_codes.eq(0)
output_width = int(committed.max())
sequences = torch.cat((input_ids, generated[:, :output_width]), dim=-1)
if streamer is not None:
streamer.end()
reason_names = {
1: "turn_end",
2: "eos",
3: "max_new_tokens",
4: "max_denoising_steps",
5: "episode_watchdog",
}
stop_reasons = tuple(
reason_names.get(code, "unknown") for code in stop_codes.detach().cpu().tolist()
)
tokens_per_forward = committed.float() / denoise_steps.clamp_min(1).float()
average_commit_len = committed.float() / shifts.clamp_min(1).float()
def scalar_or_tensor(value: torch.Tensor, *, floating: bool = False):
if batch_size > 1:
return value
item = value[0].item()
return float(item) if floating else int(item)
return ModilifyMk2GenerationOutput(
sequences=sequences,
generated_lengths=committed.clone(),
tokens_per_forward=tokens_per_forward,
past_key_values=past_key_values,
stop_reason=stop_reasons[0] if batch_size == 1 else stop_reasons,
committed_tokens=scalar_or_tensor(committed),
denoise_steps=scalar_or_tensor(denoise_steps),
no_progress_steps=scalar_or_tensor(state.latent_state.stagnation_steps),
jump_count=scalar_or_tensor(jumps),
forced_jump_bad_count=scalar_or_tensor(forced_jump_tokens),
average_commit_len=scalar_or_tensor(average_commit_len, floating=True),
state_shift_count=scalar_or_tensor(shifts),
)