"""Exact rewrites that make a graph one QNN's HTP can hold (docs/dev/inference.md §1.5). Each rewrite spells the same arithmetic in operators the Hexagon runs, and `rewrite` checks the result against the input graph on ONNX Runtime's CPU before returning it — a rewrite that changes an output by more than float rounding is refused, not shipped. - **6-D Bayer pack → SpaceToDepth(2).** The denoiser packs the mosaic with `Reshape(1,C,H/2,2,W/2,2) → Transpose(0,1,3,5,2,4) → Reshape(1,4C,…)`. QNN's tensors stop at rank 5 (error 6007 at compose); for C = 1 that sequence *is* SpaceToDepth. - **Unfold → SpaceToDepth(8).** XFeat spells its 8×8 unfold as 224 Slices, 225 Transposes and two 6-D Concats; it is SpaceToDepth(8) of the normalised image, 736 nodes down to 60. - **Computed reshape targets → constants.** `Shape → Slice → Concat` feeding a Reshape, where shape inference proves the answer. - **Bilinear Resize → two MatMuls.** The HTP refuses `ResizeBilinear` at the sizes XFeat uses (3110); a half-pixel bilinear resize between fixed sizes is `X · Rxᵀ` then `Ry ·`, with ONNX's own edge-clamped weights. - **InstanceNormalization → primitives** is not here: it is exact, but the two full-image reductions it becomes cost the HTP more than it saves. XFeat keeps its normalisation in int8, which the HTP takes. """ import numpy as np import onnx import onnxruntime as ort from onnx import helper, numpy_helper, shape_inference TOLERANCE = 1e-4 # relative to each output's largest magnitude def _prune(g): """Drop nodes nobody reads and initialisers nobody uses.""" while True: used = {i for n in g.node for i in n.input} | {o.name for o in g.output} keep = [n for n in g.node if any(o in used for o in n.output)] if len(keep) == len(g.node): break del g.node[:] g.node.extend(keep) used = {i for n in g.node for i in n.input} keep = [i for i in g.initializer if i.name in used] del g.initializer[:] g.initializer.extend(keep) def bayer_pack(m): g = m.graph prod = {o: n for n in g.node for o in n.output} users = {} for n in g.node: for i in n.input: users.setdefault(i, []).append(n) swaps = {} for n in g.node: if n.op_type != "Transpose": continue perm = next(a.ints for a in n.attribute if a.name == "perm") r1, u = prod.get(n.input[0]), users.get(n.output[0], []) if list(perm) != [0, 1, 3, 5, 2, 4] or not r1 or r1.op_type != "Reshape": continue if len(u) != 1 or u[0].op_type != "Reshape": continue s2d = helper.make_node("SpaceToDepth", [r1.input[0]], [u[0].output[0]], name=n.name + "_s2d", blocksize=2) swaps[id(r1)] = s2d swaps[id(n)] = swaps[id(u[0])] = None nodes = [swaps.get(id(n), n) for n in g.node if swaps.get(id(n), n) is not None] del g.node[:] g.node.extend(nodes) return m def fold_reshapes(m): m = shape_inference.infer_shapes(m) g = m.graph vi = {v.name: v for v in list(g.value_info) + list(g.output) + list(g.input)} inits = {i.name for i in g.initializer} dims = lambda t: [d.dim_value for d in vi[t].type.tensor_type.shape.dim] if t in vi else [] for n in g.node: if n.op_type != "Reshape" or n.input[1] in inits: continue shape = dims(n.output[0]) if not shape or 0 in shape: src = dims(n.input[0]) if len(src) != 5 or 0 in src: # (1,C,k,H,W) -> (1,C·k,H,W), the denoiser's tile continue shape = [src[0], src[1] * src[2], src[3], src[4]] name = n.output[0] + "_shape" g.initializer.append(numpy_helper.from_array(np.array(shape, np.int64), name)) n.input[1] = name _prune(g) del g.value_info[:] return m def unfold(m, block=8): """Replace the Slice/Transpose/Concat region ending in the Reshape that produces the (1, block², H/block, W/block) tensor with SpaceToDepth.""" g = m.graph prod = {o: n for n in g.node for o in n.output} region_ops = ("Slice", "Transpose", "Concat", "Unsqueeze", "Reshape") inits = {i.name for i in g.initializer} m_inf = shape_inference.infer_shapes(m) vi = {v.name: [d.dim_value for d in v.type.tensor_type.shape.dim] for v in m_inf.graph.value_info} for end in g.node: out = vi.get(end.output[0], []) if end.op_type != "Reshape" or len(out) != 4 or out[1] != block * block: continue seen, stack, leaves = set(), [end.input[0]], set() while stack: t = stack.pop() if t in seen or t in inits: continue seen.add(t) n = prod.get(t) if n is not None and n.op_type in region_ops: stack += list(n.input) else: leaves.add(t) if len(leaves) != 1 or len(seen) < 100: # the hand-written unfold, not an ordinary reshape continue region = {id(end)} | {id(prod[t]) for t in seen if t in prod and prod[t].op_type in region_ops} nodes = [] for n in g.node: if id(n) in region: if n is end: nodes.append(helper.make_node("SpaceToDepth", [next(iter(leaves))], [end.output[0]], name=end.name + "_s2d", blocksize=block)) continue nodes.append(n) del g.node[:] g.node.extend(nodes) _prune(g) del g.value_info[:] return m return m def _bilinear(n_in, n_out): r = np.zeros((n_out, n_in), np.float32) for o in range(n_out): x = min(max((o + 0.5) * n_in / n_out - 0.5, 0), n_in - 1) i0 = int(np.floor(x)) f = x - i0 r[o, i0] += 1 - f r[o, min(i0 + 1, n_in - 1)] += f return r def resize_matmul(m): """Every linear, half-pixel Resize between fixed NCHW sizes.""" m = shape_inference.infer_shapes(m) g = m.graph vi = {v.name: [d.dim_value for d in v.type.tensor_type.shape.dim] for v in list(g.value_info) + list(g.output)} nodes = [] for n in g.node: a = {x.name: helper.get_attribute_value(x) for x in n.attribute} ok = (n.op_type == "Resize" and a.get("mode") == b"linear" and a.get("coordinate_transformation_mode", b"half_pixel") == b"half_pixel" and len(vi.get(n.input[0], [])) == 4 and len(vi.get(n.output[0], [])) == 4) if not ok: nodes.append(n) continue (_, _, h, w), (_, _, h2, w2) = vi[n.input[0]], vi[n.output[0]] t = n.name rx = numpy_helper.from_array(_bilinear(w, w2).T.copy(), t + "_rxT") ry = numpy_helper.from_array(_bilinear(h, h2).T.copy(), t + "_ryT") g.initializer.extend([rx, ry]) nodes += [ helper.make_node("MatMul", [n.input[0], rx.name], [t + "_w"], name=t + "_mw"), helper.make_node("Transpose", [t + "_w"], [t + "_t"], name=t + "_t1", perm=[0, 1, 3, 2]), helper.make_node("MatMul", [t + "_t", ry.name], [t + "_h"], name=t + "_mh"), helper.make_node("Transpose", [t + "_h"], [n.output[0]], name=t + "_t2", perm=[0, 1, 3, 2]), ] del g.node[:] g.node.extend(nodes) _prune(g) del g.value_info[:] return m REWRITES = {"bayer": [bayer_pack, fold_reshapes], "unfold": [unfold], "resize": [resize_matmul]} def rewrite(path, names): """The graph at `path` with the named rewrites applied, checked exact.""" m = onnx.load(path) before = len(m.graph.node) for name in names: for step in REWRITES[name]: m = step(m) onnx.checker.check_model(m) a = ort.InferenceSession(path, providers=["CPUExecutionProvider"]) b = ort.InferenceSession(m.SerializeToString(), providers=["CPUExecutionProvider"]) rng = np.random.default_rng(3) feed = {i.name: rng.random([d if isinstance(d, int) else 1 for d in i.shape], dtype=np.float32) for i in a.get_inputs()} for o, x, y in zip(a.get_outputs(), a.run(None, feed), b.run(None, feed)): err = float(np.abs(x - y).max()) / max(float(np.abs(x).max()), 1e-6) if err > TOLERANCE: raise SystemExit(f"{path}: rewrite {names} moved output {o.name} by {err:.2e} (relative)") print(f" rewrites {'+'.join(names)}: {before} -> {len(m.graph.node)} nodes, exact") return m