Aharrypotter commented on code in PR #19639:
URL: https://github.com/apache/tvm/pull/19639#discussion_r3322181557
##########
python/tvm/relax/frontend/tflite/tflite_frontend.py:
##########
@@ -516,6 +531,228 @@ def convert_op_to_relax(self):
get_tensor_name(self.subgraph,
output_tensor.tensor_idx), ret[idx]
)
+ @staticmethod
+ def _decode_tflite_string(value):
+ """Decode a TFLite string field."""
+ if value is None:
+ return ""
+ if isinstance(value, (bytes, bytearray)):
+ return value.decode("utf-8")
+ return str(value)
+
+ def _get_var_handle_resource_key(self, op, fallback_tensor=None):
+ """Return a stable resource key for a VAR_HANDLE op."""
+ container = ""
+ shared_name = ""
+ if op.BuiltinOptions() is not None:
+ from tflite.VarHandleOptions import VarHandleOptions
+
+ opts = self._get_builtin_options(op, VarHandleOptions)
+ if hasattr(opts, "Container"):
+ container = self._decode_tflite_string(opts.Container())
+ if hasattr(opts, "SharedName"):
+ shared_name = self._decode_tflite_string(opts.SharedName())
+
+ if container or shared_name:
+ return (container, shared_name)
+ if fallback_tensor is not None:
+ return ("", get_tensor_name(self.subgraph,
fallback_tensor.tensor_idx))
+ raise tvm.error.OpNotImplemented("VAR_HANDLE requires
VarHandleOptions")
+
+ def _get_resource_key_for_handle(self, tensor, op_name):
+ tensor_name = get_tensor_name(self.subgraph, tensor.tensor_idx)
+ if tensor_name not in self.resource_handles:
+ raise tvm.error.OpNotImplemented(
+ f"{op_name} requires a VAR_HANDLE in the same TFLite subgraph"
+ )
+ return self.resource_handles[tensor_name]
+
+ def convert_var_handle(self, op):
+ """Convert a TFLite VAR_HANDLE into an importer-local resource
handle."""
+ input_tensors = self.get_input_tensors(op)
+ output_tensors = self.get_output_tensors(op)
+ if len(input_tensors) != 0 or len(output_tensors) != 1:
+ raise tvm.error.OpNotImplemented("VAR_HANDLE expects no inputs and
one output")
+
+ resource_key = self._get_var_handle_resource_key(op, output_tensors[0])
+ resource_tensor_name = get_tensor_name(self.subgraph,
output_tensors[0].tensor_idx)
+ self.resource_handles[resource_tensor_name] = resource_key
+ return None
+
+ def convert_assign_variable(self, op):
+ """Convert the CALL_ONCE initialization subset of ASSIGN_VARIABLE."""
+ if not self.conversion_state["in_call_once_init"]:
+ raise tvm.error.OpNotImplemented(
+ "ASSIGN_VARIABLE outside CALL_ONCE initialization is not
supported by the "
+ "Relax TFLite frontend yet because it requires mutable
resource state modeling."
+ )
+
+ input_tensors = self.get_input_tensors(op)
+ output_tensors = self.get_output_tensors(op)
+ if len(input_tensors) != 2 or len(output_tensors) != 0:
+ raise tvm.error.OpNotImplemented(
+ "ASSIGN_VARIABLE expects a resource handle and value input
with no outputs"
+ )
+
+ resource_key = self._get_resource_key_for_handle(input_tensors[0],
"ASSIGN_VARIABLE")
+ self.conversion_state["resource_values"][resource_key] =
self.get_tensor_expr(
+ input_tensors[1]
+ )
+ return None
+
+ def convert_read_variable(self, op):
+ """Convert READ_VARIABLE for resources initialized by CALL_ONCE."""
+ input_tensors = self.get_input_tensors(op)
+ output_tensors = self.get_output_tensors(op)
+ if len(input_tensors) != 1 or len(output_tensors) != 1:
+ raise tvm.error.OpNotImplemented("READ_VARIABLE expects one input
and one output")
+
+ resource_key = self._get_resource_key_for_handle(input_tensors[0],
"READ_VARIABLE")
+ resource_values = self.conversion_state["resource_values"]
+ if resource_key not in resource_values:
+ raise tvm.error.OpNotImplemented(
+ "READ_VARIABLE requires a resource initialized by a supported
CALL_ONCE subgraph"
+ )
+ return resource_values[resource_key]
+
+ def _get_hashtable_key(self, op, fallback_tensor=None):
+ """Return a stable key for a TFLite HASHTABLE resource."""
+ table_id = None
+ if op.BuiltinOptions() is not None:
+ from tflite.HashtableOptions import HashtableOptions
+
+ opts = self._get_builtin_options(op, HashtableOptions)
+ table_id = int(opts.TableId())
+
+ if table_id is not None:
+ return table_id
+ if fallback_tensor is not None:
+ return get_tensor_name(self.subgraph, fallback_tensor.tensor_idx)
+ raise tvm.error.OpNotImplemented("HASHTABLE requires HashtableOptions")
+
+ def _get_hashtable_key_for_handle(self, tensor, op_name):
+ tensor_name = get_tensor_name(self.subgraph, tensor.tensor_idx)
+ if tensor_name not in self.hashtable_handles:
+ raise tvm.error.OpNotImplemented(
+ f"{op_name} requires a HASHTABLE in the same TFLite subgraph"
+ )
+ return self.hashtable_handles[tensor_name]
+
+ def convert_hashtable(self, op):
+ """Convert a TFLite HASHTABLE into an importer-local table handle."""
+ input_tensors = self.get_input_tensors(op)
+ output_tensors = self.get_output_tensors(op)
+ if len(input_tensors) != 0 or len(output_tensors) != 1:
+ raise tvm.error.OpNotImplemented("HASHTABLE expects no inputs and
one output")
+
+ table_key = self._get_hashtable_key(op, output_tensors[0])
+ table_tensor_name = get_tensor_name(self.subgraph,
output_tensors[0].tensor_idx)
+ self.hashtable_handles[table_tensor_name] = table_key
+ return None
+
+ def convert_hashtable_import(self, op):
+ """Convert the CALL_ONCE initialization subset of HASHTABLE_IMPORT."""
+ if not self.conversion_state["in_call_once_init"]:
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_IMPORT outside CALL_ONCE initialization is not
supported by the "
+ "Relax TFLite frontend yet because it requires mutable
resource state modeling."
+ )
+
+ input_tensors = self.get_input_tensors(op)
+ output_tensors = self.get_output_tensors(op)
+ if len(input_tensors) != 3 or len(output_tensors) != 0:
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_IMPORT expects table, keys, and values inputs with
no outputs"
+ )
+
+ table_key = self._get_hashtable_key_for_handle(input_tensors[0],
"HASHTABLE_IMPORT")
+ keys = self.get_tensor_value(input_tensors[1])
+ values = self.get_tensor_value(input_tensors[2])
+ if keys.ndim != 1 or values.ndim < 1:
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_IMPORT requires one-dimensional keys and at least
one-dimensional values"
+ )
+ if keys.shape[0] != values.shape[0]:
+ raise tvm.error.OpNotImplemented("HASHTABLE_IMPORT keys and values
size mismatch")
+
+ self.conversion_state["hashtable_values"][table_key] = {
+ "keys": self.get_tensor_expr(input_tensors[1]),
+ "values": self.get_tensor_expr(input_tensors[2]),
+ "size": int(keys.shape[0]),
+ "value_shape": tuple(int(dim) for dim in values.shape[1:]),
+ }
+ return None
+
+ def convert_hashtable_find(self, op):
+ """Convert HASHTABLE_FIND for static tables initialized by
CALL_ONCE."""
+ input_tensors = self.get_input_tensors(op)
+ output_tensors = self.get_output_tensors(op)
+ if len(input_tensors) != 3 or len(output_tensors) != 1:
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_FIND expects table, keys, and default value inputs
with one output"
+ )
+
+ table_key = self._get_hashtable_key_for_handle(input_tensors[0],
"HASHTABLE_FIND")
+ hashtable_values = self.conversion_state["hashtable_values"]
+ if table_key not in hashtable_values:
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_FIND requires a table initialized by a supported
CALL_ONCE subgraph"
+ )
+
+ table = hashtable_values[table_key]
+ output_shape = (
+ tuple(output_tensors[0].tensor.ShapeAsNumpy())
+ if output_tensors[0].tensor.ShapeLength() > 0
+ else ()
+ )
+ value_shape = table["value_shape"]
+ if value_shape and (
+ len(output_shape) < len(value_shape)
+ or tuple(output_shape[-len(value_shape) :]) != value_shape
+ ):
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_FIND output shape must append the imported value
shape"
+ )
+ query_keys = self.get_tensor_expr(input_tensors[1])
+ table_keys = table["keys"]
+ table_values = table["values"]
+ default_value = self.get_tensor_expr(input_tensors[2])
+
+ matches = relax.op.equal(relax.op.expand_dims(query_keys, axis=-1),
table_keys)
Review Comment:
fixed
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]