# 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 dataclasses import dataclass, replace import math from typing import Any import torch from torch import nn from transformers.cache_utils import Cache from transformers.modeling_outputs import BaseModelOutputWithPast from transformers.utils import ModelOutput from transformers.models.diffusion_gemma import ( DiffusionGemmaDecoderModel, DiffusionGemmaEncoderModel, DiffusionGemmaPreTrainedModel, ) 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 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 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 only 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) self.post_init() def get_encoder(self) -> ModilifyMk1EncoderModel: """Return the standard 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 _proposal_statistics( self, logits: torch.Tensor, *, denoise_temperature: float | 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 = logits.float() / temperature probabilities = torch.softmax(scores, dim=-1) flat = probabilities.reshape(-1, probabilities.shape[-1]) proposal = torch.multinomial(flat, num_samples=1).view(logits.shape[:-1]) 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, **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, ) 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], ) __all__ = [ "ModilifyMk1BlockDiffusionOutput", "ModilifyMk1Config", "ModilifyMk1DecoderModel", "ModilifyMk1EncoderModel", "ModilifyMk1ForBlockDiffusion", "ModilifyMk1Model", ]