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