Image-Text-to-Text
Transformers
Safetensors
modilify_mk1
text-generation
diffusion
multimodal
mixture-of-experts
trust-remote-code
conversational
custom_code
Instructions to use modilify/Modilify-Mk1 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use modilify/Modilify-Mk1 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("image-text-to-text", model="modilify/Modilify-Mk1", trust_remote_code=True) messages = [ { "role": "user", "content": [ {"type": "image", "url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/p-blog/candy.JPG"}, {"type": "text", "text": "What animal is on the candy?"} ] }, ] pipe(text=messages)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("modilify/Modilify-Mk1", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use modilify/Modilify-Mk1 with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "modilify/Modilify-Mk1" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk1", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }'Use Docker
docker model run hf.co/modilify/Modilify-Mk1
- SGLang
How to use modilify/Modilify-Mk1 with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "modilify/Modilify-Mk1" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk1", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "modilify/Modilify-Mk1" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk1", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }' - Docker Model Runner
How to use modilify/Modilify-Mk1 with Docker Model Runner:
docker model run hf.co/modilify/Modilify-Mk1
| # 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, | |
| ) | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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", | |
| ] | |