File size: 2,052 Bytes
ff5f59d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Task 2: convert QDec to Core ML. usage: convert.py B T [variant...]; variants fp16 | pal8 | lin8.
Writes models/dec_b{B}_t{T}_{variant}.mlpackage and appends one line per variant to logs/convert.log."""
import sys, time, traceback, torch, numpy as np, coremltools as ct
import coremltools.optimize.coreml as cto
from pathlib import Path
from qdec import QDec, ane_kwargs

B, T = int(sys.argv[1]), int(sys.argv[2])
variants = sys.argv[3:] or ["fp16"]
Path("models").mkdir(exist_ok=True); Path("logs").mkdir(exist_ok=True)
log = open("logs/convert.log", "a")
def L(msg):
    print(msg); log.write(msg + "\n"); log.flush()

import os
CH = int(os.environ.get('CHUNK', '0')) or None
m = QDec(T, B=B, chunk=CH, **ane_kwargs()).eval()
SUF = f'c{CH}' if CH else ''
ex = torch.zeros((B, T), dtype=torch.int32)
t0 = time.time()
with torch.no_grad():
    tr = torch.jit.trace(m, ex.long())
base = None
for v in variants:
    name = f"models/dec_b{B}_t{T}_{v}{SUF}.mlpackage"
    try:
        t0 = time.time()
        if base is None:
            base = ct.convert(tr, inputs=[ct.TensorType("ids", shape=(B, T), dtype=np.int32)],
                              outputs=[ct.TensorType("lp", dtype=np.float32)],
                              compute_precision=ct.precision.FLOAT16,
                              minimum_deployment_target=ct.target.macOS15, convert_to="mlprogram")
        mm = base
        if v == "pal8":
            mm = cto.palettize_weights(base, cto.OptimizationConfig(global_config=cto.OpPalettizerConfig(nbits=8, mode="kmeans")))
        elif v == "lin8":
            mm = cto.linear_quantize_weights(base, cto.OptimizationConfig(global_config=cto.OpLinearQuantizerConfig(mode="linear_symmetric")))
        mm.save(name)
        sz = sum(f.stat().st_size for f in Path(name).rglob("*") if f.is_file()) / 1e6
        L(f"OK   b{B} t{T} {v}{SUF}: {time.time()-t0:.1f}s  {sz:.1f} MB  -> {name}")
    except Exception as e:
        L(f"FAIL b{B} t{T} {v}{SUF}: {type(e).__name__}: {str(e).splitlines()[0][:300]}")
        traceback.print_exc()