Modilify-Mk1 / modeling_modilify_mk1.py
ydy9038074's picture
Publish Modilify Mk1
9315757 verified
Raw
History Blame Contribute Delete
27.7 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 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",
]