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()
|