#!/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}, }