"""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()