modilify/Modilify-Mk1-preview
120
1# Copyright 2026 Modilify2# SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.03"""Configuration classes for Modilify Mk1."""4 5from __future__ import annotations6 7import math8from typing import Any9 10from transformers.models.diffusion_gemma import (11 DiffusionGemmaConfig,12 DiffusionGemmaTextConfig,13)14 15 16class ModilifyMk1TextConfig(DiffusionGemmaTextConfig):17 """Text configuration for the Modilify Mk1 decoder.18 19 This class preserves the standard DiffusionGemma text schema while giving20 the exported model an independent, stable model type.21 """22 23 model_type = "modilify_mk1_text"24 25 26class ModilifyMk1Config(DiffusionGemmaConfig):27 """Serializable multimodal inference configuration for Modilify Mk1.28 29 Args:30 text_config: DiffusionGemma text configuration or its serialized form.31 vision_config: Gemma 4 vision configuration or its serialized form.32 denoise_temperature: Sampling temperature used at every denoising step.33 commit_failure_budget: Maximum cumulative failure risk for normal commits.34 fused_entropy_weight: Multiplicative entropy penalty coefficient.35 jump_failure_budget: Maximum cumulative failure risk for forced jumps.36 vocab_chunk_size: Vocabulary projection planning size recorded with the37 model. Inference uses standard PyTorch tensor operations.38 latent_dim: Width of the recurrent latent state.39 latent_memory_slots: Number of persistent latent memory slots.40 latent_num_layers: Number of latent Transformer blocks.41 latent_num_heads: Number of latent attention heads.42 latent_local_attention_window: Local token-attention radius.43 latent_dropout: Latent Transformer dropout probability.44 jump_on_no_progress_after: Stagnation steps before a forced jump.45 max_ponder_steps: Maximum denoising iterations per requested token.46 min_trajectory_progress: Minimum fused-risk improvement counted as progress.47 turn_end_token_id: Native Gemma turn terminator.48 kwargs: Standard DiffusionGemma configuration values.49 """50 51 model_type = "modilify_mk1"52 sub_configs = {53 "text_config": ModilifyMk1TextConfig,54 **{55 key: value56 for key, value in DiffusionGemmaConfig.sub_configs.items()57 if key != "text_config"58 },59 }60 61 def __init__(62 self,63 text_config: (64 ModilifyMk1TextConfig65 | DiffusionGemmaTextConfig66 | dict[str, Any]67 | None68 ) = None,69 vision_config: Any | dict[str, Any] | None = None,70 *,71 denoise_temperature: float = 0.8,72 commit_failure_budget: float = 0.2,73 fused_entropy_weight: float = 0.5,74 jump_failure_budget: float = 2.0,75 vocab_chunk_size: int = 65_536,76 latent_dim: int = 1536,77 latent_memory_slots: int = 64,78 latent_num_layers: int = 4,79 latent_num_heads: int = 16,80 latent_local_attention_window: int = 128,81 latent_dropout: float = 0.0,82 jump_on_no_progress_after: int = 12,83 max_ponder_steps: int = 64,84 min_trajectory_progress: float = 0.005,85 turn_end_token_id: int = 106,86 **kwargs: Any,87 ) -> None:88 kwargs.pop("model_type", None)89 if isinstance(text_config, DiffusionGemmaTextConfig):90 text_payload = text_config.to_dict()91 text_payload.pop("model_type", None)92 text_config = ModilifyMk1TextConfig(**text_payload)93 elif isinstance(text_config, dict):94 text_payload = dict(text_config)95 text_payload.pop("model_type", None)96 text_config = ModilifyMk1TextConfig(**text_payload)97 elif text_config is None:98 text_config = ModilifyMk1TextConfig()99 100 self.denoise_temperature = float(denoise_temperature)101 self.commit_failure_budget = float(commit_failure_budget)102 self.fused_entropy_weight = float(fused_entropy_weight)103 self.jump_failure_budget = float(jump_failure_budget)104 self.vocab_chunk_size = int(vocab_chunk_size)105 self.latent_dim = int(latent_dim)106 self.latent_memory_slots = int(latent_memory_slots)107 self.latent_num_layers = int(latent_num_layers)108 self.latent_num_heads = int(latent_num_heads)109 self.latent_local_attention_window = int(latent_local_attention_window)110 self.latent_dropout = float(latent_dropout)111 self.jump_on_no_progress_after = int(jump_on_no_progress_after)112 self.max_ponder_steps = int(max_ponder_steps)113 self.min_trajectory_progress = float(min_trajectory_progress)114 self.turn_end_token_id = int(turn_end_token_id)115 super().__init__(116 text_config=text_config,117 vision_config=vision_config,118 **kwargs,119 )120 self.model_type = type(self).model_type121 if not hasattr(self, "eos_token_id"):122 self.eos_token_id = self.text_config.eos_token_id123 if not hasattr(self, "pad_token_id"):124 self.pad_token_id = self.text_config.pad_token_id125 if not hasattr(self, "bos_token_id"):126 self.bos_token_id = self.text_config.bos_token_id127 self._validate_modilify()128 129 def _validate_modilify(self) -> None:130 """Validate inference-only extension values."""131 132 policy_values = (133 self.denoise_temperature,134 self.commit_failure_budget,135 self.fused_entropy_weight,136 self.jump_failure_budget,137 self.min_trajectory_progress,138 )139 if any(not math.isfinite(value) for value in policy_values):140 raise ValueError("Modilify Mk1 policy values must be finite.")141 positive = (142 self.denoise_temperature,143 self.commit_failure_budget,144 self.jump_failure_budget,145 self.vocab_chunk_size,146 self.latent_dim,147 self.latent_memory_slots,148 self.latent_num_layers,149 self.latent_num_heads,150 self.latent_local_attention_window,151 self.jump_on_no_progress_after,152 self.max_ponder_steps,153 )154 if any(value <= 0 for value in positive):155 raise ValueError(156 "Modilify Mk1 dimensions, budgets, and intervals must be positive."157 )158 if self.fused_entropy_weight < 0:159 raise ValueError("`fused_entropy_weight` must be non-negative.")160 if self.latent_dim % self.latent_num_heads:161 raise ValueError("`latent_dim` must be divisible by `latent_num_heads`.")162 if not 0.0 <= self.latent_dropout < 1.0:163 raise ValueError("`latent_dropout` must be in [0, 1).")164 if self.min_trajectory_progress < 0:165 raise ValueError("`min_trajectory_progress` must be non-negative.")166 167 168__all__ = ["ModilifyMk1Config", "ModilifyMk1TextConfig"]169 