CoolFace
Modelpublic

modilify/Modilify-Mk1-preview

sourceHugging Faceotherupdated 2mo agoView on Hugging Face
1likes20downloads
modeling_modilify_mk1.py474 linesDownload Raw Back to root
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