"""Strict text-only configuration for the schema25 GDN2 protocol.""" from __future__ import annotations import json import math from pathlib import Path from collections.abc import Mapping, Sequence from typing import Any from transformers.configuration_utils import PreTrainedConfig from transformers.models.diffusion_gemma import DiffusionGemmaTextConfig STATE_SCHEMA_VERSION = 25 # Keep topology stable; CE scope and normalization have separate objective tags. TRAINING_SCHEME = ( "gold_prefix_shared_commit_0_256_committed_ce_calibration_" "causal_throughput_terminal_sft" ) MEMORY_SCHEME = "dual_timescale_gdn2_trajectory_memory" MEMORY_ARCHITECTURE = "compact_gdn2_v2" CONFIG_PROTOCOL_ERROR = "ModilifyMk2 requires schema25 compact_gdn2_v2; older memory topologies require a new run." HISTORY_VIEWS = 4 COMMIT_SEQUENCE_LAYERS = 2 DENOISE_TEMPERATURE = 0.8 VOCAB_CHUNK_SIZE = 32_768 COMMIT_READINESS_OBJECTIVE = "frontier_prefix_budget_v2" TOKEN_CE_SUPERVISION = "valid_canvas_v1" TOKEN_CE_NORMALIZATION = "per_sample_exposure_v1" def require_current_checkpoint_protocol(metadata: Mapping[str, Any]) -> None: """Reject checkpoints that do not implement the GDN2 state topology.""" if ( metadata.get("state_schema_version") != STATE_SCHEMA_VERSION or metadata.get("training_scheme") != TRAINING_SCHEME or metadata.get("memory_scheme") != MEMORY_SCHEME or metadata.get("memory_architecture") != MEMORY_ARCHITECTURE ): raise RuntimeError( f"ModilifyMk2 checkpoint requires schema{STATE_SCHEMA_VERSION} {MEMORY_ARCHITECTURE}; " "older memory topologies cannot be restored. Start from the base model." ) class ModilifyMk2TextConfig(DiffusionGemmaTextConfig): model_type = "modilify_mk2_text" vocab_size: int = 262_144 hidden_size: int = 2816 intermediate_size: int = 2112 num_hidden_layers: int = 30 num_attention_heads: int = 16 num_key_value_heads: int = 8 head_dim: int = 256 max_position_embeddings: int = 262_144 sliding_window: int = 1024 use_bidirectional_attention: str | None = None num_global_key_value_heads: int | None = 2 global_head_dim: int = 512 num_experts: int | None = 128 top_k_experts: int | None = 8 moe_intermediate_size: int | None = 704 class ModilifyMk2Config(PreTrainedConfig): """Configuration with recurrent GDN2 trajectory memory.""" model_type = "modilify_mk2" sub_configs = {"text_config": ModilifyMk2TextConfig} def __init__( self, text_config: ModilifyMk2TextConfig | dict[str, Any] | None = None, *, canvas_length: int = 256, initializer_range: float = 0.02, tie_word_embeddings: bool = True, state_schema_version: int = STATE_SCHEMA_VERSION, training_scheme: str = TRAINING_SCHEME, memory_scheme: str = MEMORY_SCHEME, memory_architecture: str = MEMORY_ARCHITECTURE, latent_dim: int = 2816, latent_ffn_dim: int = 7168, latent_memory_slots: int = 256, latent_num_layers: int = 4, latent_num_heads: int = 16, latent_local_attention_window: int = 128, latent_history_length: int = 16, latent_tape_probes: int = 1, latent_tape_scheme: str = "gdn2_spatial_probe_v1", latent_history_views: int = HISTORY_VIEWS, latent_history_kv_rank: int | None = None, latent_working_last_block_global: bool = True, working_memory_bus: bool = True, persistent_memory_bus: bool = True, persistent_memory_write: str = "commit_only_transformer", commit_sequence_layers: int = COMMIT_SEQUENCE_LAYERS, commit_sequence_dim: int | None = None, writer_slot_gate: str = "per_slot", latent_working_bus_unfreeze_steps: int = 0, latent_persistent_bus_unfreeze_steps: int = 0, training_bptt_steps: int = 16, kv_cache_bucket_size: int = 128, turn_end_token_id: int = 106, terminal_token_ids: Sequence[int] | None = None, channel_end_token_id: int = 101, eos_token_id: int = 1, commit_failure_budget: float = 0.2, commit_top_k: int | None = 40, commit_min_p: float | None = 0.05, commit_target_confidence: float | None = 0.5, commit_entropy_weight: float = 1.0, commit_confidence_power: float = 1.0, commit_gold_alpha: float = 0.4, commit_gold_weight: float = 1.0, commit_readiness_failure_weight: float = 1.0, token_loss_weight: float = 1.0, token_ce_supervision: str = TOKEN_CE_SUPERVISION, token_ce_normalization: str = TOKEN_CE_NORMALIZATION, confidence_calibration_loss_weight: float = 0.1, commit_readiness_objective: str = COMMIT_READINESS_OBJECTIVE, commit_readiness_target_tokens: int = 16, commit_readiness_loss_weight: float = 0.1, commit_readiness_budget_margin: float = 0.02, commit_readiness_beta: float = 0.02, terminal_stop_loss_weight: float = 1.0, terminal_stop_target_probability: float = 0.95, **kwargs: Any, ) -> None: if any(key.startswith("commit_throughput_") for key in kwargs): raise ValueError("Removed commit-throughput configuration fields.") if token_ce_supervision != TOKEN_CE_SUPERVISION: raise ValueError("Schema25 requires valid-canvas token CE.") if state_schema_version != STATE_SCHEMA_VERSION: raise RuntimeError(CONFIG_PROTOCOL_ERROR) if training_scheme != TRAINING_SCHEME or memory_scheme != MEMORY_SCHEME: raise RuntimeError(CONFIG_PROTOCOL_ERROR) if memory_architecture != MEMORY_ARCHITECTURE: raise RuntimeError("Schema25 requires compact_gdn2_v2 memory; start a new run.") if text_config is None: text_config = ModilifyMk2TextConfig() elif isinstance(text_config, dict): text_config = dict(text_config) text_config.pop("model_type", None) text_config["use_bidirectional_attention"] = None text_config = ModilifyMk2TextConfig(**text_config) elif not isinstance(text_config, ModilifyMk2TextConfig): payload = text_config.to_dict() payload.pop("model_type", None) text_config = ModilifyMk2TextConfig(**payload) self.text_config = text_config self.canvas_length = canvas_length self.initializer_range = initializer_range self.state_schema_version = STATE_SCHEMA_VERSION self.training_scheme = training_scheme self.memory_scheme = memory_scheme self.memory_architecture = memory_architecture self.latent_dim = latent_dim self.latent_ffn_dim = latent_ffn_dim self.latent_memory_slots = latent_memory_slots self.latent_num_layers = latent_num_layers self.latent_num_heads = latent_num_heads self.latent_local_attention_window = latent_local_attention_window self.latent_history_length = latent_history_length self.latent_tape_probes = int(latent_tape_probes) self.latent_tape_scheme = str(latent_tape_scheme) self.latent_history_views = int(latent_history_views) if latent_history_kv_rank is None: rank = min(1024, latent_dim) rank -= rank % max(latent_num_heads, 1) if rank <= 0: rank = latent_num_heads self.latent_history_kv_rank = rank else: self.latent_history_kv_rank = latent_history_kv_rank self.latent_working_last_block_global = bool(latent_working_last_block_global) self.working_memory_bus = bool(working_memory_bus) self.persistent_memory_bus = bool(persistent_memory_bus) self.persistent_memory_write = str(persistent_memory_write) self.commit_sequence_layers = int(commit_sequence_layers) if commit_sequence_dim is None: self.commit_sequence_dim = self.latent_history_kv_rank else: self.commit_sequence_dim = int(commit_sequence_dim) self.writer_slot_gate = str(writer_slot_gate) self.latent_working_bus_unfreeze_steps = int(latent_working_bus_unfreeze_steps) self.latent_persistent_bus_unfreeze_steps = int( latent_persistent_bus_unfreeze_steps ) self.training_bptt_steps = training_bptt_steps self.kv_cache_bucket_size = kv_cache_bucket_size self.turn_end_token_id = turn_end_token_id if terminal_token_ids is None: self.terminal_token_ids = (int(turn_end_token_id),) else: self.terminal_token_ids = tuple(int(token_id) for token_id in terminal_token_ids) self.channel_end_token_id = int(channel_end_token_id) self.commit_failure_budget = float(commit_failure_budget) self.commit_top_k = int(commit_top_k) if commit_top_k is not None else None self.commit_min_p = float(commit_min_p) if commit_min_p is not None else None self.commit_target_confidence = float(commit_target_confidence) if commit_target_confidence is not None else None self.commit_entropy_weight = float(commit_entropy_weight) self.commit_confidence_power = float(commit_confidence_power) self.commit_gold_alpha = float(commit_gold_alpha) self.commit_gold_weight = float(commit_gold_weight) self.commit_readiness_failure_weight = float(commit_readiness_failure_weight) self.token_loss_weight = token_loss_weight self.token_ce_supervision = str(token_ce_supervision) self.token_ce_normalization = str(token_ce_normalization) self.confidence_calibration_loss_weight = confidence_calibration_loss_weight self.commit_readiness_objective = str(commit_readiness_objective) self.commit_readiness_target_tokens = int(commit_readiness_target_tokens) self.commit_readiness_loss_weight = float(commit_readiness_loss_weight) self.commit_readiness_budget_margin = float(commit_readiness_budget_margin) self.commit_readiness_beta = float(commit_readiness_beta) self.terminal_stop_loss_weight = terminal_stop_loss_weight self.terminal_stop_target_probability = terminal_stop_target_probability self.vocab_chunk_size = VOCAB_CHUNK_SIZE super().__init__( tie_word_embeddings=tie_word_embeddings, eos_token_id=eos_token_id, **kwargs, ) self._validate() def _validate(self) -> None: positive = ( self.canvas_length, self.latent_dim, self.latent_ffn_dim, self.latent_memory_slots, self.latent_num_layers, self.latent_num_heads, self.latent_local_attention_window, self.latent_history_length, self.latent_tape_probes, self.latent_history_kv_rank, self.commit_sequence_layers, self.commit_sequence_dim, self.training_bptt_steps, self.kv_cache_bucket_size, ) if any(value <= 0 for value in positive): raise ValueError("All schema25 dimensions and intervals must be positive.") if self.canvas_length != 256: raise ValueError("Schema25 requires a 256-token canvas.") if self.latent_tape_scheme != "gdn2_spatial_probe_v1": raise ValueError("Schema25 requires GDN2 spatial probes.") if self.latent_tape_probes > self.canvas_length: raise ValueError("Spatial probes cannot exceed canvas positions.") if self.commit_sequence_dim % self.latent_num_heads: raise ValueError("`commit_sequence_dim` must be divisible by `latent_num_heads`.") if self.latent_working_bus_unfreeze_steps < 0: raise ValueError("`latent_working_bus_unfreeze_steps` must be non-negative.") if self.latent_persistent_bus_unfreeze_steps < 0: raise ValueError( "`latent_persistent_bus_unfreeze_steps` must be non-negative." ) if self.latent_persistent_bus_unfreeze_steps < self.latent_working_bus_unfreeze_steps: raise ValueError( "Persistent bus must unfreeze no earlier than the working bus." ) if self.latent_history_kv_rank > self.latent_dim: raise ValueError("`latent_history_kv_rank` must not exceed `latent_dim`.") if self.latent_history_kv_rank % self.latent_num_heads: raise ValueError("`latent_history_kv_rank` must be divisible by `latent_num_heads`.") if not isinstance(self.eos_token_id, int) or self.eos_token_id < 0: raise ValueError("ModilifyMk2 requires one non-negative integer EOS token ID.") if not isinstance(self.channel_end_token_id, int) or self.channel_end_token_id < 0: raise ValueError("`channel_end_token_id` must be a non-negative integer.") if not self.terminal_token_ids: raise ValueError("`terminal_token_ids` must not be empty.") if any( not isinstance(token_id, int) or token_id < 0 for token_id in self.terminal_token_ids ): raise ValueError("`terminal_token_ids` must be non-negative integers.") if self.latent_dim % self.latent_num_heads: raise ValueError("`latent_dim` must be divisible by `latent_num_heads`.") if self.commit_readiness_objective != COMMIT_READINESS_OBJECTIVE: raise ValueError( "Unsupported commit-readiness objective: " f"{self.commit_readiness_objective!r}." ) if not all(math.isfinite(v) for v in ( self.commit_failure_budget, self.commit_entropy_weight, self.commit_confidence_power, self.commit_gold_alpha, self.commit_gold_weight, self.commit_readiness_failure_weight, self.commit_readiness_loss_weight, )): raise ValueError("Commit policy parameters must be finite.") if self.commit_failure_budget <= 0 or self.commit_entropy_weight < 0: raise ValueError("Commit budget must be positive and entropy weight non-negative.") if self.commit_top_k is not None and self.commit_top_k <= 0: raise ValueError("`commit_top_k` must be a positive integer.") if self.commit_min_p is not None and not 0.0 < self.commit_min_p < 1.0: raise ValueError("`commit_min_p` must be in (0, 1).") if self.commit_target_confidence is not None and not 0.0 < self.commit_target_confidence < 1.0: raise ValueError("`commit_target_confidence` must be in (0, 1).") if self.commit_confidence_power <= 0 or not 0 < self.commit_gold_alpha < 1: raise ValueError("Commit power must be positive and gold alpha in (0, 1).") if self.commit_gold_weight <= 0: raise ValueError("Commit gold weight must be positive.") if self.commit_readiness_failure_weight < 0: raise ValueError("Commit readiness failure weight must be non-negative.") if self.commit_readiness_target_tokens <= 0: raise ValueError("`commit_readiness_target_tokens` must be positive.") if self.commit_readiness_loss_weight < 0: raise ValueError("`commit_readiness_loss_weight` must be non-negative.") if not 0.0 <= self.commit_readiness_budget_margin < self.commit_failure_budget: raise ValueError( "`commit_readiness_budget_margin` must be in [0, commit_failure_budget)." ) if self.commit_readiness_beta <= 0: raise ValueError("`commit_readiness_beta` must be positive.") if self.terminal_stop_loss_weight < 0: raise ValueError("`terminal_stop_loss_weight` must be non-negative.") if not 0.0 < self.terminal_stop_target_probability < 1.0: raise ValueError("`terminal_stop_target_probability` must be in (0, 1).") loss_weights = ( self.token_loss_weight, self.confidence_calibration_loss_weight, self.commit_readiness_loss_weight, self.terminal_stop_loss_weight, ) if any(not math.isfinite(weight) or weight < 0 for weight in loss_weights): raise ValueError("Schema25 fixed loss weights must be finite and non-negative.") if self.token_ce_normalization != TOKEN_CE_NORMALIZATION: raise ValueError("Unsupported token CE normalization.") @classmethod def from_dict(cls, config_dict: dict[str, Any], **kwargs: Any) -> "ModilifyMk2Config": return_unused_kwargs = kwargs.pop("return_unused_kwargs", False) payload = dict(config_dict) if (payload.get("state_schema_version") != STATE_SCHEMA_VERSION or payload.get("memory_architecture") != MEMORY_ARCHITECTURE): raise RuntimeError(CONFIG_PROTOCOL_ERROR) payload.pop("model_type", None) payload.pop("architectures", None) for key in tuple(kwargs): if key in payload: payload[key] = kwargs.pop(key) config = cls(**payload) for key in tuple(kwargs): if hasattr(config, key) or key == "name_or_path": setattr(config, key, kwargs.pop(key)) return (config, kwargs) if return_unused_kwargs else config __all__ = [ "ModilifyMk2Config", "ModilifyMk2TextConfig", "CONFIG_PROTOCOL_ERROR", "COMMIT_READINESS_OBJECTIVE", "COMMIT_SEQUENCE_LAYERS", "DENOISE_TEMPERATURE", "HISTORY_VIEWS", "MEMORY_SCHEME", "MEMORY_ARCHITECTURE", "STATE_SCHEMA_VERSION", "TRAINING_SCHEME", "VOCAB_CHUNK_SIZE", "TOKEN_CE_SUPERVISION", "TOKEN_CE_NORMALIZATION", "require_current_checkpoint_protocol", ] class ModilifyMk2GenerationConfig: """Native settings; no autoregressive sampler or Torch runtime is needed.""" def __init__(self, **kwargs: Any): unsupported = ( 'sampler_config', 'stability_threshold', 'confidence_threshold', 'one_token_per_denoise_step', 'compile_generation', 'sliding_denoise', 'adaptive_ponder_budget', 'force_commit_on_max_steps', 'ponder_budget_id', ) configured = [name for name in unsupported if kwargs.pop(name, None) not in (None, False)] if configured: raise ValueError(f'Unsupported generation fields: {configured}') defaults = { 'max_new_tokens': None, 'max_denoising_steps': None, 'bos_token_id': None, 'pad_token_id': None, 'eos_token_id': None, 'turn_end_token_id': None, 'max_ponder_steps': 64, 'jump_on_no_progress_after': 12, 'min_trajectory_progress': 0.005, 'repetition_penalty': 1.0, 'repetition_penalty_exclude_token_ids': [], } for name, default in defaults.items(): setattr(self, name, kwargs.pop(name, default)) self.repetition_penalty_exclude_token_ids = list(dict.fromkeys( int(value) for value in self.repetition_penalty_exclude_token_ids or () )) def validate(self) -> None: for name in ('max_new_tokens', 'max_denoising_steps', 'max_ponder_steps', 'jump_on_no_progress_after'): value = getattr(self, name) if value is not None and (isinstance(value, bool) or not isinstance(value, int) or value <= 0): raise ValueError(f'{name} must be a positive integer.') if not math.isfinite(self.repetition_penalty) or self.repetition_penalty <= 0: raise ValueError('repetition_penalty must be finite and positive.') if not math.isfinite(self.min_trajectory_progress) or self.min_trajectory_progress < 0: raise ValueError('min_trajectory_progress must be finite and nonnegative.') if self.turn_end_token_id is not None and self.turn_end_token_id < 0: raise ValueError('turn_end_token_id must be nonnegative.') if any(isinstance(value, bool) or not isinstance(value, int) or value < 0 for value in self.repetition_penalty_exclude_token_ids): raise ValueError('Excluded token IDs must be nonnegative integers.') def load_generation_config(model_dir: str | Path) -> ModilifyMk2GenerationConfig: path = Path(model_dir) / 'generation_config.json' if not path.is_file(): raise FileNotFoundError(f'Missing generation configuration: {path}') return ModilifyMk2GenerationConfig(**json.loads(path.read_text(encoding='utf-8'))) def configure_generation_config( generation_config: ModilifyMk2GenerationConfig, processor: Any, *, max_new_tokens: int, max_denoising_steps: int | None, repetition_penalty: float | None = None, ) -> ModilifyMk2GenerationConfig: if generation_config.max_denoising_steps is None: generation_config.max_denoising_steps = 48 generation_config.max_new_tokens = max_new_tokens if max_denoising_steps is not None: generation_config.max_denoising_steps = max_denoising_steps if repetition_penalty is not None: generation_config.repetition_penalty = repetition_penalty tokenizer = getattr(processor, 'tokenizer', processor) if generation_config.bos_token_id is None: generation_config.bos_token_id = tokenizer.bos_token_id if generation_config.pad_token_id is None: generation_config.pad_token_id = tokenizer.pad_token_id turn_id = tokenizer.convert_tokens_to_ids('') if turn_id is None or turn_id == getattr(tokenizer, 'unk_token_id', None): raise ValueError('Tokenizer must define .') generation_config.turn_end_token_id = int(turn_id) if generation_config.eos_token_id is None: generation_config.eos_token_id = [value for value in [tokenizer.eos_token_id] if value is not None] tool_id = tokenizer.convert_tokens_to_ids('<|tool_response>') if tool_id is not None and tool_id != getattr(tokenizer, 'unk_token_id', None): eos = generation_config.eos_token_id eos = [eos] if isinstance(eos, int) else list(eos or ()) if int(tool_id) not in eos: generation_config.eos_token_id = [*eos, int(tool_id)] generation_config.repetition_penalty_exclude_token_ids = list(dict.fromkeys( int(value) for value in [*generation_config.repetition_penalty_exclude_token_ids, *(getattr(tokenizer, 'all_special_ids', None) or ())] if value is not None )) generation_config.validate() return generation_config