"""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), )