gump2049's picture
Publish complete APUS-OpenJev-v1 models and Technical Report v1.1
68e5880 verified
Raw History Blame
8.08 kB
"""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),
}