File size: 9,415 Bytes
43d6004 d330a0b 43d6004 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | 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)
|