Modilify-Mk1-preview / modilify_mk1 /modeling_modilify_mk1.py
ydy9038074's picture
Publish Modilify Mk1 Preview
164d101 verified
Raw
History Blame Contribute Delete
18.6 kB
# 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",
]