Download clef_exl3.py from ramgpt/clef-flash-EXL3: direct link, hf CLI and curl.
- Browser
- Download file 3.88 kB
-
https://huggingface.co/ramgpt/clef-flash-EXL3/resolve/main/clef_exl3.py
- Command line
-
hf download hf://ramgpt/clef-flash-EXL3/clef_exl3.py
-
curl -L -o clef_exl3.py https://huggingface.co/ramgpt/clef-flash-EXL3/resolve/main/clef_exl3.py
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] | |
| 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 | |
| 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}, | |
| } | |