File size: 3,555 Bytes
f5864f8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Portable NF4 loader for xuhaodev/Qwen3.5-4B-Jev."""
import hashlib
import json
from pathlib import Path

import torch
from decision_model import GeneralDecisionModel
from decision_schema import SystemOneRequest, answer
from huggingface_hub import snapshot_download


def verify(directory, files):
    root = Path(directory).resolve()
    for name, expected in files.items():
        path = (root / name).resolve()
        if not path.is_relative_to(root):
            raise ValueError(f"Invalid artifact path: {name}")
        with path.open('rb') as stream:
            if hashlib.file_digest(stream, 'sha256').hexdigest() != expected:
                raise ValueError(f"Artifact hash mismatch: {name}")


class JevModel:
    def __init__(self, directory, *, base_path=None, device='cuda'):
        self.directory = Path(directory)
        manifest = json.loads((self.directory / 'manifest.json').read_text())
        verify(self.directory, manifest['files'])
        self.config = json.loads((self.directory / 'config.json').read_text())
        self.temperatures = json.loads(
            (self.directory / 'calibration.json').read_text())['temperatures']
        if any(not 0 < t < float('inf') for t in self.temperatures.values()):
            raise ValueError('Invalid temperatures')
        if not str(device).startswith('cuda') or not torch.cuda.is_available():
            raise ValueError('This NF4/FLA release requires a supported NVIDIA CUDA GPU')
        if base_path is None:
            base_path = snapshot_download(
                self.config['base_model'], revision=self.config['base_revision'],
                allow_patterns=['*.json', '*.safetensors', '*.jinja', '*.txt', 'LICENSE'])
        verify(base_path, manifest['base_files'])
        self.model = GeneralDecisionModel(
            base_path, adapter=self.directory, mode='baseline', device=device).eval()

    @classmethod
    def from_pretrained(cls, repo_id='xuhaodev/Qwen3.5-4B-Jev', *, revision=None, **kwargs):
        directory = (repo_id if Path(repo_id).is_dir()
                     else snapshot_download(repo_id, revision=revision))
        return cls(directory, **kwargs)

    @torch.inference_mode()
    def predict(self, state, questions):
        request = SystemOneRequest(state=state, model=self.config['model_name'],
                                   questions=questions)
        if len(request.questions) > 32:
            raise ValueError('At most 32 questions per request')
        answers, tokens, prepared = {}, 0, []
        for name, question in request.questions.items():
            q = question.model_dump(exclude_none=True)
            inputs = self.model.encode(state, q, max_length=self.config['max_length'])
            count = sum(item['input_ids'].shape[-1] for item in inputs)
            if count > self.config['max_expanded_tokens']:
                raise ValueError('Question expanded token budget exceeded')
            tokens += count
            if tokens > self.config.get('max_request_expanded_tokens', 262144):
                raise ValueError('Request expanded token budget exceeded')
            prepared.append((name, q, inputs))
        for name, q, inputs in prepared:
            logits = self.model.logits(inputs)
            p = torch.softmax(logits.float() / self.temperatures[q['type']], -1)
            answers[name] = answer(q, p.cpu().tolist())
        return {'model': self.config['model_name'], 'answers': answers,
                'usage': {'input_tokens': tokens, 'output_tokens': 0}}

    system_one = predict