"""Portable, single-request merged Qwen3.5 decision reference runtime.""" import json from pathlib import Path import torch import transformers from .contracts import PROMPT_VERSION, format_response, label_mapping, render_prompt from .early_exit import QwenEarlyExit class OpenJet: """Explicit low/high, not automatic routing. Text-only; never truncates.""" def __init__(self, model, tokenizer, depth_config, max_length=8192): if transformers.__version__ != "5.16.1": raise RuntimeError("Layer execution is audited for transformers==5.16.1") self.model = model.eval() self.tokenizer = tokenizer self.wrapper = QwenEarlyExit(self.model) self.device = self.model.get_input_embeddings().weight.device self.max_length = max_length if type(max_length) is not int or not 1 <= max_length <= 8192: raise ValueError("max_length must be an integer in 1..8192") if depth_config.get("prompt_version") != PROMPT_VERSION: raise ValueError("checkpoint prompt version does not match runtime") full = depth_config.get("full_depth") low = depth_config.get("exit_depth") if full != self.wrapper.full_depth: raise ValueError("checkpoint depth and model depth disagree") if type(low) is not int or not 0 < low < full: raise ValueError("checkpoint does not declare a trained shallow exit") if self.wrapper.backbone.config.layer_types[low - 1] != "full_attention": raise ValueError("shallow exit must be a full-attention boundary") self.depths = {"low": low, "high": full} @classmethod def from_pretrained(cls, directory, device="cuda:0", dtype="bfloat16"): """Load a local HF snapshot (download explicitly with a pinned revision).""" from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer directory = Path(directory) if (directory / "adapter_config.json").exists(): raise ValueError("Expected a merged snapshot, not an adapter directory") if dtype not in ("float32", "bfloat16"): raise ValueError("supported dtypes: float32, bfloat16") config = AutoConfig.from_pretrained(directory, local_files_only=True) loader = AutoModelForCausalLM if config.model_type == "qwen3_5": from transformers import Qwen3_5ForConditionalGeneration loader = Qwen3_5ForConditionalGeneration elif config.model_type != "qwen3_5_text": raise ValueError("Expected Qwen3.5 text or conditional-generation model") model = loader.from_pretrained( directory, local_files_only=True, dtype=getattr(torch, dtype), attn_implementation="sdpa", ).to(device) tokenizer = AutoTokenizer.from_pretrained(directory, local_files_only=True) depth_config = json.loads((directory / "depth_config.json").read_text()) return cls(model, tokenizer, depth_config) def _depth(self, effort): if effort not in self.depths: raise ValueError("effort must be low or high") return self.depths[effort] def compile(self, request): """Exact chat/no-thinking contract used in original native evaluation.""" mapping = label_mapping(request) prompt = self.tokenizer.apply_chat_template( [{"role": "user", "content": render_prompt(request)}], tokenize=False, add_generation_prompt=True, enable_thinking=False, ) ids = self.tokenizer.encode(prompt, add_special_tokens=False) if not ids or len(ids) > self.max_length: raise ValueError("input exceeds runtime limit; no truncation permitted") for name in ( "image_token_id", "video_token_id", "vision_start_token_id", "vision_end_token_id", ): token = getattr(self.model.config, name, None) if token is not None and token in ids: raise ValueError("multimodal placeholders are unsupported") tokens = [] for label in mapping: token = self.tokenizer.encode(label, add_special_tokens=False) joint = self.tokenizer.encode(prompt + label, add_special_tokens=False) if len(token) != 1 or joint != ids + token: raise ValueError( "candidate label is not single-token at answer boundary" ) if token[0] in self.tokenizer.all_special_ids: raise ValueError("candidate label must not be special token") tokens.append(token[0]) if len(set(tokens)) != len(tokens): raise ValueError("candidate token IDs must be unique") return ids, tokens @torch.inference_mode() def decide(self, request, effort="high"): depth = self._depth(effort) ids, candidates = self.compile(request) input_ids = torch.tensor([ids], dtype=torch.long, device=self.device) candidate_ids = torch.tensor(candidates, dtype=torch.long, device=self.device) if effort == "high": # Match the original reference high/full-vocabulary head path. output = self.model( input_ids=input_ids, attention_mask=torch.ones_like(input_ids), position_ids=torch.arange(len(ids), device=self.device).unsqueeze(0), past_key_values=None, use_cache=True, return_dict=True, logits_to_keep=1, ) logits = output.logits[0, -1].index_select(0, candidate_ids).float() projection = "full_head" else: (decision,) = self.wrapper(input_ids, candidate_ids, depths=(depth,)) logits = decision.logits[0].float() projection = decision.projection_mode response = format_response(request, logits.softmax(-1).tolist()) response.update( effort=effort, executed_layers=depth, prompt_tokens=len(ids), logits=logits.tolist(), projection=projection, calibrated=False, ) return response @torch.inference_mode() def generate_text(self, user_text, effort="high", max_new_tokens=128): """TYPE greedy reference; replays prefix each token, not optimized serving.""" depth = self._depth(effort) if not isinstance(user_text, str) or not user_text.strip(): raise ValueError("user_text must be nonempty") if type(max_new_tokens) is not int or max_new_tokens < 1: raise ValueError("max_new_tokens must be positive integer") ids = self.tokenizer.apply_chat_template( [{"role": "user", "content": user_text}], tokenize=True, add_generation_prompt=True, enable_thinking=False, return_dict=False, ) if not ids or len(ids) + max_new_tokens > self.max_length: raise ValueError("prompt plus generation reservation exceeds limit") eos = self.model.generation_config.eos_token_id eos = [eos] if isinstance(eos, int) else list(eos or []) generated = [] for _ in range(max_new_tokens): tensor = torch.tensor([ids + generated], device=self.device) state = self.wrapper.advance(self.wrapper.begin(tensor), depth) hidden = self.wrapper.backbone.norm(state.hidden[:, -1]) token = self.model.get_output_embeddings()(hidden)[0].argmax().item() generated.append(token) if token in eos: break return { "text": self.tokenizer.decode(generated, skip_special_tokens=True), "token_ids": generated, "effort": effort, "executed_layers_per_token": depth, "finish_reason": "eos" if generated[-1] in eos else "length", "prompt_tokens": len(ids), }