Modilify-Mk2-preview-mlx / modilify_mk2 /configuration_modilify_mk2.py
ydy9038074's picture
Publish Modilify Mk2 Preview MLX
e4f7326 verified
Raw History Blame Contribute Delete
22.6 kB
"""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('<turn|>')
if turn_id is None or turn_id == getattr(tokenizer, 'unk_token_id', None):
raise ValueError('Tokenizer must define <turn|>.')
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