modilify/Modilify-Mk1-preview
120
1# Copyright 2026 Modilify2# SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.03"""Standard PyTorch multimodal model implementation for Modilify Mk1."""4 5from __future__ import annotations6 7from dataclasses import dataclass, replace8import math9from typing import Any10 11import torch12from torch import nn13from transformers.cache_utils import Cache14from transformers.modeling_outputs import BaseModelOutputWithPast15from transformers.utils import ModelOutput16from transformers.models.diffusion_gemma import (17 DiffusionGemmaDecoderModel,18 DiffusionGemmaEncoderModel,19 DiffusionGemmaPreTrainedModel,20)21 22from .configuration_modilify_mk1 import ModilifyMk1Config23from .generation_modilify_mk1 import (24 ModilifyMk1GenerationConfig,25 ModilifyMk1GenerationMixin,26)27from .latent_deliberation import (28 LatentDeliberationState,29 LatentDeliberationTransformer,30)31 32 33@dataclass34class ModilifyMk1DecoderOutput(BaseModelOutputWithPast):35 """Decoder hidden states and latent-context diagnostics."""36 37 token_embeddings: torch.FloatTensor | None = None38 latent_residual_diagnostics: dict[str, torch.Tensor] | None = None39 40 41@dataclass42class ModilifyMk1ModelOutput(BaseModelOutputWithPast):43 """Combined multimodal encoder and diffusion decoder output."""44 45 token_embeddings: torch.FloatTensor | None = None46 encoder_last_hidden_state: torch.FloatTensor | None = None47 latent_residual_diagnostics: dict[str, torch.Tensor] | None = None48 49 50@dataclass51class ModilifyMk1BlockDiffusionOutput(ModelOutput):52 """Inference output used by the rolling diffusion generator."""53 54 logits: torch.FloatTensor | None = None55 heavy_hidden_state: torch.FloatTensor | None = None56 next_latent_state: LatentDeliberationState | None = None57 past_key_values: Cache | None = None58 encoder_last_hidden_state: torch.FloatTensor | None = None59 temporal_context: torch.FloatTensor | None = None60 latent_residual_diagnostics: dict[str, torch.Tensor] | None = None61 proposal: torch.LongTensor | None = None62 proposal_confidence: torch.FloatTensor | None = None63 token_entropy: torch.FloatTensor | None = None64 greedy_proposal: torch.LongTensor | None = None65 greedy_confidence: torch.FloatTensor | None = None66 67 68class ModilifyMk1EncoderModel(DiffusionGemmaEncoderModel):69 """Unmodified Transformers DiffusionGemma multimodal encoder."""70 71 config_class = ModilifyMk1Config72 73 74class ModilifyMk1DecoderModel(DiffusionGemmaDecoderModel):75 """DiffusionGemma decoder conditioned by recurrent latent embeddings."""76 77 config_class = ModilifyMk1Config78 latent_residual_rms_ratio_cap = 0.579 80 def merge_latent_context(81 self,82 token_embeddings: torch.Tensor,83 latent_context: torch.Tensor | None,84 ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:85 """Apply the native self-conditioning bridge to latent context.86 87 Args:88 token_embeddings: Embedded noisy canvas tokens.89 latent_context: Context emitted by the latent Transformer.90 91 Returns:92 Merged embeddings and scalar diagnostic tensors.93 """94 95 context = (96 torch.zeros_like(token_embeddings)97 if latent_context is None98 else latent_context.to(token_embeddings)99 )100 if context.shape != token_embeddings.shape:101 raise ValueError("Latent context must match the canvas embedding shape.")102 mapper = self.self_conditioning103 normalized = mapper.pre_norm(context)104 mapped = mapper.down_proj(105 mapper.act_fn(mapper.gate_proj(normalized)) * mapper.up_proj(normalized)106 )107 mapped_rms_per_token = mapped.float().square().mean(dim=-1, keepdim=True).sqrt()108 token_rms_per_token = token_embeddings.float().square().mean(dim=-1, keepdim=True).sqrt()109 cap = self.latent_residual_rms_ratio_cap * token_rms_per_token110 scale = cap / torch.sqrt(mapped_rms_per_token.square() + cap.square() + 1.0e-12)111 mapped = mapped * scale.to(mapped)112 combined = mapper.post_norm(token_embeddings + mapped)113 token_rms = token_embeddings.detach().float().square().mean().sqrt()114 mapped_rms = mapped.detach().float().square().mean().sqrt()115 diagnostics = {116 "token_embedding_rms": token_rms,117 "latent_context_rms": context.detach().float().square().mean().sqrt(),118 "mapped_context_rms": mapped_rms,119 "latent_to_embedding_rms_ratio": mapped_rms / token_rms.clamp_min(1.0e-12),120 }121 return combined, diagnostics122 123 def forward(124 self,125 decoder_input_ids: torch.LongTensor,126 past_key_values: Cache | None = None,127 temporal_context_embeddings: torch.FloatTensor | None = None,128 decoder_attention_mask: torch.Tensor | dict | None = None,129 decoder_position_ids: torch.LongTensor | None = None,130 **kwargs: Any,131 ) -> ModilifyMk1DecoderOutput:132 """Decode one noisy canvas using only Transformers and PyTorch operations."""133 134 token_embeddings = self.embed_tokens(decoder_input_ids)135 inputs_embeds, diagnostics = self.merge_latent_context(136 token_embeddings,137 temporal_context_embeddings,138 )139 if decoder_position_ids is None:140 prefix = past_key_values.get_seq_length(0) if past_key_values is not None else 0141 decoder_position_ids = torch.arange(142 prefix,143 prefix + inputs_embeds.shape[1],144 device=inputs_embeds.device,145 ).unsqueeze(0)146 if not isinstance(mask_mapping := decoder_attention_mask, dict):147 mask_mapping = self.create_diffusion_decoder_attention_mask(148 config=self.text_config,149 inputs_embeds=inputs_embeds,150 past_key_values=past_key_values,151 decoder_attention_mask=decoder_attention_mask,152 )153 position_embeddings = {154 layer_type: self.rotary_emb(inputs_embeds, decoder_position_ids, layer_type)155 for layer_type in self.unique_layer_types156 }157 hidden_states = inputs_embeds158 for index, layer in enumerate(self.layers[: self.text_config.num_hidden_layers]):159 layer_type = self.text_config.layer_types[index]160 hidden_states = layer(161 hidden_states,162 position_embeddings=position_embeddings[layer_type],163 attention_mask=mask_mapping[layer_type],164 position_ids=decoder_position_ids,165 past_key_values=past_key_values,166 **kwargs,167 )168 return ModilifyMk1DecoderOutput(169 last_hidden_state=self.norm(hidden_states),170 past_key_values=past_key_values,171 token_embeddings=token_embeddings,172 latent_residual_diagnostics=diagnostics,173 )174 175 176class ModilifyMk1Model(DiffusionGemmaPreTrainedModel):177 """Multimodal encoder plus latent-conditioned block diffusion decoder."""178 179 config_class = ModilifyMk1Config180 _tied_weights_keys = {181 "encoder.language_model.norm.weight": "decoder.norm.weight",182 r"encoder.language_model.layers\.(?:[^.]+\.)*weight": r"decoder.layers\.(?:[^.]+\.)*weight",183 r"encoder.language_model.layers\.(?:[^.]+\.)*scale": r"decoder.layers\.(?:[^.]+\.)*scale",184 (185 r"encoder.language_model.layers\.(?:[^.]+\.)*per_expert_scale"186 ): r"decoder.layers\.(?:[^.]+\.)*per_expert_scale",187 (188 r"encoder.language_model.layers\.(?:[^.]+\.)*gate_up_proj"189 ): r"decoder.layers\.(?:[^.]+\.)*gate_up_proj",190 (191 r"encoder.language_model.layers\.(?:[^.]+\.)*down_proj"192 ): r"decoder.layers\.(?:[^.]+\.)*down_proj",193 "encoder.language_model.embed_tokens.weight": "decoder.embed_tokens.weight",194 }195 196 def __init__(self, config: ModilifyMk1Config) -> None:197 super().__init__(config)198 self.encoder = ModilifyMk1EncoderModel(config)199 self.decoder = ModilifyMk1DecoderModel(config)200 self.post_init()201 202 def get_encoder(self) -> ModilifyMk1EncoderModel:203 """Return the standard multimodal encoder."""204 205 return self.encoder206 207 def get_decoder(self) -> ModilifyMk1DecoderModel:208 """Return the diffusion decoder."""209 210 return self.decoder211 212 def get_input_embeddings(self) -> nn.Module:213 """Return the shared text embedding module."""214 215 return self.encoder.get_input_embeddings()216 217 def set_input_embeddings(self, value: nn.Module) -> None:218 """Set the shared text embedding module."""219 220 self.encoder.set_input_embeddings(value)221 self.decoder.embed_tokens = value222 223 def forward(224 self,225 *,226 input_ids: torch.LongTensor | None = None,227 attention_mask: torch.Tensor | dict | None = None,228 past_key_values: Cache | None = None,229 position_ids: torch.LongTensor | None = None,230 decoder_input_ids: torch.LongTensor,231 temporal_context_embeddings: torch.FloatTensor | None = None,232 decoder_attention_mask: torch.Tensor | dict | None = None,233 decoder_position_ids: torch.LongTensor | None = None,234 **kwargs: Any,235 ) -> ModilifyMk1ModelOutput:236 """Encode multimodal context and decode one canvas."""237 238 encoder_hidden_state = None239 encoder_keys = ("pixel_values", "mm_token_type_ids", "image_position_ids", "inputs_embeds")240 encoder_kwargs = {key: kwargs.pop(key) for key in encoder_keys if key in kwargs}241 if input_ids is not None:242 encoded = self.encoder(243 input_ids=input_ids,244 attention_mask=attention_mask,245 past_key_values=past_key_values,246 position_ids=position_ids,247 **encoder_kwargs,248 )249 past_key_values = encoded.past_key_values250 encoder_hidden_state = encoded.last_hidden_state251 elif past_key_values is None:252 raise ValueError("Either `input_ids` or `past_key_values` is required.")253 decoded = self.decoder(254 decoder_input_ids=decoder_input_ids,255 past_key_values=past_key_values,256 temporal_context_embeddings=temporal_context_embeddings,257 decoder_attention_mask=decoder_attention_mask,258 decoder_position_ids=decoder_position_ids,259 **kwargs,260 )261 return ModilifyMk1ModelOutput(262 last_hidden_state=decoded.last_hidden_state,263 past_key_values=past_key_values,264 token_embeddings=decoded.token_embeddings,265 encoder_last_hidden_state=encoder_hidden_state,266 latent_residual_diagnostics=decoded.latent_residual_diagnostics,267 )268 269 270class ModilifyMk1ForBlockDiffusion(271 DiffusionGemmaPreTrainedModel,272 ModilifyMk1GenerationMixin,273):274 """Inference-only multimodal Modilify Mk1 model."""275 276 config_class = ModilifyMk1Config277 _tied_weights_keys = {"lm_head.weight": "model.decoder.embed_tokens.weight"}278 generation_config_class = ModilifyMk1GenerationConfig279 280 @torch.no_grad()281 def _init_weights(self, module: nn.Module) -> None:282 super()._init_weights(module)283 if isinstance(module, LatentDeliberationTransformer):284 module.reset_memory_slot_identity()285 286 def __init__(self, config: ModilifyMk1Config) -> None:287 super().__init__(config)288 self.model = ModilifyMk1Model(config)289 self.latent_deliberation = LatentDeliberationTransformer(290 hidden_size=config.text_config.hidden_size,291 latent_dim=config.latent_dim,292 memory_slots=config.latent_memory_slots,293 num_layers=config.latent_num_layers,294 num_heads=config.latent_num_heads,295 local_attention_window=config.latent_local_attention_window,296 dropout=config.latent_dropout,297 )298 self.lm_head = nn.Linear(299 config.text_config.hidden_size,300 config.text_config.vocab_size,301 bias=False,302 )303 self.final_logit_softcapping = config.text_config.final_logit_softcapping304 self.post_init()305 306 def _prepare_latent_context(307 self,308 decoder_input_ids: torch.LongTensor,309 *,310 history_hidden_state: torch.Tensor | None,311 confidence: torch.Tensor | None,312 entropy: torch.Tensor | None,313 age: torch.Tensor | None,314 latent_state: LatentDeliberationState | None,315 ) -> tuple[torch.Tensor, LatentDeliberationState]:316 """Advance recurrent latent state for the current canvas."""317 318 batch_size, canvas_length = decoder_input_ids.shape319 dtype = self.model.decoder.embed_tokens.weight.dtype320 if latent_state is None:321 latent_state = LatentDeliberationState.empty(322 batch_size=batch_size,323 canvas_length=canvas_length,324 latent_dim=self.config.latent_dim,325 memory_slots=self.config.latent_memory_slots,326 device=decoder_input_ids.device,327 dtype=dtype,328 )329 confidence = (330 latent_state.confidence331 if confidence is None332 else confidence.squeeze(-1).float()333 )334 entropy = latent_state.entropy if entropy is None else entropy.squeeze(-1).float()335 if age is not None:336 latent_state = replace(337 latent_state,338 age=age.to(device=decoder_input_ids.device, dtype=torch.int32),339 )340 token_embeddings = self.model.decoder.embed_tokens(decoder_input_ids)341 history = (342 torch.zeros_like(token_embeddings)343 if history_hidden_state is None344 else history_hidden_state345 )346 return self.latent_deliberation(347 heavy_hidden=history,348 token_embeddings=token_embeddings,349 confidence=confidence,350 entropy=entropy,351 state=latent_state,352 )353 354 def _proposal_statistics(355 self,356 logits: torch.Tensor,357 *,358 denoise_temperature: float | None = None,359 ) -> tuple[360 torch.LongTensor,361 torch.Tensor,362 torch.Tensor,363 torch.LongTensor,364 torch.Tensor,365 ]:366 """Compute exact proposal statistics with standard PyTorch operations."""367 368 temperature = (369 self.config.denoise_temperature370 if denoise_temperature is None371 else float(denoise_temperature)372 )373 if not math.isfinite(temperature) or temperature <= 0.0:374 raise ValueError("`denoise_temperature` must be positive.")375 scores = logits.float() / temperature376 probabilities = torch.softmax(scores, dim=-1)377 flat = probabilities.reshape(-1, probabilities.shape[-1])378 proposal = torch.multinomial(flat, num_samples=1).view(logits.shape[:-1])379 proposal_confidence = probabilities.gather(380 -1,381 proposal.unsqueeze(-1),382 ).squeeze(-1)383 greedy_proposal = probabilities.argmax(dim=-1)384 greedy_confidence = probabilities.gather(385 -1,386 greedy_proposal.unsqueeze(-1),387 ).squeeze(-1)388 token_entropy = -(389 probabilities * probabilities.clamp_min(1.0e-30).log()390 ).sum(dim=-1)391 return (392 proposal,393 proposal_confidence,394 token_entropy,395 greedy_proposal,396 greedy_confidence,397 )398 399 def forward(400 self,401 *,402 input_ids: torch.LongTensor | None = None,403 attention_mask: torch.Tensor | dict | None = None,404 past_key_values: Cache | None = None,405 position_ids: torch.LongTensor | None = None,406 decoder_input_ids: torch.LongTensor,407 previous_confidence: torch.FloatTensor | None = None,408 previous_entropy: torch.FloatTensor | None = None,409 token_age: torch.Tensor | None = None,410 latent_state: LatentDeliberationState | None = None,411 history_hidden_state: torch.FloatTensor | None = None,412 decoder_attention_mask: torch.Tensor | dict | None = None,413 decoder_position_ids: torch.LongTensor | None = None,414 return_proposal_statistics: bool = False,415 denoise_temperature: float | None = None,416 **kwargs: Any,417 ) -> ModilifyMk1BlockDiffusionOutput:418 """Run one inference step over a noisy diffusion canvas."""419 420 latent_context, next_state = self._prepare_latent_context(421 decoder_input_ids,422 history_hidden_state=history_hidden_state,423 confidence=previous_confidence,424 entropy=previous_entropy,425 age=token_age,426 latent_state=latent_state,427 )428 outputs = self.model(429 input_ids=input_ids,430 attention_mask=attention_mask,431 past_key_values=past_key_values,432 position_ids=position_ids,433 decoder_input_ids=decoder_input_ids,434 temporal_context_embeddings=latent_context,435 decoder_attention_mask=decoder_attention_mask,436 decoder_position_ids=decoder_position_ids,437 **kwargs,438 )439 logits = self.lm_head(outputs.last_hidden_state)440 logits = (441 torch.tanh(logits / self.final_logit_softcapping)442 * self.final_logit_softcapping443 )444 statistics = (None, None, None, None, None)445 if return_proposal_statistics:446 statistics = self._proposal_statistics(447 logits,448 denoise_temperature=denoise_temperature,449 )450 return ModilifyMk1BlockDiffusionOutput(451 logits=None if return_proposal_statistics else logits,452 heavy_hidden_state=outputs.last_hidden_state,453 next_latent_state=next_state,454 past_key_values=outputs.past_key_values,455 encoder_last_hidden_state=outputs.encoder_last_hidden_state,456 temporal_context=latent_context,457 latent_residual_diagnostics=outputs.latent_residual_diagnostics,458 proposal=statistics[0],459 proposal_confidence=statistics[1],460 token_entropy=statistics[2],461 greedy_proposal=statistics[3],462 greedy_confidence=statistics[4],463 )464 465 466__all__ = [467 "ModilifyMk1BlockDiffusionOutput",468 "ModilifyMk1Config",469 "ModilifyMk1DecoderModel",470 "ModilifyMk1EncoderModel",471 "ModilifyMk1ForBlockDiffusion",472 "ModilifyMk1Model",473]474 