[Frontend] [Tensorflow2] Added test infrastructure for TF2 frozen models (#8074)
* added test infrastructure for frozen TF2 models * linting with black * removing some comments * change in comment in sequential test * addressed the comments * refactored to place vmobj_to_list in a common file * Added helper function in python/tvm/relay/testing/tf.py Co-authored-by: David Huang <davhuan@amazon.com> Co-authored-by: Rohan Mukherjee <mukrohan@amazon.com> Co-authored-by: Xiao <weix@amazon.com> * Refactor tf according to CI error Co-authored-by: David Huang <davhuan@amazon.com> Co-authored-by: Rohan Mukherjee <mukrohan@amazon.com> Co-authored-by: Xiao <weix@amazon.com> * Added docstring Co-authored-by: David Huang <davhuan@amazon.com> Co-authored-by: Rohan Mukherjee <mukrohan@amazon.com> Co-authored-by: Xiao <weix@amazon.com> * removing print Co-authored-by: David Huang <davhuan@amazon.com> Co-authored-by: Xiao <weix@amazon.com>
This commit is contained in:
@@ -28,6 +28,8 @@ import numpy as np
|
||||
# Tensorflow imports
|
||||
import tensorflow as tf
|
||||
from tensorflow.core.framework import graph_pb2
|
||||
|
||||
import tvm
|
||||
from tvm.contrib.download import download_testdata
|
||||
|
||||
try:
|
||||
@@ -73,6 +75,46 @@ def convert_to_list(x):
|
||||
return x
|
||||
|
||||
|
||||
def vmobj_to_list(o):
|
||||
"""Converts TVM objects returned by VM execution to Python List.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
o : Obj
|
||||
VM Object as output from VM runtime executor.
|
||||
|
||||
Returns
|
||||
-------
|
||||
result : list
|
||||
Numpy objects as list with equivalent values to the input object.
|
||||
|
||||
"""
|
||||
|
||||
if isinstance(o, tvm.nd.NDArray):
|
||||
result = [o.asnumpy()]
|
||||
elif isinstance(o, tvm.runtime.container.ADT):
|
||||
result = []
|
||||
for f in o:
|
||||
result.extend(vmobj_to_list(f))
|
||||
elif isinstance(o, tvm.relay.backend.interpreter.ConstructorValue):
|
||||
if o.constructor.name_hint == "Cons":
|
||||
tl = vmobj_to_list(o.fields[1])
|
||||
hd = vmobj_to_list(o.fields[0])
|
||||
hd.extend(tl)
|
||||
result = hd
|
||||
elif o.constructor.name_hint == "Nil":
|
||||
result = []
|
||||
elif "tensor_nil" in o.constructor.name_hint:
|
||||
result = [0]
|
||||
elif "tensor" in o.constructor.name_hint:
|
||||
result = [o.fields[0].asnumpy()]
|
||||
else:
|
||||
raise RuntimeError("Unknown object type: %s" % o.constructor.name_hint)
|
||||
else:
|
||||
raise RuntimeError("Unknown object type: %s" % type(o))
|
||||
return result
|
||||
|
||||
|
||||
def AddShapesToGraphDef(session, out_node):
|
||||
"""Add shapes attribute to nodes of the graph.
|
||||
Input graph here is the default graph in context.
|
||||
|
||||
Reference in New Issue
Block a user