# Copyright 2026 Modilify # SPDX-License-Identifier: LicenseRef-Modilify-Open-Model-1.0 """Standard PyTorch multimodal model implementation for Modilify Mk1.""" from __future__ import annotations from collections.abc import Sequence from dataclasses import dataclass, replace import math from typing import Any import torch from torch import nn from torch.nn import functional as F from transformers.cache_utils import Cache from transformers.masking_utils import ( ALL_MASK_ATTENTION_FUNCTIONS, bidirectional_mask_function, ) from transformers.modeling_outputs import BaseModelOutputWithPast from transformers.utils import ModelOutput from transformers.models.diffusion_gemma import ( DiffusionGemmaDecoderModel, DiffusionGemmaEncoderModel, DiffusionGemmaPreTrainedModel, ) from transformers.models.diffusion_gemma.modeling_diffusion_gemma import ( DiffusionGemmaRMSNorm, DiffusionGemmaTextRouter, ) from .configuration_modilify_mk1 import ModilifyMk1Config from .generation_modilify_mk1 import ( ModilifyMk1GenerationConfig, ModilifyMk1GenerationMixin, ) from .latent_deliberation import ( LatentDeliberationState, LatentDeliberationTransformer, ) @dataclass class ModilifyMk1DecoderOutput(BaseModelOutputWithPast): """Decoder hidden states and latent-context diagnostics.""" token_embeddings: torch.FloatTensor | None = None latent_residual_diagnostics: dict[str, torch.Tensor] | None = None @dataclass class ModilifyMk1ModelOutput(BaseModelOutputWithPast): """Combined multimodal encoder and diffusion decoder output.""" token_embeddings: torch.FloatTensor | None = None encoder_last_hidden_state: torch.FloatTensor | None = None latent_residual_diagnostics: dict[str, torch.Tensor] | None = None @dataclass class ModilifyMk1BlockDiffusionOutput(ModelOutput): """Inference output used by the rolling diffusion generator.""" logits: torch.FloatTensor | None = None heavy_hidden_state: torch.FloatTensor | None = None next_latent_state: LatentDeliberationState | None = None past_key_values: Cache | None = None encoder_last_hidden_state: torch.FloatTensor | None = None temporal_context: torch.FloatTensor | None = None latent_residual_diagnostics: dict[str, torch.Tensor] | None = None proposal: torch.LongTensor | None = None proposal_confidence: torch.FloatTensor | None = None token_entropy: torch.FloatTensor | None = None greedy_proposal: torch.LongTensor | None = None greedy_confidence: torch.FloatTensor | None = None class ModilifyMk1RMSNorm(DiffusionGemmaRMSNorm): """Official RMSNorm parameters with a same-dtype residual forward.""" def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: """Normalize ``hidden_states`` and restore the input dtype.""" normed_output = self._norm(hidden_states) if self.with_scale: normed_output = normed_output * self.weight.to(dtype=normed_output.dtype) return normed_output.type_as(hidden_states) class ModilifyMk1TextRouter(DiffusionGemmaTextRouter): """Official router parameters with a log-softmax top-k route.""" def __init__(self, config: Any) -> None: super().__init__(config) self.norm = ModilifyMk1RMSNorm(self.hidden_size, eps=self.eps, with_scale=False) def forward( self, hidden_states: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Return route probabilities, top-k weights, and expert indices.""" hidden_states = self.norm(hidden_states) hidden_states = hidden_states * self.scale * self.scalar_root_size expert_scores = self.proj(hidden_states) router_probabilities = F.log_softmax(expert_scores, dim=-1).exp() top_k_weights, top_k_index = torch.topk( router_probabilities, k=self.config.top_k_experts, dim=-1, ) top_k_weights = top_k_weights / top_k_weights.sum(dim=-1, keepdim=True) top_k_weights = top_k_weights * self.per_expert_scale[top_k_index] return router_probabilities, top_k_weights, top_k_index def install_modilify_mk1_trunk_semantics(module: nn.Module) -> None: """Replace official leaf modules on this instance only. Args: module: Encoder, decoder, or parent module whose children should be swapped to the instance-scoped RMSNorm and router implementations. """ for name, child in list(module.named_children()): if type(child) is DiffusionGemmaRMSNorm: dim = int(child.weight.shape[0]) if child.with_scale else 1 replacement = ModilifyMk1RMSNorm( dim, eps=child.eps, with_scale=child.with_scale ) replacement.load_state_dict(child.state_dict()) setattr(module, name, replacement) elif type(child) is DiffusionGemmaTextRouter: replacement = ModilifyMk1TextRouter(child.config) replacement.load_state_dict(child.state_dict()) setattr(module, name, replacement) else: install_modilify_mk1_trunk_semantics(child) class ModilifyMk1EncoderModel(DiffusionGemmaEncoderModel): """Unmodified Transformers DiffusionGemma multimodal encoder.""" config_class = ModilifyMk1Config class ModilifyMk1DecoderModel(DiffusionGemmaDecoderModel): """DiffusionGemma decoder conditioned by recurrent latent embeddings.""" config_class = ModilifyMk1Config latent_residual_rms_ratio_cap = 0.5 @staticmethod def create_diffusion_decoder_attention_mask( config: Any, inputs_embeds: torch.Tensor, past_key_values: Cache, decoder_attention_mask: torch.Tensor | dict | None = None, ) -> dict[str, torch.Tensor | None]: """Build bidirectional canvas masks without skipping sliding layers. Args: config: Text configuration used for layer types and window size. inputs_embeds: Canvas embeddings that define query length and dtype. past_key_values: Prefix cache used to size the key/value axis. decoder_attention_mask: Optional 2-D mask or precomputed 4-D maps. Returns: A mapping from layer pattern to attention mask. """ if past_key_values is None: raise ValueError( "The diffusion mask requires `past_key_values` to construct the " "next attention mask correctly." ) if ( decoder_attention_mask is None or config._attn_implementation not in ALL_MASK_ATTENTION_FUNCTIONS._global_mapping ): return {"full_attention": None, "sliding_attention": None} if isinstance(decoder_attention_mask, dict) and all( mask.ndim == 4 for mask in decoder_attention_mask.values() ): return decoder_attention_mask text_config = config.get_text_config() if hasattr(config, "get_text_config") else config q_length = inputs_embeds.shape[1] q_offset = past_key_values.get_seq_length() if isinstance(q_offset, torch.Tensor): q_offset = q_offset.to(inputs_embeds.device) additional_kv_length = ( getattr(config, "canvas_length", 0) if past_key_values.is_compileable else 0 ) mask_mapping: dict[str, torch.Tensor | None] = {} for layer_pattern in set(text_config.layer_types): layer_idx = past_key_values.is_sliding.index( layer_pattern == "sliding_attention" ) kv_length, kv_offset = past_key_values.get_mask_sizes(q_length, layer_idx) kv_length += additional_kv_length if layer_pattern == "sliding_attention" and past_key_values.is_compileable: sliding_layer = past_key_values.layers[layer_idx] max_length = sliding_layer.get_max_length() + additional_kv_length if kv_length >= max_length: kv_length = max_length mask_mapping[layer_pattern] = ALL_MASK_ATTENTION_FUNCTIONS[ config._attn_implementation ]( batch_size=inputs_embeds.shape[0], q_length=q_length, kv_length=kv_length, q_offset=q_offset, kv_offset=kv_offset, mask_function=bidirectional_mask_function, attention_mask=decoder_attention_mask, allow_is_causal_skip=False, allow_is_bidirectional_skip=True, local_size=getattr(text_config, "sliding_window", None), dtype=inputs_embeds.dtype, config=text_config, use_vmap=False, device=inputs_embeds.device, ) return mask_mapping def merge_latent_context( self, token_embeddings: torch.Tensor, latent_context: torch.Tensor | None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Apply the native self-conditioning bridge to latent context. Args: token_embeddings: Embedded noisy canvas tokens. latent_context: Context emitted by the latent Transformer. Returns: Merged embeddings and scalar diagnostic tensors. """ context = ( torch.zeros_like(token_embeddings) if latent_context is None else latent_context.to(token_embeddings) ) if context.shape != token_embeddings.shape: raise ValueError("Latent context must match the canvas embedding shape.") mapper = self.self_conditioning normalized = mapper.pre_norm(context) mapped = mapper.down_proj( mapper.act_fn(mapper.gate_proj(normalized)) * mapper.up_proj(normalized) ) mapped_rms_per_token = mapped.float().square().mean(dim=-1, keepdim=True).sqrt() token_rms_per_token = ( token_embeddings.float().square().mean(dim=-1, keepdim=True).sqrt() ) cap = self.latent_residual_rms_ratio_cap * token_rms_per_token scale = cap / torch.sqrt(mapped_rms_per_token.square() + cap.square() + 1.0e-12) mapped = mapped * scale.to(mapped) combined = mapper.post_norm(token_embeddings + mapped) token_rms = token_embeddings.detach().float().square().mean().sqrt() mapped_rms = mapped.detach().float().square().mean().sqrt() diagnostics = { "token_embedding_rms": token_rms, "latent_context_rms": context.detach().float().square().mean().sqrt(), "mapped_context_rms": mapped_rms, "latent_to_embedding_rms_ratio": mapped_rms / token_rms.clamp_min(1.0e-12), } return combined, diagnostics def forward( self, decoder_input_ids: torch.LongTensor, past_key_values: Cache | None = None, temporal_context_embeddings: torch.FloatTensor | None = None, decoder_attention_mask: torch.Tensor | dict | None = None, decoder_position_ids: torch.LongTensor | None = None, **kwargs: Any, ) -> ModilifyMk1DecoderOutput: """Decode one noisy canvas using Transformers and PyTorch operations.""" token_embeddings = self.embed_tokens(decoder_input_ids) inputs_embeds, diagnostics = self.merge_latent_context( token_embeddings, temporal_context_embeddings, ) if decoder_position_ids is None: prefix = past_key_values.get_seq_length(0) if past_key_values is not None else 0 decoder_position_ids = torch.arange( prefix, prefix + inputs_embeds.shape[1], device=inputs_embeds.device, ).unsqueeze(0) if not isinstance(mask_mapping := decoder_attention_mask, dict): mask_mapping = self.create_diffusion_decoder_attention_mask( config=self.text_config, inputs_embeds=inputs_embeds, past_key_values=past_key_values, decoder_attention_mask=decoder_attention_mask, ) position_embeddings = { layer_type: self.rotary_emb(inputs_embeds, decoder_position_ids, layer_type) for layer_type in self.unique_layer_types } hidden_states = inputs_embeds for index, layer in enumerate(self.layers[: self.text_config.num_hidden_layers]): layer_type = self.text_config.layer_types[index] hidden_states = layer( hidden_states, position_embeddings=position_embeddings[layer_type], attention_mask=mask_mapping[layer_type], position_ids=decoder_position_ids, past_key_values=past_key_values, **kwargs, ) return ModilifyMk1DecoderOutput( last_hidden_state=self.norm(hidden_states), past_key_values=past_key_values, token_embeddings=token_embeddings, latent_residual_diagnostics=diagnostics, ) class ModilifyMk1Model(DiffusionGemmaPreTrainedModel): """Multimodal encoder plus latent-conditioned block diffusion decoder.""" config_class = ModilifyMk1Config _tied_weights_keys = { "encoder.language_model.norm.weight": "decoder.norm.weight", r"encoder.language_model.layers\.(?:[^.]+\.)*weight": r"decoder.layers\.(?:[^.]+\.)*weight", r"encoder.language_model.layers\.(?:[^.]+\.)*scale": r"decoder.layers\.(?:[^.]+\.)*scale", ( r"encoder.language_model.layers\.(?:[^.]+\.)*per_expert_scale" ): r"decoder.layers\.(?:[^.]+\.)*per_expert_scale", ( r"encoder.language_model.layers\.(?:[^.]+\.)*gate_up_proj" ): r"decoder.layers\.(?:[^.]+\.)*gate_up_proj", ( r"encoder.language_model.layers\.(?:[^.]+\.)*down_proj" ): r"decoder.layers\.(?:[^.]+\.)*down_proj", "encoder.language_model.embed_tokens.weight": "decoder.embed_tokens.weight", } def __init__(self, config: ModilifyMk1Config) -> None: super().__init__(config) self.encoder = ModilifyMk1EncoderModel(config) self.decoder = ModilifyMk1DecoderModel(config) install_modilify_mk1_trunk_semantics(self) self.post_init() def get_encoder(self) -> ModilifyMk1EncoderModel: """Return the multimodal encoder.""" return self.encoder def get_decoder(self) -> ModilifyMk1DecoderModel: """Return the diffusion decoder.""" return self.decoder def get_input_embeddings(self) -> nn.Module: """Return the shared text embedding module.""" return self.encoder.get_input_embeddings() def set_input_embeddings(self, value: nn.Module) -> None: """Set the shared text embedding module.""" self.encoder.set_input_embeddings(value) self.decoder.embed_tokens = value def forward( self, *, input_ids: torch.LongTensor | None = None, attention_mask: torch.Tensor | dict | None = None, past_key_values: Cache | None = None, position_ids: torch.LongTensor | None = None, decoder_input_ids: torch.LongTensor, temporal_context_embeddings: torch.FloatTensor | None = None, decoder_attention_mask: torch.Tensor | dict | None = None, decoder_position_ids: torch.LongTensor | None = None, **kwargs: Any, ) -> ModilifyMk1ModelOutput: """Encode multimodal context and decode one canvas.""" encoder_hidden_state = None encoder_keys = ("pixel_values", "mm_token_type_ids", "image_position_ids", "inputs_embeds") encoder_kwargs = {key: kwargs.pop(key) for key in encoder_keys if key in kwargs} if input_ids is not None: encoded = self.encoder( input_ids=input_ids, attention_mask=attention_mask, past_key_values=past_key_values, position_ids=position_ids, **encoder_kwargs, ) past_key_values = encoded.past_key_values encoder_hidden_state = encoded.last_hidden_state elif past_key_values is None: raise ValueError("Either `input_ids` or `past_key_values` is required.") decoded = self.decoder( decoder_input_ids=decoder_input_ids, past_key_values=past_key_values, temporal_context_embeddings=temporal_context_embeddings, decoder_attention_mask=decoder_attention_mask, decoder_position_ids=decoder_position_ids, **kwargs, ) return ModilifyMk1ModelOutput( last_hidden_state=decoded.last_hidden_state, past_key_values=past_key_values, token_embeddings=decoded.token_embeddings, encoder_last_hidden_state=encoder_hidden_state, latent_residual_diagnostics=decoded.latent_residual_diagnostics, ) class ModilifyMk1ForBlockDiffusion( DiffusionGemmaPreTrainedModel, ModilifyMk1GenerationMixin, ): """Inference-only multimodal Modilify Mk1 model.""" config_class = ModilifyMk1Config _tied_weights_keys = {"lm_head.weight": "model.decoder.embed_tokens.weight"} generation_config_class = ModilifyMk1GenerationConfig @torch.no_grad() def _init_weights(self, module: nn.Module) -> None: super()._init_weights(module) if isinstance(module, LatentDeliberationTransformer): module.reset_memory_slot_identity() def __init__(self, config: ModilifyMk1Config) -> None: super().__init__(config) self.model = ModilifyMk1Model(config) self.latent_deliberation = LatentDeliberationTransformer( hidden_size=config.text_config.hidden_size, latent_dim=config.latent_dim, memory_slots=config.latent_memory_slots, num_layers=config.latent_num_layers, num_heads=config.latent_num_heads, local_attention_window=config.latent_local_attention_window, dropout=config.latent_dropout, ) self.lm_head = nn.Linear( config.text_config.hidden_size, config.text_config.vocab_size, bias=False, ) self.final_logit_softcapping = config.text_config.final_logit_softcapping self.post_init() def _prepare_latent_context( self, decoder_input_ids: torch.LongTensor, *, history_hidden_state: torch.Tensor | None, confidence: torch.Tensor | None, entropy: torch.Tensor | None, age: torch.Tensor | None, latent_state: LatentDeliberationState | None, ) -> tuple[torch.Tensor, LatentDeliberationState]: """Advance recurrent latent state for the current canvas.""" batch_size, canvas_length = decoder_input_ids.shape dtype = self.model.decoder.embed_tokens.weight.dtype if latent_state is None: latent_state = LatentDeliberationState.empty( batch_size=batch_size, canvas_length=canvas_length, latent_dim=self.config.latent_dim, memory_slots=self.config.latent_memory_slots, device=decoder_input_ids.device, dtype=dtype, ) confidence = ( latent_state.confidence if confidence is None else confidence.squeeze(-1).float() ) entropy = latent_state.entropy if entropy is None else entropy.squeeze(-1).float() if age is not None: latent_state = replace( latent_state, age=age.to(device=decoder_input_ids.device, dtype=torch.int32), ) token_embeddings = self.model.decoder.embed_tokens(decoder_input_ids) history = ( torch.zeros_like(token_embeddings) if history_hidden_state is None else history_hidden_state ) return self.latent_deliberation( heavy_hidden=history, token_embeddings=token_embeddings, confidence=confidence, entropy=entropy, state=latent_state, ) def _apply_repetition_penalty( self, logits: torch.Tensor, *, repetition_token_mask: torch.BoolTensor | None, repetition_penalty: float, ) -> torch.Tensor: """Apply a sign-aware Transformers repetition penalty. Args: logits: Soft-capped scores, shape ``[batch, canvas, vocab]``. repetition_token_mask: Tokens already seen, shape ``[batch, vocab]``. repetition_penalty: Penalty factor. ``1.0`` leaves logits unchanged. Returns: Penalized logits with the same shape as ``logits``. """ if ( repetition_token_mask is None or not math.isfinite(repetition_penalty) or repetition_penalty == 1.0 ): return logits if repetition_penalty <= 0: raise ValueError("`repetition_penalty` must be a positive finite number.") if repetition_token_mask.shape != (logits.shape[0], logits.shape[-1]): raise ValueError( "`repetition_token_mask` must have shape [batch, vocab]." ) scores = logits.float() penalized = torch.where(scores < 0, scores * repetition_penalty, scores / repetition_penalty) mask = repetition_token_mask.to(device=scores.device).unsqueeze(1) return torch.where(mask, penalized, scores).to(dtype=logits.dtype) def _proposal_statistics( self, logits: torch.Tensor, *, denoise_temperature: float | None = None, repetition_token_mask: torch.BoolTensor | None = None, repetition_penalty: float = 1.0, sampling_generators: Sequence[torch.Generator] | None = None, ) -> tuple[ torch.LongTensor, torch.Tensor, torch.Tensor, torch.LongTensor, torch.Tensor, ]: """Compute exact proposal statistics with standard PyTorch operations.""" temperature = ( self.config.denoise_temperature if denoise_temperature is None else float(denoise_temperature) ) if not math.isfinite(temperature) or temperature <= 0.0: raise ValueError("`denoise_temperature` must be positive.") scores = self._apply_repetition_penalty( logits, repetition_token_mask=repetition_token_mask, repetition_penalty=repetition_penalty, ).float() / temperature probabilities = torch.softmax(scores, dim=-1) if sampling_generators is None: proposal = torch.multinomial( probabilities.reshape(-1, probabilities.shape[-1]), num_samples=1, ).view(logits.shape[:-1]) else: if len(sampling_generators) != logits.shape[0]: raise ValueError("Sampling requires one generator per batch row.") rows = [] for row, generator in enumerate(sampling_generators): rows.append( torch.multinomial( probabilities[row], num_samples=1, generator=generator, ).squeeze(-1) ) proposal = torch.stack(rows, dim=0) proposal_confidence = probabilities.gather(-1, proposal.unsqueeze(-1)).squeeze(-1) greedy_proposal = probabilities.argmax(dim=-1) greedy_confidence = probabilities.gather( -1, greedy_proposal.unsqueeze(-1) ).squeeze(-1) token_entropy = -( probabilities * probabilities.clamp_min(1.0e-30).log() ).sum(dim=-1) return ( proposal, proposal_confidence, token_entropy, greedy_proposal, greedy_confidence, ) def forward( self, *, input_ids: torch.LongTensor | None = None, attention_mask: torch.Tensor | dict | None = None, past_key_values: Cache | None = None, position_ids: torch.LongTensor | None = None, decoder_input_ids: torch.LongTensor, previous_confidence: torch.FloatTensor | None = None, previous_entropy: torch.FloatTensor | None = None, token_age: torch.Tensor | None = None, latent_state: LatentDeliberationState | None = None, history_hidden_state: torch.FloatTensor | None = None, decoder_attention_mask: torch.Tensor | dict | None = None, decoder_position_ids: torch.LongTensor | None = None, return_proposal_statistics: bool = False, denoise_temperature: float | None = None, repetition_token_mask: torch.BoolTensor | None = None, repetition_penalty: float = 1.0, sampling_generators: Sequence[torch.Generator] | None = None, **kwargs: Any, ) -> ModilifyMk1BlockDiffusionOutput: """Run one inference step over a noisy diffusion canvas.""" latent_context, next_state = self._prepare_latent_context( decoder_input_ids, history_hidden_state=history_hidden_state, confidence=previous_confidence, entropy=previous_entropy, age=token_age, latent_state=latent_state, ) outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, past_key_values=past_key_values, position_ids=position_ids, decoder_input_ids=decoder_input_ids, temporal_context_embeddings=latent_context, decoder_attention_mask=decoder_attention_mask, decoder_position_ids=decoder_position_ids, **kwargs, ) logits = self.lm_head(outputs.last_hidden_state) logits = ( torch.tanh(logits / self.final_logit_softcapping) * self.final_logit_softcapping ) statistics = (None, None, None, None, None) if return_proposal_statistics: statistics = self._proposal_statistics( logits, denoise_temperature=denoise_temperature, repetition_token_mask=repetition_token_mask, repetition_penalty=repetition_penalty, sampling_generators=sampling_generators, ) return ModilifyMk1BlockDiffusionOutput( logits=None if return_proposal_statistics else logits, heavy_hidden_state=outputs.last_hidden_state, next_latent_state=next_state, past_key_values=outputs.past_key_values, encoder_last_hidden_state=outputs.encoder_last_hidden_state, temporal_context=latent_context, latent_residual_diagnostics=outputs.latent_residual_diagnostics, proposal=statistics[0], proposal_confidence=statistics[1], token_entropy=statistics[2], greedy_proposal=statistics[3], greedy_confidence=statistics[4], ) ModilifyMk1Model.register_for_auto_class("AutoModel") ModilifyMk1ForBlockDiffusion.register_for_auto_class("AutoModelForCausalLM") ModilifyMk1ForBlockDiffusion.register_for_auto_class("AutoModelForMultimodalLM") __all__ = [ "ModilifyMk1BlockDiffusionOutput", "ModilifyMk1Config", "ModilifyMk1DecoderModel", "ModilifyMk1EncoderModel", "ModilifyMk1ForBlockDiffusion", "ModilifyMk1Model", ]