File size: 3,786 Bytes
29f25be
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Manual reduced-output ANE experiment; not the deployment interface.

Usage: python scripts/deployment/ane_output_probe.py NEW_OUTPUT.mlpackage
Conversion, timing and accuracy checks start only when this command is run.
"""
import argparse
import sys
import time
from pathlib import Path

ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(ROOT / "src"))


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("output", type=Path, help="New scratch .mlpackage path")
    args = parser.parse_args()
    out = args.output
    if out.exists():
        parser.error("Output exists; choose a new experiment path.")
    if out.suffix != ".mlpackage":
        parser.error("Output must end with .mlpackage.")
    if sys.platform != "darwin":
        parser.error("Run this experiment manually on Mac.")
    import numpy as np, torch, coremltools as ct
    from torch import nn
    from vimeml.deployment.bundle import BundleLM
    from vimeml.deployment.graph import ConversionGraph
    lm = BundleLM(ROOT/"artifacts/deployment/tiny-ja-v1-inference-v1")
    K = 16
    class Head(nn.Module):
        def __init__(s, model): super().__init__(); s.g = ConversionGraph(model)
        def forward(s, input_ids, target_ids):
            logp = torch.log_softmax(s.g(input_ids), dim=-1)
            target = torch.gather(logp, 2, target_ids.long().unsqueeze(-1)).squeeze(-1)
            top, idx = torch.topk(logp, K, dim=-1)
            return target, top, idx.int()
    h = Head(lm.model).eval()
    ex = (torch.zeros(1,16,dtype=torch.int32), torch.zeros(1,16,dtype=torch.int32))
    tr = torch.jit.trace(h, ex)
    shapes = ct.EnumeratedShapes(shapes=[(1,L) for L in (16,32,64,128)], default=(1,16))
    m = ct.convert(tr, convert_to="mlprogram", minimum_deployment_target=ct.target.iOS18, compute_precision=ct.precision.FLOAT16,
        inputs=[ct.TensorType(name="input_ids", shape=shapes, dtype=np.int32), ct.TensorType(name="target_ids", shape=shapes, dtype=np.int32)],
        outputs=[ct.TensorType(name="target_logprobs", dtype=np.float32), ct.TensorType(name="top_logprobs", dtype=np.float32), ct.TensorType(name="top_ids", dtype=np.int32)])
    m.save(str(out))
    for units in ("CPU_ONLY","CPU_AND_NE"):
        mm = ct.models.MLModel(str(out), compute_units=getattr(ct.ComputeUnit, units)); row=[units]
        for L in (16,64,128):
            x={"input_ids":np.full((1,L),5,np.int32),"target_ids":np.full((1,L),6,np.int32)}; mm.predict(x); ts=[]
            for _ in range(20): t=time.perf_counter(); mm.predict(x); ts.append((time.perf_counter()-t)*1000)
            row.append(f"T{L}={np.median(ts):.2f}ms")
        print(" ".join(row))
    # accuracy vs FP32 on AJIMEE text
    import json
    items=json.load(open(ROOT/"artifacts/benchmarks/ajimee-jwtd-v2-v1/evaluation_items.json"))
    mm = ct.models.MLModel(str(out), compute_units=ct.ComputeUnit.CPU_AND_NE)
    diffs=[]; topagree=[]
    for it in items[:100]:
        ids=[2]+lm.processor.encode((it.get("context_text") or "")+it["expected_output"][0])[:127]
        n=len(ids)-1; L=next(b for b in (16,32,64,128) if b>=n)
        inp=np.zeros((1,L),np.int32); inp[0,:n]=ids[:-1]; tgt=np.zeros((1,L),np.int32); tgt[0,:n]=ids[1:]
        o=mm.predict({"input_ids":inp,"target_ids":tgt})
        with torch.no_grad(): ref=lm.model(torch.tensor([ids[:-1]])).float().log_softmax(-1)[0]
        r=ref.gather(1,torch.tensor(ids[1:])[:,None]).squeeze(1).numpy()
        diffs.append(abs(o["target_logprobs"][0,:n]-r).max())
        topagree.append(np.mean(o["top_ids"][0,:n,0]==ref.argmax(-1).numpy()))
    print(f"ANE target logprob max|diff| mean={np.mean(diffs):.4f} max={np.max(diffs):.4f}; top1 agree={np.mean(topagree):.4f}")


if __name__ == "__main__":
    main()