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)