File size: 2,464 Bytes
a9f8c07
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Load a sapiens2-int4-safetensors artifact back into a PyTorch state dict."""

from __future__ import annotations

import argparse
import json
from pathlib import Path

import torch
from safetensors import safe_open
from safetensors.torch import save_file


def unpack_int4(packed: torch.Tensor, elements: int) -> torch.Tensor:
    packed = packed.cpu().to(torch.uint8)
    lo = (packed & 0x0F).to(torch.int16)
    hi = ((packed >> 4) & 0x0F).to(torch.int16)
    vals = torch.empty(packed.numel() * 2, dtype=torch.int16)
    vals[0::2] = lo
    vals[1::2] = hi
    vals = vals[:elements]
    vals = torch.where(vals >= 8, vals - 16, vals)
    return vals.to(torch.float32)


def dequantize_tensor(packed: torch.Tensor, scales: torch.Tensor, shape: list[int], group_size: int, dtype: str) -> torch.Tensor:
    elements = 1
    for dim in shape:
        elements *= dim
    pad = (-elements) % group_size
    q = unpack_int4(packed, elements + pad).view(-1, group_size)
    out = (q * scales.to(torch.float32)[:, None]).flatten()[:elements].view(*shape)
    target_dtype = getattr(torch, dtype, torch.float16)
    return out.to(target_dtype if target_dtype.is_floating_point else torch.float32)


def load_state_dict(path: str | Path, device: str = "cpu") -> dict[str, torch.Tensor]:
    with safe_open(str(path), framework="pt", device="cpu") as f:
        metadata = f.metadata() or {}
        manifest = json.loads(metadata["manifest_json"])
        state = {}
        for name, info in manifest.items():
            if info.get("quantized"):
                state[name] = dequantize_tensor(
                    f.get_tensor(info["qweight"]),
                    f.get_tensor(info["scales"]),
                    info["shape"],
                    int(info["group_size"]),
                    info.get("dtype", "float16"),
                ).to(device)
            else:
                state[name] = f.get_tensor(name).to(device)
    return state


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("artifact")
    parser.add_argument("--save-dequantized", help="Optional safetensors path for the dequantized state dict")
    args = parser.parse_args()

    state = load_state_dict(args.artifact)
    print(f"loaded {len(state)} tensors")
    if args.save_dequantized:
        save_file(state, args.save_dequantized)
        print(f"saved {args.save_dequantized}")


if __name__ == "__main__":
    main()