"""Clef: a multimodal Qwen backbone with a joint schema head for typed decisions. A record provides a ``state`` (any JSON value), optional ``images`` and ``videos``, and ``questions``. Each question has a ``type`` (``noul``, ``choice``, or ``score``), ``instructions``, and, for ``choice`` and ``score``, ``criteria`` describing the allowed options. The model returns one logit per allowed option for every question. ``systemone`` answers a Jev/SystemOne ``/v1/systemone`` request body with the same response body. """ from __future__ import annotations import json import math from dataclasses import dataclass from dataclasses import field as dataclass_field from pathlib import Path from typing import Any import torch import torch.nn.functional as functional SYSTEM_PROMPT = ( "Read the complete state and schema. Decide every field jointly. Each answer " "must be exactly one of that field's allowed options." ) IMAGE_PLACEHOLDER = "<|vision_start|><|image_pad|><|vision_end|>" VIDEO_PLACEHOLDER = "<|vision_start|><|video_pad|><|vision_end|>" MEDIA_BATCH_KEYS = ("pixel_values", "image_grid_thw", "pixel_values_videos", "video_grid_thw") MEDIA_TOKEN_KEYS = ("mm_token_type_ids",) QUESTION_TYPES = {"noul": 0, "choice": 1, "score": 2} def render(value: Any) -> str: if isinstance(value, str): return value return json.dumps( value, ensure_ascii=False, separators=(",", ":"), sort_keys=True, ) def question_options(question: dict[str, Any]) -> list[tuple[str, Any]]: question_type = str(question["type"]) if question_type == "noul": criteria = { "true": "The proposition is true or the answer is yes.", "false": "The proposition is false or the answer is no.", } criteria.update(question.get("criteria") or {}) return [(key, criteria[key]) for key in ("true", "false")] if question_type == "choice": return sorted((str(key), value) for key, value in question["criteria"].items()) return [(str(index), value) for index, value in enumerate(question["criteria"])] @dataclass(frozen=True) class EncodedQuestion: question_id: str question_type: int question_span: tuple[int, int] option_spans: tuple[tuple[int, int], ...] option_ids: tuple[str, ...] @dataclass(frozen=True) class EncodedRecord: input_ids: tuple[int, ...] questions: tuple[EncodedQuestion, ...] record_id: str media: dict[str, Any] | None = dataclass_field(default=None, compare=False, repr=False) def _tokens(tokenizer: Any, text: str) -> list[int]: return tokenizer(text, add_special_tokens=False).input_ids def _encode_media(processor: Any, record: dict[str, Any]) -> tuple[list[int], dict[str, Any] | None]: images = list(record.get("images") or []) videos = list(record.get("videos") or []) if not images and not videos: return [], None if processor is None: raise ValueError("records with images or videos require a processor") text = IMAGE_PLACEHOLDER * len(images) + VIDEO_PLACEHOLDER * len(videos) + "\n" encoded = processor( text=[text], images=images or None, videos=videos or None, return_tensors="pt", **(record.get("media_kwargs") or {}), ) media = {key: encoded[key] for key in MEDIA_BATCH_KEYS if key in encoded} for key in MEDIA_TOKEN_KEYS: if key in encoded: media[key] = encoded[key][0].tolist() return encoded["input_ids"][0].tolist(), media def encode_record( tokenizer: Any, record: dict[str, Any], max_length: int = 16384, max_state_tokens: int | None = None, processor: Any | None = None, ) -> EncodedRecord: schema_ids = _tokens(tokenizer, "\n\nSCHEMA FIELDS:\n") questions: list[EncodedQuestion] = [] for question_index, (question_id, question) in enumerate(record["questions"].items()): schema_ids.extend( _tokens( tokenizer, f"\nFIELD {question_index + 1}\nID: {question_id}\nTYPE: {question['type']}\nINSTRUCTION: ", ) ) question_start = len(schema_ids) instructions = question.get("instructions") if instructions is None or instructions == "": instructions = str(question_id) schema_ids.extend(_tokens(tokenizer, render(instructions))) question_end = len(schema_ids) schema_ids.extend(_tokens(tokenizer, "\nALLOWED OPTIONS:\n")) option_spans: list[tuple[int, int]] = [] option_ids: list[str] = [] for option_index, (option_id, description) in enumerate(question_options(question)): schema_ids.extend(_tokens(tokenizer, f"OPTION {option_index + 1}: ")) option_start = len(schema_ids) semantics = {"option_id": option_id} if description is not None: semantics["description"] = description schema_ids.extend(_tokens(tokenizer, render(semantics))) option_spans.append((option_start, len(schema_ids))) option_ids.append(option_id) schema_ids.extend(_tokens(tokenizer, "\n")) schema_ids.extend(_tokens(tokenizer, "END FIELD\n")) questions.append( EncodedQuestion( question_id=str(question_id), question_type=QUESTION_TYPES[str(question["type"])], question_span=(question_start, question_end), option_spans=tuple(option_spans), option_ids=tuple(option_ids), ) ) prefix_ids = _tokens( tokenizer, f"<|im_start|>system\n{SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\nSTATE:\n", ) suffix_ids = _tokens( tokenizer, "\n<|im_end|>\n<|im_start|>assistant\n\n\n\n\nJOINT SCHEMA DECISIONS:", ) media_ids, media = _encode_media(processor, record) if media is not None: media["token_offset"] = len(prefix_ids) prefix_ids = prefix_ids + media_ids state_ids = _tokens(tokenizer, render(record["state"])) if max_state_tokens is not None: state_ids = state_ids[:max_state_tokens] fixed_length = len(prefix_ids) + len(schema_ids) + len(suffix_ids) if fixed_length > max_length: raise ValueError( f"schema requires {fixed_length} tokens before state; maximum is {max_length}" ) state_ids = state_ids[: max_length - fixed_length] schema_offset = len(prefix_ids) + len(state_ids) shifted_questions = tuple( EncodedQuestion( question_id=question.question_id, question_type=question.question_type, question_span=( question.question_span[0] + schema_offset, question.question_span[1] + schema_offset, ), option_spans=tuple( (start + schema_offset, end + schema_offset) for start, end in question.option_spans ), option_ids=question.option_ids, ) for question in questions ) input_ids = tuple(prefix_ids + state_ids + schema_ids + suffix_ids) if not input_ids or not shifted_questions: raise ValueError("record produced no model input or questions") return EncodedRecord( input_ids=input_ids, questions=shifted_questions, record_id=str(record.get("id", "unknown")), media=media, ) def collate_records( records: list[EncodedRecord], pad_token_id: int, device: torch.device, ) -> dict[str, Any]: maximum_length = max(len(record.input_ids) for record in records) input_ids = torch.full( (len(records), maximum_length), pad_token_id, dtype=torch.long, device=device, ) attention_mask = torch.zeros( (len(records), maximum_length), dtype=torch.long, device=device, ) for index, record in enumerate(records): length = len(record.input_ids) input_ids[index, :length] = torch.tensor(record.input_ids, device=device) attention_mask[index, :length] = 1 media: dict[str, torch.Tensor] = {} for key in MEDIA_BATCH_KEYS: values = [record.media[key] for record in records if record.media and key in record.media] if values: media[key] = torch.cat(values, dim=0).to(device) for key in MEDIA_TOKEN_KEYS: if any(record.media and key in record.media for record in records): token_values = torch.zeros((len(records), maximum_length), dtype=torch.long, device=device) for index, record in enumerate(records): if record.media and key in record.media: offset = record.media["token_offset"] values = torch.tensor(record.media[key], dtype=torch.long, device=device) token_values[index, offset : offset + len(values)] = values media[key] = token_values return { "input_ids": input_ids, "attention_mask": attention_mask, "records": records, "media": media, } class EvidenceRoutingLayer(torch.nn.Module): def __init__( self, width: int, heads: int, feedforward: int, dropout: float = 0.0, ) -> None: super().__init__() self.query_norm = torch.nn.LayerNorm(width) self.memory_norm = torch.nn.LayerNorm(width) self.attention = torch.nn.MultiheadAttention( width, heads, dropout=dropout, batch_first=True, ) self.attention_dropout = torch.nn.Dropout(dropout) self.feedforward_norm = torch.nn.LayerNorm(width) self.feedforward = torch.nn.Sequential( torch.nn.Linear(width, feedforward), torch.nn.GELU(), torch.nn.Dropout(dropout), torch.nn.Linear(feedforward, width), torch.nn.Dropout(dropout), ) def forward(self, queries: torch.Tensor, memory: torch.Tensor) -> torch.Tensor: normalized_queries = self.query_norm(queries) routed, _ = self.attention( normalized_queries, self.memory_norm(memory), self.memory_norm(memory), need_weights=False, ) queries = queries + self.attention_dropout(routed) return queries + self.feedforward(self.feedforward_norm(queries)) class JointSchemaHead(torch.nn.Module): def __init__( self, hidden_size: int, width: int, routing_layers: int, layers: int, heads: int, feedforward: int, dropout: float = 0.0, ) -> None: super().__init__() self.hidden_norm = torch.nn.LayerNorm(hidden_size) self.memory_projection = torch.nn.Linear(hidden_size, width, bias=False) self.question_projection = torch.nn.Linear(hidden_size, width, bias=False) self.option_question_projection = torch.nn.Linear(hidden_size, width, bias=False) self.global_projection = torch.nn.Linear(hidden_size, width, bias=False) self.option_context_projection = torch.nn.Linear(hidden_size, width, bias=False) self.option_lexical_projection = torch.nn.Linear(hidden_size, width, bias=False) self.type_embedding = torch.nn.Embedding(3, width) self.evidence_layers = torch.nn.ModuleList( [ EvidenceRoutingLayer( width=width, heads=heads, feedforward=feedforward, dropout=dropout, ) for _ in range(routing_layers) ] ) self.option_summary_norm = torch.nn.LayerNorm(width) self.layers = torch.nn.ModuleList( [ torch.nn.TransformerDecoderLayer( d_model=width, nhead=heads, dim_feedforward=feedforward, dropout=dropout, activation="gelu", batch_first=True, norm_first=True, ) for _ in range(layers) ] ) self.field_norm = torch.nn.LayerNorm(width) self.option_norm = torch.nn.LayerNorm(width) self.residual_scorer = torch.nn.Sequential( torch.nn.Linear(width * 4, width), torch.nn.GELU(), torch.nn.Dropout(dropout), torch.nn.Linear(width, 1), ) self.prior_logit_scale = torch.nn.Parameter(torch.zeros(())) self.joint_logit_scale = torch.nn.Parameter(torch.zeros(())) self.residual_gate = torch.nn.Parameter(torch.zeros(())) @staticmethod def _mean_span(values: torch.Tensor, span: tuple[int, int]) -> torch.Tensor: start, end = span return values[start:end].mean(dim=0) def forward( self, hidden_states: torch.Tensor, input_ids: torch.Tensor, attention_mask: torch.Tensor, records: list[EncodedRecord], output_embedding_weight: torch.Tensor, ) -> list[list[torch.Tensor]]: results: list[list[torch.Tensor]] = [] normalized_hidden = self.hidden_norm(hidden_states) for batch_index, record in enumerate(records): sequence_length = int(attention_mask[batch_index].sum().item()) sequence_hidden = normalized_hidden[batch_index, :sequence_length] memory = self.memory_projection(sequence_hidden).unsqueeze(0) global_vector = sequence_hidden[-1] question_vectors = torch.stack( [ self._mean_span(sequence_hidden, question.question_span) for question in record.questions ] ) type_ids = torch.tensor( [question.question_type for question in record.questions], device=hidden_states.device, ) option_contexts: list[torch.Tensor] = [] lexical_options: list[torch.Tensor] = [] option_counts = [] for question in record.questions: context_vectors = torch.stack( [ self._mean_span(sequence_hidden, span) for span in question.option_spans ] ) lexical_vectors = [] for start, end in question.option_spans: token_ids = input_ids[batch_index, start:end] lexical_vectors.append(output_embedding_weight[token_ids].mean(dim=0)) lexical = torch.stack(lexical_vectors) option_contexts.append(context_vectors) lexical_options.append(lexical) option_counts.append(len(question.option_spans)) option_queries = [] for question_index, (context_vectors, lexical) in enumerate( zip(option_contexts, lexical_options) ): option_queries.append( self.option_context_projection(context_vectors) + self.option_lexical_projection(lexical) + self.option_question_projection( question_vectors[question_index] ).unsqueeze(0) ) routed_options = torch.cat(option_queries, dim=0).unsqueeze(0) for layer in self.evidence_layers: routed_options = layer(routed_options, memory) routed_options = routed_options[0] split_options = list(torch.split(routed_options, option_counts, dim=0)) base_fields = self.question_projection(question_vectors) option_summaries = [] for field, options in zip(base_fields, split_options): routing_weights = torch.softmax( torch.matmul(options, field) / math.sqrt(options.shape[-1]), dim=0, ) option_summaries.append( torch.sum(routing_weights.unsqueeze(-1) * options, dim=0) ) fields = ( base_fields + self.option_summary_norm(torch.stack(option_summaries)) + self.global_projection(global_vector).unsqueeze(0) + self.type_embedding(type_ids) ) fields = fields.unsqueeze(0) for layer in self.layers: fields = layer(fields, memory) fields = self.field_norm(fields[0]) record_logits: list[torch.Tensor] = [] for field, question, lexical, routed in zip( fields, record.questions, lexical_options, split_options, ): anchor = functional.normalize( question_vectors[len(record_logits)] + global_vector, dim=-1, ) lexical_anchor = functional.normalize(lexical, dim=-1) prior_scale = self.prior_logit_scale.clamp(max=math.log(100.0)).exp() prior = prior_scale * torch.matmul(lexical_anchor, anchor) options = self.option_norm(routed) repeated_field = field.unsqueeze(0).expand_as(options) cosine = functional.cosine_similarity(repeated_field, options, dim=-1) features = torch.cat( [ repeated_field, options, repeated_field * options, torch.abs(repeated_field - options), ], dim=-1, ) residual = self.residual_scorer(features).squeeze(-1) joint_scale = self.joint_logit_scale.clamp(max=math.log(100.0)).exp() joint = joint_scale * cosine + residual record_logits.append( prior + torch.sigmoid(self.residual_gate) * joint ) results.append(record_logits) return results class ClefModel(torch.nn.Module): def __init__(self, language_model: Any, head: JointSchemaHead) -> None: super().__init__() self.language_model = language_model self.head = head def forward(self, batch: dict[str, Any]) -> list[list[torch.Tensor]]: base_model = ( self.language_model.get_base_model() if hasattr(self.language_model, "get_base_model") else self.language_model ) media = batch.get("media") or {} text_model = base_model.model if not media and hasattr(text_model, "language_model"): text_model = text_model.language_model outputs = text_model( input_ids=batch["input_ids"], attention_mask=batch["attention_mask"], use_cache=False, return_dict=True, **media, ) return self.head( outputs.last_hidden_state, batch["input_ids"], batch["attention_mask"], batch["records"], base_model.get_output_embeddings().weight, ) def load_release_model( model_path: str | Path, device: str | torch.device = "cuda", dtype: torch.dtype = torch.bfloat16, **from_pretrained_kwargs: Any, ) -> tuple[ClefModel, Any]: """Load a Clef release (merged backbone, joint schema head, and processor).""" from huggingface_hub import snapshot_download from safetensors.torch import load_file from transformers import AutoProcessor, Qwen3_5ForConditionalGeneration path = Path(model_path) if not path.is_dir(): path = Path(snapshot_download(str(model_path))) backbone = Qwen3_5ForConditionalGeneration.from_pretrained( path, dtype=dtype, device_map={"": str(device)}, **from_pretrained_kwargs, ) backbone.config.use_cache = False head_config = json.loads((path / "joint_head_config.json").read_text()) head = JointSchemaHead(**head_config) head.load_state_dict(load_file(path / "joint_head.safetensors"), strict=True) head = head.to(device=device, dtype=dtype) processor = AutoProcessor.from_pretrained(path) return ClefModel(backbone, head).eval(), processor def systemone_answer(question: dict[str, Any], probabilities: dict[str, float]) -> dict[str, Any]: """Convert per-option probabilities for one question into a SystemOne answer.""" if question["type"] == "noul": return {"type": "noul", "noul": round(probabilities["true"], 4)} if question["type"] == "choice": options = [str(option) for option in question["criteria"]] choice = max(options, key=probabilities.__getitem__) return { "type": "choice", "choice": choice, "confidence": round(probabilities[choice], 4), "probabilities": {option: round(probabilities[option], 4) for option in options}, } levels = [str(index) for index in range(len(question["criteria"]))] return { "type": "score", "score": round(sum(index * probabilities[level] for index, level in enumerate(levels)), 4), "confidence": round(max(probabilities[level] for level in levels), 4), "legend": dict(zip(levels, question["criteria"])), "probabilities": {level: round(probabilities[level], 4) for level in levels}, } @torch.inference_mode() def systemone(model: ClefModel, processor: Any, request: dict[str, Any], max_length: int = 16384) -> dict[str, Any]: """Answer a Jev/SystemOne ``/v1/systemone`` request body with a SystemOne response body. The request has ``model``, ``state``, and ``questions``, plus optional ``images`` and ``videos``. """ questions = request.get("questions") if not isinstance(request.get("model"), str) or "state" not in request: raise ValueError("model and state are required") if not isinstance(questions, dict) or not questions: raise ValueError("at least one question is required") for question_id, question in questions.items(): if question.get("type") not in QUESTION_TYPES: raise ValueError(f"{question_id}: type must be noul, choice, or score") if question["type"] != "noul" and not question.get("criteria"): raise ValueError(f"{question_id}: criteria must not be empty") encoded = encode_record(processor.tokenizer, request, max_length=max_length, processor=processor) device = next(model.parameters()).device logits = model(collate_records([encoded], processor.tokenizer.pad_token_id, device))[0] answers = { question.question_id: systemone_answer( questions[question.question_id], dict(zip(question.option_ids, question_logits.float().softmax(-1).tolist())), ) for question, question_logits in zip(encoded.questions, logits) } return { "model": request["model"], "answers": answers, "usage": {"input_tokens": len(encoded.input_ids), "output_tokens": 0}, }