File size: 5,276 Bytes
fd3090c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
"""
Load the frozen Semantic-Lite-2 encoder and extract the two conditioning signals.

    Data A : (B, 256)      global semantic vector   -> decoder prefix token
    Data B : (B, L, 2048)  per-token hidden states  -> cross-attention key/value

WHY THIS FILE EXISTS
--------------------
`ukung/semantic-lite-2` ships custom modelling code (`model_type` is not
registered in transformers), so it cannot be loaded with a plain
`AutoModel.from_pretrained`. It also needs two source patches to run under
transformers 4.57.1. Both patches are applied here, at load time, on a copy of
the downloaded file. See NOTES.md for the full list of gotchas.

The encoder is ALWAYS frozen and ALWAYS in eval() mode. See `set_train_mode()`.
"""

import os
import subprocess
import sys

import torch
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from transformers import AutoTokenizer

ENCODER_ID = "ukung/semantic-lite-2"

# Where the patched custom code is staged. Must be importable.
STAGE_DIR = "/content/semantic_lite"
STAGE_PARENT = "/content"

# --- Encoder architecture (must match the published checkpoint) --------------
ENCODER_CONFIG = dict(
    vocab_size=131072,
    hidden_size=2048,
    intermediate_size=6656,
    num_hidden_layers=2,
    num_attention_heads=8,
    num_key_value_heads=2,
    head_dim=256,
    hidden_act="gelu",
    max_position_embeddings=1048576,
    initializer_range=0.0221,
    rms_norm_eps=1e-6,
    use_cache=True,
    pad_token_id=2,
    bos_token_id=0,
    eos_token_id=1,
    tie_word_embeddings=True,
    rope_parameters=None,  # set manually below; None bypasses the validator
    attention_bias=False,
    attention_dropout=0.0,
    mlp_bias=False,
    headwise_attn_output_gate=True,
    gate_attn_act_mode="sigmoid",
    sliding_window=512,
    layer_types=["sliding_attention", "sliding_attention"],
    num_attn_layers=2,
    out_dim=256,
)

ROPE_PARAMETERS = {
    "full_attention": {"partial_rotary_factor": 0.25, "rope_theta": 5000000},
    "sliding_attention": {"partial_rotary_factor": 1.0, "rope_theta": 10000},
}


def _stage_and_patch_custom_code(stage_dir=STAGE_DIR):
    """Download the encoder's custom code and apply the 4.57.1 compatibility patches."""
    os.makedirs(stage_dir, exist_ok=True)

    for fname in ("modeling_semantic_lite.py", "configuration_semantic_lite.py"):
        src = hf_hub_download(repo_id=ENCODER_ID, filename=fname)
        subprocess.run(["cp", src, os.path.join(stage_dir, fname)], check=True)

    path = os.path.join(stage_dir, "modeling_semantic_lite.py")
    with open(path) as f:
        code = f.read()

    # PATCH 1: create_causal_mask in 4.57.1 expects `input_embeds`, not `inputs_embeds`.
    code = code.replace('"inputs_embeds": inputs_embeds', '"input_embeds": inputs_embeds')

    # PATCH 2: 4.57.1 requires `cache_position` inside mask_kwargs.
    code = code.replace(
        '"past_key_values": past_key_values,\n                "position_ids": position_ids,',
        '"past_key_values": past_key_values,\n                "position_ids": position_ids,\n'
        '                "cache_position": cache_position,',
    )

    with open(path, "w") as f:
        f.write(code)

    # Drop any previously imported copy so the patched file is what gets imported.
    for mod in list(sys.modules):
        if mod.startswith("semantic_lite"):
            del sys.modules[mod]

    with open(os.path.join(stage_dir, "__init__.py"), "w") as f:
        f.write("from .configuration_semantic_lite import SemanticLiteConfig\n")
        f.write("from .modeling_semantic_lite import SemanticLiteEmbedder\n")

    if STAGE_PARENT not in sys.path:
        sys.path.insert(0, STAGE_PARENT)


def load_encoder(device="cuda"):
    """Return (encoder, tokenizer). Encoder is frozen and in eval() mode."""
    _stage_and_patch_custom_code()

    from semantic_lite.configuration_semantic_lite import SemanticLiteConfig
    from semantic_lite.modeling_semantic_lite import SemanticLiteEmbedder

    config = SemanticLiteConfig(**ENCODER_CONFIG)
    config.rope_parameters = ROPE_PARAMETERS

    tokenizer = AutoTokenizer.from_pretrained(ENCODER_ID)

    encoder = SemanticLiteEmbedder(config)
    state_dict = load_file(hf_hub_download(repo_id=ENCODER_ID, filename="model.safetensors"))
    encoder.load_state_dict(state_dict, strict=False)

    for param in encoder.parameters():
        param.requires_grad = False
    encoder.eval()

    return encoder.to(device), tokenizer


@torch.no_grad()
def extract_conditioning(encoder, input_ids, attention_mask):
    """Run the frozen encoder and return (data_a, data_b)."""
    data_a = encoder(input_ids=input_ids, attention_mask=attention_mask)  # (B, 256)
    data_b = encoder.backbone(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
    return data_a, data_b  # (B, 256), (B, L, 2048)


def set_train_mode(model):
    """
    Put the decoder in train mode while keeping the frozen encoder in eval().

    GOTCHA: `model.train()` recurses into every submodule, including the frozen
    encoder, whose attention heads carry dropout=0.1. That would inject noise
    into Data A / Data B on every step. Always call this instead of `.train()`.
    """
    model.train()
    model.encoder.eval()