Xunzhuo's picture
Release measured v1.2 update
d330a0b verified
Raw History Blame
9.42 kB
from __future__ import annotations
import hashlib
import importlib.util
import inspect
import json
import math
from pathlib import Path
import warnings
def _sha(path):
digest = hashlib.sha256()
with path.open('rb') as stream:
for block in iter(lambda: stream.read(8 << 20), b''):
digest.update(block)
return digest.hexdigest()
def _resolve(source, revision, local_files_only, cache_dir):
path = Path(source).expanduser()
if path.is_dir():
return path.resolve()
if path.is_absolute() or str(source).startswith(('.', '~')):
raise FileNotFoundError(f'Local model directory does not exist: {source}')
if not revision:
raise ValueError('A Hugging Face repo requires revision=; use a commit SHA for reproducibility.')
try:
from huggingface_hub import snapshot_download
except ImportError as exc:
raise RuntimeError('Hub loading requires huggingface_hub; install this package with [hub].') from exc
return Path(snapshot_download(repo_id=str(source), revision=revision,
local_files_only=local_files_only, cache_dir=cache_dir)).resolve()
def _verify_bundle(path):
manifest_path = path / 'bundle-manifest.json'
manifest = json.loads(manifest_path.read_text())
if manifest.get('format') != 'research-pointer-bundle-v1':
raise ValueError('Unsupported Decision bundle format')
seen = set()
for item in manifest['files']:
relative = Path(item['file'])
if relative.is_absolute() or '..' in relative.parts or item['file'] in seen:
raise ValueError('Invalid or duplicate bundle manifest path')
seen.add(item['file'])
file = path / relative
if not file.is_file() or file.stat().st_size != item['bytes'] or _sha(file) != item['sha256']:
raise ValueError(f'Bundle integrity check failed: {item["file"]}')
required = {'code/decision_api.py', 'code/decision_model.py', 'decision_config.json',
'decision_head.safetensors', 'temperature.json', 'runtime.json',
'backbone/config.json', 'tokenizer.json', 'tokenizer_config.json'}
if not required <= seen or not any(name.startswith('backbone/') and name.endswith('.safetensors') for name in seen):
raise ValueError('Bundle manifest omits required inference files')
return manifest
def _runtime_report(expected, device, allow_unvalidated_runtime):
if not str(device).startswith('cuda'):
raise RuntimeError('CPU/MPS inference is not implemented by this engine; use a GPU cuda device, including ROCm.')
try:
import torch
import transformers
import fla
import tokenizers
import safetensors
import triton
except ImportError as exc:
raise RuntimeError('Install the inference runtime recorded in runtime.json before loading weights; '
'pip install . installs only the lightweight wrapper.') from exc
if not torch.cuda.is_available():
raise RuntimeError('This release supports GPU inference through PyTorch cuda devices, including ROCm. '
'CPU/MPS inference is not implemented by the packaged numerical engine.')
from transformers.models.qwen3_5 import modeling_qwen3_5 as qwen
closure = inspect.getclosurevars(qwen.torch_chunk_gated_delta_rule).nonlocals
implementation = closure.get('implementation')
dispatch = getattr(implementation, '__module__', None)
current = {'torch': str(torch.__version__), 'hip': torch.version.hip,
'transformers': transformers.__version__, 'fla': fla.__version__,
'tokenizers': tokenizers.__version__, 'safetensors': safetensors.__version__,
'triton': triton.__version__, 'gated_delta': dispatch}
differences = {key: {'expected': expected.get(key), 'actual': value}
for key, value in current.items() if expected.get(key) != value}
if not closure.get('is_new_implementation'):
differences['gdn_new_implementation'] = {'expected': True, 'actual': False}
if differences:
message = 'Runtime differs from the validated bundle: ' + json.dumps(differences, sort_keys=True)
if not allow_unvalidated_runtime:
raise RuntimeError(message + '; explicitly set allow_unvalidated_runtime=True for exploratory use.')
warnings.warn(message + '. Numerical or runtime compatibility is not established.', RuntimeWarning, stacklevel=2)
return {'actual': current, 'differences': differences, 'matches_validated_runtime': not differences}
def _load_api(path):
api_path = path / 'code/decision_api.py'
spec = importlib.util.spec_from_file_location('decision_bundle_api_' + _sha(api_path)[:16], api_path)
api = importlib.util.module_from_spec(spec)
spec.loader.exec_module(api)
return api
class DecisionModel:
"""Load an exported bundle without modifying its prompt, head, or calibration."""
def __init__(self, engine, bundle_path, manifest, runtime):
self._engine = engine
self.bundle_path = bundle_path
self.manifest = manifest
self.runtime = runtime
@classmethod
def from_pretrained(cls, source, *, revision=None, device='cuda:0', local_files_only=False,
cache_dir=None, model_name=None, max_length=None,
allow_unvalidated_runtime=False):
"""Local paths never download; Hub sources require an explicit revision.
The manifest fixes the batch size and maximum complete-question length.
A lower max_length may be selected to reject larger requests. Calibration
and parameter dtypes come from the bundle, with no caller override.
"""
path = _resolve(source, revision, local_files_only, cache_dir)
manifest = _verify_bundle(path)
limit = manifest['input_length_limit']
maximum = limit if max_length is None else max_length
if isinstance(maximum, bool) or not isinstance(maximum, int) or not 1 <= maximum <= limit:
raise ValueError(f'max_length must be an integer in 1..{limit}')
config = json.loads((path / 'decision_config.json').read_text())
calibration = json.loads((path / 'temperature.json').read_text())
temperatures = calibration['temperatures']
if set(temperatures) != {'choice', 'noul', 'score'} or any(
isinstance(v, bool) or not isinstance(v, (int, float)) or not math.isfinite(v) or v <= 0
for v in temperatures.values()):
raise ValueError('Bundle must contain finite positive temperatures for all three types')
api = _load_api(path)
expected_runtime = json.loads((path / 'runtime.json').read_text())
if expected_runtime.get('normalization_profile') is not None:
if not hasattr(api, 'prepare_runtime_profile'):
raise RuntimeError('Profiled bundle omits its automatic runtime entrypoint')
profile = api.prepare_runtime_profile(path, device=device)
else:
import sys
if '_decision_process_normalization_profile_v1' in sys.modules:
raise RuntimeError('Use separate processes for profiled and unprofiled models')
profile = None
runtime = _runtime_report(expected_runtime, device, allow_unvalidated_runtime)
runtime['normalization_profile'] = profile
names = {'Qwen/Qwen3.5-2B': 'Decision-1.0-Sol', 'Qwen/Qwen3.5-4B': 'Decision-1.0-Nox'}
name = model_name or config.get('model_name') or names.get(config.get('base_model'), path.name)
engine = api.DecisionEngine(path, path / 'code', device=device, max_length=maximum,
batch_size=manifest['production_batch_size'], temperatures=temperatures,
model_name=name)
import torch
if {p.dtype for p in engine.model.backbone.parameters()} != {torch.bfloat16}:
raise RuntimeError('Backbone must remain BF16')
if {p.dtype for p in engine.model.head.parameters()} != {torch.float32}:
raise RuntimeError('Candidate head must remain FP32')
return cls(engine, path, manifest, runtime)
def decide(self, state, questions):
"""Return native Choice, Noul, and Score answers from the frozen engine.
Every question includes the complete state, instructions and candidate
descriptions. Any overflowing question raises ValueError before forward;
there is no truncation. Question names and candidate order are preserved.
"""
if not isinstance(questions, dict) or not questions:
raise ValueError('questions must be a nonempty mapping')
if not all(isinstance(name, str) and isinstance(question, dict) for name, question in questions.items()):
raise ValueError('Question names must be strings and questions must be mappings')
# Reject non-JSON inputs/NaN instead of accepting implementation-dependent text.
try:
json.dumps({'state': state, 'questions': questions}, ensure_ascii=False, allow_nan=False)
except (TypeError, ValueError) as exc:
raise ValueError('state and questions must be finite JSON-compatible values') from exc
return self._engine.decide(state, questions)