File size: 4,028 Bytes
69cfd78
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
af24c1d
69cfd78
 
 
 
 
 
af24c1d
69cfd78
 
 
 
 
 
 
 
 
 
 
 
 
af24c1d
69cfd78
9558c4e
69cfd78
 
 
 
 
af24c1d
69cfd78
 
af24c1d
 
69cfd78
af24c1d
 
 
 
69cfd78
af24c1d
69cfd78
 
 
 
 
 
 
 
 
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
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
from __future__ import annotations

import argparse
from pathlib import Path

import onnx
from onnx import TensorProto, helper

from custom_ops import (
    mesh_topology_extract,
    model,
    scalar,
    sparse_conv3d,
    sparse_downsample,
    sparse_subdivide,
    sparse_upsample,
    sparse_window_attention,
    vi,
)


def save(m, path: Path):
    path.parent.mkdir(parents=True, exist_ok=True)
    onnx.checker.check_model(m, full_check=False)
    onnx.save(m, path)
    print(path)


def build_all(out: Path):
    g = helper.make_graph(
        [sparse_subdivide(["feats", "coords"], ["out_feats", "out_coords"])],
        "SparseSubdivide",
        [vi("feats", TensorProto.FLOAT16, ["N", "C"]), vi("coords", TensorProto.INT32, ["N", 4])],
        [vi("out_feats", TensorProto.FLOAT16, ["N8", "C"]), vi("out_coords", TensorProto.INT32, ["N8", 4])],
    )
    save(model(g), out / "SparseSubdivide.onnx")

    g = helper.make_graph(
        [sparse_downsample(["feats", "coords"], ["out_feats", "out_coords", "inverse", "count"], factor_x=2, factor_y=2, factor_z=2)],
        "SparseDownsample",
        [vi("feats", TensorProto.FLOAT16, ["N", "C"]), vi("coords", TensorProto.INT32, ["N", 4])],
        [vi("out_feats", TensorProto.FLOAT16, ["M", "C"]), vi("out_coords", TensorProto.INT32, ["M", 4]), vi("inverse", TensorProto.INT32, ["N"]), scalar("count")],
    )
    save(model(g), out / "SparseDownsample.onnx")

    g = helper.make_graph(
        [sparse_upsample(["feats", "target_coords", "inverse"], ["out_feats", "out_coords"])],
        "SparseUpsample",
        [vi("feats", TensorProto.FLOAT16, ["M", "C"]), vi("target_coords", TensorProto.INT32, ["N", 4]), vi("inverse", TensorProto.INT32, ["N"])],
        [vi("out_feats", TensorProto.FLOAT16, ["N", "C"]), vi("out_coords", TensorProto.INT32, ["N", 4])],
    )
    save(model(g), out / "SparseUpsample.onnx")

    g = helper.make_graph(
        [sparse_window_attention(["qkv", "coords"], ["out"], window_size=8, shift_x=0, shift_y=0, shift_z=0)],
        "SparseWindowAttention",
        [vi("qkv", TensorProto.FLOAT16, ["N", 3, "H", "D"]), vi("coords", TensorProto.INT32, ["N", 4])],
        [vi("out", TensorProto.FLOAT16, ["N", "H", "D"])],
    )
    save(model(g), out / "SparseWindowAttention.onnx")

    g = helper.make_graph(
        [sparse_conv3d(["feats", "coords", "weight", "bias"], ["out_feats", "out_coords", "count"], out_channels=16, kernel_size=3, stride=1, dilation=1, padding=0, subm=True, spatial_x=64, spatial_y=64, spatial_z=64, batch_size=1)],
        "SparseConv3D",
        [vi("feats", TensorProto.FLOAT16, ["N", 8]), vi("coords", TensorProto.INT32, ["N", 4]), vi("weight", TensorProto.FLOAT16, [16, 3, 3, 3, 8]), vi("bias", TensorProto.FLOAT16, [16])],
        [vi("out_feats", TensorProto.FLOAT16, ["M", 16]), vi("out_coords", TensorProto.INT32, ["M", 4]), scalar("count")],
    )
    save(model(g), out / "SparseConv3D.onnx")

    g = helper.make_graph(
        [mesh_topology_extract(["verts_grid", "sdf", "cube_idx", "beta", "alpha", "gamma", "colors_grid"], ["vertices", "faces", "colors", "vertex_count", "face_count"], resolution=64, no_sigmoid=True)],
        "MeshTopologyExtract",
        [
            vi("verts_grid", TensorProto.FLOAT, ["Vg", 3]),
            vi("sdf", TensorProto.FLOAT, ["Vg"]),
            vi("cube_idx", TensorProto.INT64, ["Nc", 8]),
            vi("beta", TensorProto.FLOAT, ["Nc", 12]),
            vi("alpha", TensorProto.FLOAT, ["Nc", 8]),
            vi("gamma", TensorProto.FLOAT, ["Nc"]),
            vi("colors_grid", TensorProto.FLOAT, ["Vg", 6]),
        ],
        [vi("vertices", TensorProto.FLOAT, ["Nv", 3]), vi("faces", TensorProto.INT32, ["Nf", 3]), vi("colors", TensorProto.FLOAT, ["Nv", 6]), scalar("vertex_count"), scalar("face_count")],
    )
    save(model(g), out / "MeshTopologyExtract.onnx")


if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("--out", default="custom-onnx")
    args = ap.parse_args()
    build_all(Path(args.out))