clef-flash-EXL3 / clef_exl3.py
ramgpt's picture
Add files using upload-large-folder tool
ff690e8 verified
Raw History Blame Contribute Delete
3.88 kB
#!/usr/bin/env python3
import importlib.util
import json
import sys
from pathlib import Path
from types import SimpleNamespace
import torch
from exllamav3 import Config, Model, Tokenizer
from safetensors.torch import load_file
def load_joint_module(path):
spec = importlib.util.spec_from_file_location("clef_joint_schema_model", path / "joint_schema_model.py")
module = importlib.util.module_from_spec(spec)
sys.modules[spec.name] = module
spec.loader.exec_module(module)
return module
class TokenizerAdapter:
def __init__(self, tokenizer):
self.tokenizer = tokenizer
def __call__(self, text, add_special_tokens=False):
ids = self.tokenizer.encode(text, encode_special_tokens=True)
return SimpleNamespace(input_ids=ids[0].tolist())
class ClefEXL3:
def __init__(self, model_path, device="cuda:0"):
self.path = Path(model_path)
self.device = torch.device(device)
self.joint = load_joint_module(self.path)
config = Config.from_directory(str(self.path))
self.tokenizer = TokenizerAdapter(Tokenizer.from_config(config))
self.model = Model.from_config(config, component="text")
self.model.load(progressbar=True, device=device)
head_config = json.loads((self.path / "joint_head_config.json").read_text())
self.head = self.joint.JointSchemaHead(**head_config)
self.head.load_state_dict(load_file(self.path / "joint_head.safetensors"), strict=True)
self.head = self.head.to(device=self.device, dtype=torch.float16).eval()
lm_head = self.model.modules[self.model.logit_layer_idx]
if getattr(lm_head, "quant_type", None) != "fp16":
raise RuntimeError("Clef EXL3 requires an FP16 LM head; quantize with --head_bits 16")
raw = json.loads((self.path / "config.json").read_text())
text_config = raw.get("text_config", raw)
vocab_size = int(text_config["vocab_size"])
self.output_embedding_weight = lm_head.inner.weight.T[:vocab_size]
@torch.inference_mode()
def hidden_states(self, input_ids):
params = {"attn_mode": "flash_attn_nc"}
x = self.model.prepare_inputs(input_ids, params)
for module, instance, _ in self.model.fwd_modules:
if module.caps.get("logits_output"):
break
params["layer_instance"] = instance
x = module.prepare_for_device(x, params)
x = module.forward(x, params)
return x
@torch.inference_mode()
def systemone(self, request, max_length=16384):
if request.get("images") or request.get("videos"):
raise NotImplementedError("Clef EXL3 adapter currently supports text/JSON state only")
encoded = self.joint.encode_record(
self.tokenizer,
request,
max_length=max_length,
processor=None,
)
input_ids = torch.tensor([encoded.input_ids], dtype=torch.long, device=self.device)
attention_mask = torch.ones_like(input_ids)
hidden = self.hidden_states(input_ids).to(torch.float16)
logits = self.head(
hidden,
input_ids,
attention_mask,
[encoded],
self.output_embedding_weight,
)[0]
questions = request["questions"]
answers = {}
for question, question_logits in zip(encoded.questions, logits):
probabilities = dict(zip(question.option_ids, question_logits.float().softmax(-1).tolist()))
answers[question.question_id] = self.joint.systemone_answer(
questions[question.question_id], probabilities
)
return {
"model": request.get("model", "clef-flash-exl3"),
"answers": answers,
"usage": {"input_tokens": len(encoded.input_ids), "output_tokens": 0},
}