ukung's picture
Add source code (encoder_loader, model, data, train, ablation, generate) + NOTES
fd3090c verified
Raw History Blame Contribute Delete
5.28 kB
"""
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()