File size: 5,617 Bytes
c69aaec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""FP4 E2M1 packed weights, BF16 activations, unchanged native JEV API.

Portable backend dequantizes one matrix with Triton before BF16 GEMM. This
reduces resident weight memory but is not native FP4 Tensor Core GEMM.
Reference backend caches dequantized BF16 weights for numerical comparison.
"""
import json
import math
from pathlib import Path

import torch
from accelerate import init_empty_weights
from safetensors.torch import load_file
from transformers import AutoConfig, AutoProcessor
from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5Model
from kev.model import DecisionModel, MAX_OPTIONS, answer_codes


def unpack_weight(packed, scale, global_scale):
    table = torch.tensor([0,.5,1,1.5,2,3,4,6,0,-.5,-1,-1.5,-2,-3,-4,-6], device=packed.device)
    codes = torch.stack([packed & 15, packed >> 4], dim=-1).flatten(-2)
    shape = codes.shape
    values = table[codes.long()].reshape(shape[0], -1, 16)
    return (values * (scale.float() * global_scale).unsqueeze(-1)).reshape(shape).to(torch.bfloat16)


class FP4Linear(torch.nn.Module):
    def __init__(self, in_features, out_features, backend):
        super().__init__()
        assert in_features % 16 == 0 and out_features % 16 == 0
        self.in_features, self.out_features, self.backend = in_features, out_features, backend
        self.register_buffer('weight', torch.empty(out_features, in_features//2, dtype=torch.uint8, device='meta'))
        self.register_buffer('weight_scale', torch.empty(out_features, in_features//16, dtype=torch.float8_e4m3fn, device='meta'))
        self.register_buffer('weight_scale_2', torch.empty((), dtype=torch.float32, device='meta'))
        self.register_buffer('_reference_weight', None, persistent=False)
        self.register_buffer('_kernel_scales', None, persistent=False)

    def prepare_backend(self):
        if self.backend == 'reference':
            self._reference_weight = unpack_weight(self.weight, self.weight_scale, self.weight_scale_2)
        elif self.backend != 'portable':
            raise ValueError('backend must be portable or reference')

    def forward(self, x):
        if self.backend == 'reference':
            weight = self._reference_weight
        else:
            from fp4_kernels import dequantize_weight
            weight = dequantize_weight(self.weight, self.weight_scale, self.weight_scale_2)
        return torch.nn.functional.linear(x, weight)


class FP4DecisionModel(DecisionModel):
    def __init__(self, checkpoint, *, device='cuda:0', backend='portable', cpu_threads=8):
        torch.nn.Module.__init__(self)
        if backend not in ['portable', 'reference']:
            raise ValueError('backend must be portable or reference')
        torch.set_num_threads(cpu_threads)
        torch.backends.cuda.enable_cudnn_sdp(False)
        path = Path(checkpoint)
        quant = json.loads((path/'jev_quantization.json').read_text())
        if quant['format'] != 'jev_fp4_w4a16_v1':
            raise ValueError('Unsupported JEV quantization format')
        saved = json.loads((path/'decision_config.json').read_text())
        assert saved['format_version'] == 1
        self.device_name = device
        self.base_model, self.revision = saved['base_model'], saved['revision']
        self.temperature = saved['temperature']
        assert math.isfinite(self.temperature) and self.temperature > 0
        self.processor = AutoProcessor.from_pretrained(path, local_files_only=True)
        self.processor.tokenizer.padding_side = 'left'
        self.processor.image_processor.size = {'shortest_edge':65536, 'longest_edge':262144}
        self.codes, self.token_ids = answer_codes(self.processor.tokenizer)
        assert self.codes == saved['codes'] and self.token_ids == saved['token_ids']
        config = AutoConfig.from_pretrained(path, local_files_only=True)
        config._attn_implementation = 'sdpa'
        config.text_config._attn_implementation = 'sdpa'
        with init_empty_weights(include_buffers=False):
            self.backbone = Qwen3_5Model(config)
        for name, spec in quant['modules'].items():
            parent, _, child = name.rpartition('.')
            old = self.backbone.get_submodule(name)
            assert isinstance(old, torch.nn.Linear) and old.bias is None
            assert [old.out_features, old.in_features] == spec['shape']
            setattr(self.backbone.get_submodule(parent), child, FP4Linear(old.in_features, old.out_features, backend))
        index = json.loads((path/'model.safetensors.index.json').read_text())['weight_map']
        expected = set(self.backbone.state_dict())
        assert set(index) == expected, {'missing': sorted(expected-set(index)), 'extra': sorted(set(index)-expected)}
        seen = set()
        for filename in sorted(set(index.values())):
            state = load_file(str(path/filename))
            assert all(index[k] == filename for k in state) and not seen.intersection(state)
            seen.update(state)
            result = self.backbone.load_state_dict(state, strict=False, assign=True)
            assert not result.unexpected_keys
        assert seen == expected and not any(t.is_meta for t in self.backbone.state_dict().values())
        self.readout = torch.nn.Linear(config.text_config.hidden_size, MAX_OPTIONS, bias=False, dtype=torch.bfloat16)
        self.readout.load_state_dict(load_file(str(path/'readout.safetensors')))
        self.to(device)
        for module in self.backbone.modules():
            if isinstance(module, FP4Linear):
                module.prepare_backend()
        self.requires_grad_(False)
        self.eval()