File size: 4,346 Bytes
b88f761
53d5244
 
 
b88f761
 
53d5244
 
 
 
 
b88f761
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
53d5244
b88f761
53d5244
 
b88f761
53d5244
 
b88f761
 
 
53d5244
 
 
 
 
b88f761
 
53d5244
 
 
b88f761
53d5244
b88f761
53d5244
b88f761
 
 
 
 
 
 
 
53d5244
 
b88f761
 
 
 
53d5244
b88f761
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
"""Text-only schema25 inference configuration."""
from __future__ import annotations
from collections.abc import Sequence
from typing import Any
from transformers import PreTrainedConfig
from transformers.models.diffusion_gemma import DiffusionGemmaTextConfig

DENOISE_TEMPERATURE = 0.8

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):
    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 = 25,
        memory_architecture: str = "compact_gdn2_v2",
        latent_dim: int = 2816,
        latent_ffn_dim: int = 7168,
        latent_num_layers: int = 4,
        latent_num_heads: int = 16,
        latent_local_attention_window: int = 128,
        latent_tape_probes: int = 4,
        latent_history_kv_rank: int = 1024,
        latent_working_last_block_global: bool = True,
        working_memory_bus: bool = True,
        persistent_memory_bus: bool = True,
        commit_sequence_dim: int = 1024,
        kv_cache_bucket_size: int = 128,
        vocab_chunk_size: int = 32768,
        turn_end_token_id: int = 106,
        terminal_token_ids: Sequence[int] = (106, 50),
        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,
        **kwargs: Any,
    ) -> None:
        if state_schema_version != 25 or memory_architecture != "compact_gdn2_v2":
            raise ValueError("This release requires schema25 compact_gdn2_v2 weights.")
        if text_config is None:
            text_config = ModilifyMk2TextConfig()
        elif isinstance(text_config, dict):
            payload = dict(text_config)
            payload.pop("model_type", None)
            text_config = ModilifyMk2TextConfig(**payload)
        text_config.use_bidirectional_attention = None
        self.text_config = text_config
        self.canvas_length = canvas_length
        self.initializer_range = initializer_range
        self.state_schema_version = state_schema_version
        self.memory_architecture = memory_architecture
        self.latent_dim = latent_dim
        self.latent_ffn_dim = latent_ffn_dim
        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_tape_probes = latent_tape_probes
        self.latent_history_kv_rank = latent_history_kv_rank
        self.latent_working_last_block_global = latent_working_last_block_global
        self.working_memory_bus = working_memory_bus
        self.persistent_memory_bus = persistent_memory_bus
        self.commit_sequence_dim = commit_sequence_dim
        self.kv_cache_bucket_size = kv_cache_bucket_size
        self.vocab_chunk_size = vocab_chunk_size
        self.turn_end_token_id = turn_end_token_id
        self.terminal_token_ids = terminal_token_ids
        self.commit_failure_budget = commit_failure_budget
        self.commit_top_k = commit_top_k
        self.commit_min_p = commit_min_p
        self.commit_target_confidence = commit_target_confidence
        self.commit_entropy_weight = commit_entropy_weight
        self.commit_confidence_power = commit_confidence_power
        super().__init__(tie_word_embeddings=tie_word_embeddings,
                         eos_token_id=eos_token_id, **kwargs)