Text Generation
MLX
Safetensors
modilify_mk2
diffusion
mixture-of-experts
custom-code
modilify-mk2
conversational
Instructions to use modilify/Modilify-Mk2-preview-mlx with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use modilify/Modilify-Mk2-preview-mlx with MLX:
# Make sure mlx-lm is installed # pip install --upgrade mlx-lm # Generate text with mlx-lm from mlx_lm import load, generate model, tokenizer = load("modilify/Modilify-Mk2-preview-mlx") prompt = "Write a story about Einstein" messages = [{"role": "user", "content": prompt}] prompt = tokenizer.apply_chat_template( messages, add_generation_prompt=True ) text = generate(model, tokenizer, prompt=prompt, verbose=True) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Pi
How to use modilify/Modilify-Mk2-preview-mlx with Pi:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "modilify/Modilify-Mk2-preview-mlx"
Configure the model in Pi
# Install Pi: npm install -g @earendil-works/pi-coding-agent # Add to ~/.pi/agent/models.json: { "providers": { "mlx-lm": { "baseUrl": "http://localhost:8080/v1", "api": "openai-completions", "apiKey": "none", "models": [ { "id": "modilify/Modilify-Mk2-preview-mlx" } ] } } }Run Pi
# Start Pi in your project directory: pi
- MLX LM
How to use modilify/Modilify-Mk2-preview-mlx with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Interactive chat REPL mlx_lm.chat --model "modilify/Modilify-Mk2-preview-mlx"
Run an OpenAI-compatible server
# Install MLX LM uv tool install mlx-lm # Start the server mlx_lm.server --model "modilify/Modilify-Mk2-preview-mlx" # Calling the OpenAI-compatible server with curl curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "modilify/Modilify-Mk2-preview-mlx", "messages": [ {"role": "user", "content": "Hello"} ] }' - Hermes Agent
How to use modilify/Modilify-Mk2-preview-mlx with Hermes Agent:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "modilify/Modilify-Mk2-preview-mlx"
Configure Hermes
# Install Hermes: curl -fsSL https://hermes-agent.nousresearch.com/install.sh | bash hermes setup # Point Hermes at the local server: hermes config set model.provider custom hermes config set model.base_url http://127.0.0.1:8080/v1 hermes config set model.default modilify/Modilify-Mk2-preview-mlx
Run Hermes
hermes
- Atomic Chat
- OpenClaw
How to use modilify/Modilify-Mk2-preview-mlx with OpenClaw:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "modilify/Modilify-Mk2-preview-mlx"
Configure OpenClaw
# Install OpenClaw: npm install -g openclaw@latest # Register the local server and set it as the default model: openclaw onboard --non-interactive --mode local \ --auth-choice custom-api-key \ --custom-base-url http://127.0.0.1:8080/v1 \ --custom-model-id "modilify/Modilify-Mk2-preview-mlx" \ --custom-provider-id mlx-lm \ --custom-compatibility openai \ --custom-text-input \ --accept-risk \ --skip-health
Run OpenClaw
openclaw agent --local --agent main --message "Hello from Hugging Face"
Download modilify_mk2/configuration_modilify_mk2.py from modilify/Modilify-Mk2-preview-mlx: direct link, hf CLI and curl.
- Browser
- Download file 22.6 kB
-
https://huggingface.co/modilify/Modilify-Mk2-preview-mlx/resolve/main/modilify_mk2/configuration_modilify_mk2.py
- Command line
-
hf download hf://modilify/Modilify-Mk2-preview-mlx/modilify_mk2/configuration_modilify_mk2.py
-
curl -L -o configuration_modilify_mk2.py https://huggingface.co/modilify/Modilify-Mk2-preview-mlx/resolve/main/modilify_mk2/configuration_modilify_mk2.py
22.6 kB
| """Strict text-only configuration for the schema25 GDN2 protocol.""" | |
| from __future__ import annotations | |
| import json | |
| import math | |
| from pathlib import Path | |
| from collections.abc import Mapping, Sequence | |
| from typing import Any | |
| from transformers.configuration_utils import PreTrainedConfig | |
| from transformers.models.diffusion_gemma import DiffusionGemmaTextConfig | |
| STATE_SCHEMA_VERSION = 25 | |
| # Keep topology stable; CE scope and normalization have separate objective tags. | |
| TRAINING_SCHEME = ( | |
| "gold_prefix_shared_commit_0_256_committed_ce_calibration_" | |
| "causal_throughput_terminal_sft" | |
| ) | |
| MEMORY_SCHEME = "dual_timescale_gdn2_trajectory_memory" | |
| MEMORY_ARCHITECTURE = "compact_gdn2_v2" | |
| CONFIG_PROTOCOL_ERROR = "ModilifyMk2 requires schema25 compact_gdn2_v2; older memory topologies require a new run." | |
| HISTORY_VIEWS = 4 | |
| COMMIT_SEQUENCE_LAYERS = 2 | |
| DENOISE_TEMPERATURE = 0.8 | |
| VOCAB_CHUNK_SIZE = 32_768 | |
| COMMIT_READINESS_OBJECTIVE = "frontier_prefix_budget_v2" | |
| TOKEN_CE_SUPERVISION = "valid_canvas_v1" | |
| TOKEN_CE_NORMALIZATION = "per_sample_exposure_v1" | |
| def require_current_checkpoint_protocol(metadata: Mapping[str, Any]) -> None: | |
| """Reject checkpoints that do not implement the GDN2 state topology.""" | |
| if ( | |
| metadata.get("state_schema_version") != STATE_SCHEMA_VERSION | |
| or metadata.get("training_scheme") != TRAINING_SCHEME | |
| or metadata.get("memory_scheme") != MEMORY_SCHEME | |
| or metadata.get("memory_architecture") != MEMORY_ARCHITECTURE | |
| ): | |
| raise RuntimeError( | |
| f"ModilifyMk2 checkpoint requires schema{STATE_SCHEMA_VERSION} {MEMORY_ARCHITECTURE}; " | |
| "older memory topologies cannot be restored. Start from the base model." | |
| ) | |
| class ModilifyMk2TextConfig(DiffusionGemmaTextConfig): | |
| model_type = "modilify_mk2_text" | |
| vocab_size: int = 262_144 | |
| hidden_size: int = 2816 | |
| intermediate_size: int = 2112 | |
| num_hidden_layers: int = 30 | |
| num_attention_heads: int = 16 | |
| num_key_value_heads: int = 8 | |
| head_dim: int = 256 | |
| max_position_embeddings: int = 262_144 | |
| sliding_window: int = 1024 | |
| use_bidirectional_attention: str | None = None | |
| num_global_key_value_heads: int | None = 2 | |
| global_head_dim: int = 512 | |
| num_experts: int | None = 128 | |
| top_k_experts: int | None = 8 | |
| moe_intermediate_size: int | None = 704 | |
| class ModilifyMk2Config(PreTrainedConfig): | |
| """Configuration with recurrent GDN2 trajectory memory.""" | |
| model_type = "modilify_mk2" | |
| sub_configs = {"text_config": ModilifyMk2TextConfig} | |
| def __init__( | |
| self, | |
| text_config: ModilifyMk2TextConfig | dict[str, Any] | None = None, | |
| *, | |
| canvas_length: int = 256, | |
| initializer_range: float = 0.02, | |
| tie_word_embeddings: bool = True, | |
| state_schema_version: int = STATE_SCHEMA_VERSION, | |
| training_scheme: str = TRAINING_SCHEME, | |
| memory_scheme: str = MEMORY_SCHEME, | |
| memory_architecture: str = MEMORY_ARCHITECTURE, | |
| latent_dim: int = 2816, | |
| latent_ffn_dim: int = 7168, | |
| latent_memory_slots: int = 256, | |
| latent_num_layers: int = 4, | |
| latent_num_heads: int = 16, | |
| latent_local_attention_window: int = 128, | |
| latent_history_length: int = 16, | |
| latent_tape_probes: int = 1, | |
| latent_tape_scheme: str = "gdn2_spatial_probe_v1", | |
| latent_history_views: int = HISTORY_VIEWS, | |
| latent_history_kv_rank: int | None = None, | |
| latent_working_last_block_global: bool = True, | |
| working_memory_bus: bool = True, | |
| persistent_memory_bus: bool = True, | |
| persistent_memory_write: str = "commit_only_transformer", | |
| commit_sequence_layers: int = COMMIT_SEQUENCE_LAYERS, | |
| commit_sequence_dim: int | None = None, | |
| writer_slot_gate: str = "per_slot", | |
| latent_working_bus_unfreeze_steps: int = 0, | |
| latent_persistent_bus_unfreeze_steps: int = 0, | |
| training_bptt_steps: int = 16, | |
| kv_cache_bucket_size: int = 128, | |
| turn_end_token_id: int = 106, | |
| terminal_token_ids: Sequence[int] | None = None, | |
| channel_end_token_id: int = 101, | |
| eos_token_id: int = 1, | |
| commit_failure_budget: float = 0.2, | |
| commit_top_k: int | None = 40, | |
| commit_min_p: float | None = 0.05, | |
| commit_target_confidence: float | None = 0.5, | |
| commit_entropy_weight: float = 1.0, | |
| commit_confidence_power: float = 1.0, | |
| commit_gold_alpha: float = 0.4, | |
| commit_gold_weight: float = 1.0, | |
| commit_readiness_failure_weight: float = 1.0, | |
| token_loss_weight: float = 1.0, | |
| token_ce_supervision: str = TOKEN_CE_SUPERVISION, | |
| token_ce_normalization: str = TOKEN_CE_NORMALIZATION, | |
| confidence_calibration_loss_weight: float = 0.1, | |
| commit_readiness_objective: str = COMMIT_READINESS_OBJECTIVE, | |
| commit_readiness_target_tokens: int = 16, | |
| commit_readiness_loss_weight: float = 0.1, | |
| commit_readiness_budget_margin: float = 0.02, | |
| commit_readiness_beta: float = 0.02, | |
| terminal_stop_loss_weight: float = 1.0, | |
| terminal_stop_target_probability: float = 0.95, | |
| **kwargs: Any, | |
| ) -> None: | |
| if any(key.startswith("commit_throughput_") for key in kwargs): | |
| raise ValueError("Removed commit-throughput configuration fields.") | |
| if token_ce_supervision != TOKEN_CE_SUPERVISION: | |
| raise ValueError("Schema25 requires valid-canvas token CE.") | |
| if state_schema_version != STATE_SCHEMA_VERSION: | |
| raise RuntimeError(CONFIG_PROTOCOL_ERROR) | |
| if training_scheme != TRAINING_SCHEME or memory_scheme != MEMORY_SCHEME: | |
| raise RuntimeError(CONFIG_PROTOCOL_ERROR) | |
| if memory_architecture != MEMORY_ARCHITECTURE: | |
| raise RuntimeError("Schema25 requires compact_gdn2_v2 memory; start a new run.") | |
| if text_config is None: | |
| text_config = ModilifyMk2TextConfig() | |
| elif isinstance(text_config, dict): | |
| text_config = dict(text_config) | |
| text_config.pop("model_type", None) | |
| text_config["use_bidirectional_attention"] = None | |
| text_config = ModilifyMk2TextConfig(**text_config) | |
| elif not isinstance(text_config, ModilifyMk2TextConfig): | |
| payload = text_config.to_dict() | |
| payload.pop("model_type", None) | |
| text_config = ModilifyMk2TextConfig(**payload) | |
| self.text_config = text_config | |
| self.canvas_length = canvas_length | |
| self.initializer_range = initializer_range | |
| self.state_schema_version = STATE_SCHEMA_VERSION | |
| self.training_scheme = training_scheme | |
| self.memory_scheme = memory_scheme | |
| self.memory_architecture = memory_architecture | |
| self.latent_dim = latent_dim | |
| self.latent_ffn_dim = latent_ffn_dim | |
| self.latent_memory_slots = latent_memory_slots | |
| self.latent_num_layers = latent_num_layers | |
| self.latent_num_heads = latent_num_heads | |
| self.latent_local_attention_window = latent_local_attention_window | |
| self.latent_history_length = latent_history_length | |
| self.latent_tape_probes = int(latent_tape_probes) | |
| self.latent_tape_scheme = str(latent_tape_scheme) | |
| self.latent_history_views = int(latent_history_views) | |
| if latent_history_kv_rank is None: | |
| rank = min(1024, latent_dim) | |
| rank -= rank % max(latent_num_heads, 1) | |
| if rank <= 0: | |
| rank = latent_num_heads | |
| self.latent_history_kv_rank = rank | |
| else: | |
| self.latent_history_kv_rank = latent_history_kv_rank | |
| self.latent_working_last_block_global = bool(latent_working_last_block_global) | |
| self.working_memory_bus = bool(working_memory_bus) | |
| self.persistent_memory_bus = bool(persistent_memory_bus) | |
| self.persistent_memory_write = str(persistent_memory_write) | |
| self.commit_sequence_layers = int(commit_sequence_layers) | |
| if commit_sequence_dim is None: | |
| self.commit_sequence_dim = self.latent_history_kv_rank | |
| else: | |
| self.commit_sequence_dim = int(commit_sequence_dim) | |
| self.writer_slot_gate = str(writer_slot_gate) | |
| self.latent_working_bus_unfreeze_steps = int(latent_working_bus_unfreeze_steps) | |
| self.latent_persistent_bus_unfreeze_steps = int( | |
| latent_persistent_bus_unfreeze_steps | |
| ) | |
| self.training_bptt_steps = training_bptt_steps | |
| self.kv_cache_bucket_size = kv_cache_bucket_size | |
| self.turn_end_token_id = turn_end_token_id | |
| if terminal_token_ids is None: | |
| self.terminal_token_ids = (int(turn_end_token_id),) | |
| else: | |
| self.terminal_token_ids = tuple(int(token_id) for token_id in terminal_token_ids) | |
| self.channel_end_token_id = int(channel_end_token_id) | |
| self.commit_failure_budget = float(commit_failure_budget) | |
| self.commit_top_k = int(commit_top_k) if commit_top_k is not None else None | |
| self.commit_min_p = float(commit_min_p) if commit_min_p is not None else None | |
| self.commit_target_confidence = float(commit_target_confidence) if commit_target_confidence is not None else None | |
| self.commit_entropy_weight = float(commit_entropy_weight) | |
| self.commit_confidence_power = float(commit_confidence_power) | |
| self.commit_gold_alpha = float(commit_gold_alpha) | |
| self.commit_gold_weight = float(commit_gold_weight) | |
| self.commit_readiness_failure_weight = float(commit_readiness_failure_weight) | |
| self.token_loss_weight = token_loss_weight | |
| self.token_ce_supervision = str(token_ce_supervision) | |
| self.token_ce_normalization = str(token_ce_normalization) | |
| self.confidence_calibration_loss_weight = confidence_calibration_loss_weight | |
| self.commit_readiness_objective = str(commit_readiness_objective) | |
| self.commit_readiness_target_tokens = int(commit_readiness_target_tokens) | |
| self.commit_readiness_loss_weight = float(commit_readiness_loss_weight) | |
| self.commit_readiness_budget_margin = float(commit_readiness_budget_margin) | |
| self.commit_readiness_beta = float(commit_readiness_beta) | |
| self.terminal_stop_loss_weight = terminal_stop_loss_weight | |
| self.terminal_stop_target_probability = terminal_stop_target_probability | |
| self.vocab_chunk_size = VOCAB_CHUNK_SIZE | |
| super().__init__( | |
| tie_word_embeddings=tie_word_embeddings, | |
| eos_token_id=eos_token_id, | |
| **kwargs, | |
| ) | |
| self._validate() | |
| def _validate(self) -> None: | |
| positive = ( | |
| self.canvas_length, self.latent_dim, self.latent_ffn_dim, | |
| self.latent_memory_slots, | |
| self.latent_num_layers, self.latent_num_heads, | |
| self.latent_local_attention_window, self.latent_history_length, | |
| self.latent_tape_probes, | |
| self.latent_history_kv_rank, | |
| self.commit_sequence_layers, self.commit_sequence_dim, | |
| self.training_bptt_steps, self.kv_cache_bucket_size, | |
| ) | |
| if any(value <= 0 for value in positive): | |
| raise ValueError("All schema25 dimensions and intervals must be positive.") | |
| if self.canvas_length != 256: | |
| raise ValueError("Schema25 requires a 256-token canvas.") | |
| if self.latent_tape_scheme != "gdn2_spatial_probe_v1": | |
| raise ValueError("Schema25 requires GDN2 spatial probes.") | |
| if self.latent_tape_probes > self.canvas_length: | |
| raise ValueError("Spatial probes cannot exceed canvas positions.") | |
| if self.commit_sequence_dim % self.latent_num_heads: | |
| raise ValueError("`commit_sequence_dim` must be divisible by `latent_num_heads`.") | |
| if self.latent_working_bus_unfreeze_steps < 0: | |
| raise ValueError("`latent_working_bus_unfreeze_steps` must be non-negative.") | |
| if self.latent_persistent_bus_unfreeze_steps < 0: | |
| raise ValueError( | |
| "`latent_persistent_bus_unfreeze_steps` must be non-negative." | |
| ) | |
| if self.latent_persistent_bus_unfreeze_steps < self.latent_working_bus_unfreeze_steps: | |
| raise ValueError( | |
| "Persistent bus must unfreeze no earlier than the working bus." | |
| ) | |
| if self.latent_history_kv_rank > self.latent_dim: | |
| raise ValueError("`latent_history_kv_rank` must not exceed `latent_dim`.") | |
| if self.latent_history_kv_rank % self.latent_num_heads: | |
| raise ValueError("`latent_history_kv_rank` must be divisible by `latent_num_heads`.") | |
| if not isinstance(self.eos_token_id, int) or self.eos_token_id < 0: | |
| raise ValueError("ModilifyMk2 requires one non-negative integer EOS token ID.") | |
| if not isinstance(self.channel_end_token_id, int) or self.channel_end_token_id < 0: | |
| raise ValueError("`channel_end_token_id` must be a non-negative integer.") | |
| if not self.terminal_token_ids: | |
| raise ValueError("`terminal_token_ids` must not be empty.") | |
| if any( | |
| not isinstance(token_id, int) or token_id < 0 | |
| for token_id in self.terminal_token_ids | |
| ): | |
| raise ValueError("`terminal_token_ids` must be non-negative integers.") | |
| if self.latent_dim % self.latent_num_heads: | |
| raise ValueError("`latent_dim` must be divisible by `latent_num_heads`.") | |
| if self.commit_readiness_objective != COMMIT_READINESS_OBJECTIVE: | |
| raise ValueError( | |
| "Unsupported commit-readiness objective: " | |
| f"{self.commit_readiness_objective!r}." | |
| ) | |
| if not all(math.isfinite(v) for v in ( | |
| self.commit_failure_budget, self.commit_entropy_weight, | |
| self.commit_confidence_power, self.commit_gold_alpha, self.commit_gold_weight, | |
| self.commit_readiness_failure_weight, self.commit_readiness_loss_weight, | |
| )): | |
| raise ValueError("Commit policy parameters must be finite.") | |
| if self.commit_failure_budget <= 0 or self.commit_entropy_weight < 0: | |
| raise ValueError("Commit budget must be positive and entropy weight non-negative.") | |
| if self.commit_top_k is not None and self.commit_top_k <= 0: | |
| raise ValueError("`commit_top_k` must be a positive integer.") | |
| if self.commit_min_p is not None and not 0.0 < self.commit_min_p < 1.0: | |
| raise ValueError("`commit_min_p` must be in (0, 1).") | |
| if self.commit_target_confidence is not None and not 0.0 < self.commit_target_confidence < 1.0: | |
| raise ValueError("`commit_target_confidence` must be in (0, 1).") | |
| if self.commit_confidence_power <= 0 or not 0 < self.commit_gold_alpha < 1: | |
| raise ValueError("Commit power must be positive and gold alpha in (0, 1).") | |
| if self.commit_gold_weight <= 0: | |
| raise ValueError("Commit gold weight must be positive.") | |
| if self.commit_readiness_failure_weight < 0: | |
| raise ValueError("Commit readiness failure weight must be non-negative.") | |
| if self.commit_readiness_target_tokens <= 0: | |
| raise ValueError("`commit_readiness_target_tokens` must be positive.") | |
| if self.commit_readiness_loss_weight < 0: | |
| raise ValueError("`commit_readiness_loss_weight` must be non-negative.") | |
| if not 0.0 <= self.commit_readiness_budget_margin < self.commit_failure_budget: | |
| raise ValueError( | |
| "`commit_readiness_budget_margin` must be in [0, commit_failure_budget)." | |
| ) | |
| if self.commit_readiness_beta <= 0: | |
| raise ValueError("`commit_readiness_beta` must be positive.") | |
| if self.terminal_stop_loss_weight < 0: | |
| raise ValueError("`terminal_stop_loss_weight` must be non-negative.") | |
| if not 0.0 < self.terminal_stop_target_probability < 1.0: | |
| raise ValueError("`terminal_stop_target_probability` must be in (0, 1).") | |
| loss_weights = ( | |
| self.token_loss_weight, | |
| self.confidence_calibration_loss_weight, | |
| self.commit_readiness_loss_weight, | |
| self.terminal_stop_loss_weight, | |
| ) | |
| if any(not math.isfinite(weight) or weight < 0 for weight in loss_weights): | |
| raise ValueError("Schema25 fixed loss weights must be finite and non-negative.") | |
| if self.token_ce_normalization != TOKEN_CE_NORMALIZATION: | |
| raise ValueError("Unsupported token CE normalization.") | |
| def from_dict(cls, config_dict: dict[str, Any], **kwargs: Any) -> "ModilifyMk2Config": | |
| return_unused_kwargs = kwargs.pop("return_unused_kwargs", False) | |
| payload = dict(config_dict) | |
| if (payload.get("state_schema_version") != STATE_SCHEMA_VERSION | |
| or payload.get("memory_architecture") != MEMORY_ARCHITECTURE): | |
| raise RuntimeError(CONFIG_PROTOCOL_ERROR) | |
| payload.pop("model_type", None) | |
| payload.pop("architectures", None) | |
| for key in tuple(kwargs): | |
| if key in payload: | |
| payload[key] = kwargs.pop(key) | |
| config = cls(**payload) | |
| for key in tuple(kwargs): | |
| if hasattr(config, key) or key == "name_or_path": | |
| setattr(config, key, kwargs.pop(key)) | |
| return (config, kwargs) if return_unused_kwargs else config | |
| __all__ = [ | |
| "ModilifyMk2Config", "ModilifyMk2TextConfig", "CONFIG_PROTOCOL_ERROR", | |
| "COMMIT_READINESS_OBJECTIVE", "COMMIT_SEQUENCE_LAYERS", "DENOISE_TEMPERATURE", | |
| "HISTORY_VIEWS", "MEMORY_SCHEME", "MEMORY_ARCHITECTURE", "STATE_SCHEMA_VERSION", | |
| "TRAINING_SCHEME", "VOCAB_CHUNK_SIZE", | |
| "TOKEN_CE_SUPERVISION", "TOKEN_CE_NORMALIZATION", | |
| "require_current_checkpoint_protocol", | |
| ] | |
| class ModilifyMk2GenerationConfig: | |
| """Native settings; no autoregressive sampler or Torch runtime is needed.""" | |
| def __init__(self, **kwargs: Any): | |
| unsupported = ( | |
| 'sampler_config', 'stability_threshold', 'confidence_threshold', | |
| 'one_token_per_denoise_step', 'compile_generation', 'sliding_denoise', | |
| 'adaptive_ponder_budget', 'force_commit_on_max_steps', 'ponder_budget_id', | |
| ) | |
| configured = [name for name in unsupported if kwargs.pop(name, None) not in (None, False)] | |
| if configured: | |
| raise ValueError(f'Unsupported generation fields: {configured}') | |
| defaults = { | |
| 'max_new_tokens': None, 'max_denoising_steps': None, | |
| 'bos_token_id': None, 'pad_token_id': None, | |
| 'eos_token_id': None, 'turn_end_token_id': None, 'max_ponder_steps': 64, | |
| 'jump_on_no_progress_after': 12, 'min_trajectory_progress': 0.005, | |
| 'repetition_penalty': 1.0, 'repetition_penalty_exclude_token_ids': [], | |
| } | |
| for name, default in defaults.items(): | |
| setattr(self, name, kwargs.pop(name, default)) | |
| self.repetition_penalty_exclude_token_ids = list(dict.fromkeys( | |
| int(value) for value in self.repetition_penalty_exclude_token_ids or () | |
| )) | |
| def validate(self) -> None: | |
| for name in ('max_new_tokens', 'max_denoising_steps', | |
| 'max_ponder_steps', 'jump_on_no_progress_after'): | |
| value = getattr(self, name) | |
| if value is not None and (isinstance(value, bool) or not isinstance(value, int) or value <= 0): | |
| raise ValueError(f'{name} must be a positive integer.') | |
| if not math.isfinite(self.repetition_penalty) or self.repetition_penalty <= 0: | |
| raise ValueError('repetition_penalty must be finite and positive.') | |
| if not math.isfinite(self.min_trajectory_progress) or self.min_trajectory_progress < 0: | |
| raise ValueError('min_trajectory_progress must be finite and nonnegative.') | |
| if self.turn_end_token_id is not None and self.turn_end_token_id < 0: | |
| raise ValueError('turn_end_token_id must be nonnegative.') | |
| if any(isinstance(value, bool) or not isinstance(value, int) or value < 0 | |
| for value in self.repetition_penalty_exclude_token_ids): | |
| raise ValueError('Excluded token IDs must be nonnegative integers.') | |
| def load_generation_config(model_dir: str | Path) -> ModilifyMk2GenerationConfig: | |
| path = Path(model_dir) / 'generation_config.json' | |
| if not path.is_file(): | |
| raise FileNotFoundError(f'Missing generation configuration: {path}') | |
| return ModilifyMk2GenerationConfig(**json.loads(path.read_text(encoding='utf-8'))) | |
| def configure_generation_config( | |
| generation_config: ModilifyMk2GenerationConfig, processor: Any, *, | |
| max_new_tokens: int, max_denoising_steps: int | None, | |
| repetition_penalty: float | None = None, | |
| ) -> ModilifyMk2GenerationConfig: | |
| if generation_config.max_denoising_steps is None: | |
| generation_config.max_denoising_steps = 48 | |
| generation_config.max_new_tokens = max_new_tokens | |
| if max_denoising_steps is not None: | |
| generation_config.max_denoising_steps = max_denoising_steps | |
| if repetition_penalty is not None: | |
| generation_config.repetition_penalty = repetition_penalty | |
| tokenizer = getattr(processor, 'tokenizer', processor) | |
| if generation_config.bos_token_id is None: | |
| generation_config.bos_token_id = tokenizer.bos_token_id | |
| if generation_config.pad_token_id is None: | |
| generation_config.pad_token_id = tokenizer.pad_token_id | |
| turn_id = tokenizer.convert_tokens_to_ids('<turn|>') | |
| if turn_id is None or turn_id == getattr(tokenizer, 'unk_token_id', None): | |
| raise ValueError('Tokenizer must define <turn|>.') | |
| generation_config.turn_end_token_id = int(turn_id) | |
| if generation_config.eos_token_id is None: | |
| generation_config.eos_token_id = [value for value in [tokenizer.eos_token_id] if value is not None] | |
| tool_id = tokenizer.convert_tokens_to_ids('<|tool_response>') | |
| if tool_id is not None and tool_id != getattr(tokenizer, 'unk_token_id', None): | |
| eos = generation_config.eos_token_id | |
| eos = [eos] if isinstance(eos, int) else list(eos or ()) | |
| if int(tool_id) not in eos: | |
| generation_config.eos_token_id = [*eos, int(tool_id)] | |
| generation_config.repetition_penalty_exclude_token_ids = list(dict.fromkeys( | |
| int(value) for value in [*generation_config.repetition_penalty_exclude_token_ids, | |
| *(getattr(tokenizer, 'all_special_ids', None) or ())] | |
| if value is not None | |
| )) | |
| generation_config.validate() | |
| return generation_config | |