Skip to content

Commit b07cd11

Browse files
committed
WIP
1 parent dafcf53 commit b07cd11

1 file changed

Lines changed: 16 additions & 31 deletions

File tree

qa/common/gen_qa_models.py

Lines changed: 16 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,8 @@
4747
np_dtype_string = np.dtype(object)
4848
from typing import List, Tuple
4949

50+
51+
_color_blue = "\033[94m"
5052
_color_green = "\033[32m"
5153
_color_magenta = "\033[35m"
5254
_color_red = "\033[31m"
@@ -1655,7 +1657,7 @@ def create_libtorch_pt2_modelconfig(
16551657
dims: [ {tu.shape_to_dims_str(output1_shape)} ]
16561658
}}
16571659
]
1658-
instance_group [{{ kind: "{"KIND_GPU" if torch.cuda.is_available() else "KIND_CPU"}" }}]
1660+
instance_group [{{ kind: {"KIND_GPU" if torch.cuda.is_available() else "KIND_CPU"} }}]
16591661
"""
16601662

16611663
try:
@@ -1987,6 +1989,7 @@ def create_models(
19871989
output0_label_cnt,
19881990
version_policy=None,
19891991
):
1992+
print(f"{_color_blue}Creating models in {models_dir}{_color_reset}")
19901993
model_version = 1
19911994
if FLAGS.tensorrt:
19921995
print(f"{_color_magenta}TensorRT model generation requested{_color_reset}")
@@ -2266,18 +2269,6 @@ def create_models(
22662269
output1_dtype,
22672270
)
22682271

2269-
if FLAGS.torchvision_aoti:
2270-
print(f"{_color_magenta}TorchVision AOTI model generation requested{_color_reset}")
2271-
if create_torchvision_aoti_modelfile(
2272-
models_dir,
2273-
1,
2274-
model_version,
2275-
):
2276-
create_torchvision_aoti_modelconfig(
2277-
models_dir,
2278-
1,
2279-
)
2280-
22812272
if FLAGS.openvino:
22822273
print(f"{_color_magenta}OpenVINO model generation requested{_color_reset}")
22832274
# max-batch 8
@@ -2692,24 +2683,6 @@ def create_fixed_models(
26922683
FLAGS.models_dir, 0, 3, (16,), (16,), (16,), vt, vt, vt, swap=True
26932684
)
26942685

2695-
if FLAGS.torch_aoti:
2696-
for vt in [np.float32, np.int32, np.int16, np.int8]:
2697-
create_torch_aoti_modelfile(
2698-
FLAGS.models_dir, 8, 2, (16,), (16,), (16,), vt, vt, vt, swap=True
2699-
)
2700-
create_torch_aoti_modelfile(
2701-
FLAGS.models_dir, 8, 3, (16,), (16,), (16,), vt, vt, vt, swap=True
2702-
)
2703-
create_torch_aoti_modelfile(
2704-
FLAGS.models_dir, 0, 2, (16,), (16,), (16,), vt, vt, vt, swap=True
2705-
)
2706-
create_torch_aoti_modelfile(
2707-
FLAGS.models_dir, 0, 3, (16,), (16,), (16,), vt, vt, vt, swap=True
2708-
)
2709-
2710-
if FLAGS.torchvision_aoti:
2711-
create_torchvision_aoti_modelfile(FLAGS.models_dir, 1, 1)
2712-
27132686
if FLAGS.openvino:
27142687
for vt in [np.float16, np.float32, np.int8, np.int16, np.int32]:
27152688
create_openvino_modelfile(
@@ -2925,3 +2898,15 @@ def create_fixed_models(
29252898
# to fixed size model is not safe but doable
29262899
for model_shape in [(-1,), (-1, -1), (-1, -1, -1)]:
29272900
emu.create_nop_modelconfig(FLAGS.models_dir, model_shape, model_dtype)
2901+
2902+
if FLAGS.torchvision_aoti:
2903+
print(f"{_color_blue}TorchVision AOTI model generation requested{_color_reset}")
2904+
if create_torchvision_aoti_modelfile(
2905+
FLAGS.models_dir,
2906+
1,
2907+
1,
2908+
):
2909+
create_torchvision_aoti_modelconfig(
2910+
FLAGS.models_dir,
2911+
1,
2912+
)

0 commit comments

Comments
 (0)