gemini-code-assist[bot] commented on code in PR #19879:
URL: https://github.com/apache/tvm/pull/19879#discussion_r3464323707
##########
python/tvm/relax/frontend/tflite/tflite_frontend.py:
##########
@@ -711,6 +711,40 @@ def _has_tensor_buffer_data(tensor_wrapper):
and tensor_wrapper.buffer.DataLength() > 0
)
+ def _get_string_tensor_value(self, tensor_wrapper, op_name):
+ """Decode a constant TFLite string tensor buffer."""
+ if not self._is_tflite_string_type(tensor_wrapper.tensor.Type()):
+ raise tvm.error.OpNotImplemented(f"{op_name} requires a
TensorType.STRING tensor")
+ if not self._has_tensor_buffer_data(tensor_wrapper):
+ raise tvm.error.OpNotImplemented(f"{op_name} requires a constant
string tensor")
+
+ data = bytes(tensor_wrapper.buffer.DataAsNumpy())
+ if len(data) < 4:
+ raise tvm.error.OpNotImplemented(f"{op_name} has an invalid string
tensor buffer")
+
+ count = int(np.frombuffer(data, dtype="<i4", count=1)[0])
+ if count < 0:
+ raise tvm.error.OpNotImplemented(f"{op_name} has an invalid string
tensor count")
+
+ header_size = 4 * (count + 2)
+ if len(data) < header_size:
+ raise tvm.error.OpNotImplemented(f"{op_name} has an invalid string
tensor offsets")
+
+ offsets = np.frombuffer(data, dtype="<i4", count=count + 1,
offset=4).astype(np.int64)
+ if np.any(offsets < header_size) or np.any(offsets > len(data)):
+ raise tvm.error.OpNotImplemented(f"{op_name} has out-of-bounds
string tensor offsets")
+ if np.any(offsets[:-1] > offsets[1:]):
+ raise tvm.error.OpNotImplemented(f"{op_name} has non-monotonic
string tensor offsets")
+
+ values = [data[int(offsets[i]) : int(offsets[i + 1])].decode("utf-8")
for i in range(count)]
Review Comment:

If the constant string tensor contains invalid UTF-8 data, calling
`.decode("utf-8")` will raise a raw `UnicodeDecodeError`. It is more robust to
catch this exception and raise a clean `tvm.error.OpNotImplemented` compiler
error with a descriptive message.
```suggestion
try:
values = [
data[int(offsets[i]) : int(offsets[i + 1])].decode("utf-8")
for i in range(count)
]
except UnicodeDecodeError as e:
raise tvm.error.OpNotImplemented(
f"{op_name} has invalid UTF-8 string data: {e}"
) from e
```
##########
python/tvm/relax/frontend/tflite/tflite_frontend.py:
##########
@@ -759,21 +793,103 @@ def convert_hashtable_import(self, op):
):
raise tvm.error.OpNotImplemented("HASHTABLE_IMPORT requires
constant keys and values")
+ if self._is_tflite_string_type(table_info["key_dtype"]):
+ keys = self._get_string_tensor_value(key_tensor,
"HASHTABLE_IMPORT")
+ else:
+ keys = self.get_tensor_value(key_tensor)
+ if self._is_tflite_string_type(table_info["value_dtype"]):
+ values = self._get_string_tensor_value(value_tensor,
"HASHTABLE_IMPORT")
+ else:
+ values = self.get_tensor_value(value_tensor)
+
+ flat_keys = keys.reshape(-1).tolist()
+ if len(set(flat_keys)) != len(flat_keys):
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_IMPORT with duplicate keys is not supported"
+ )
+
hashtable_values = self.conversion_state["hashtable_values"]
table_key = table_info["table_key"]
if table_key not in hashtable_values:
hashtable_values[table_key] = {
"size": math.prod(key_shape) if key_shape else 1,
"key_dtype": table_info["key_dtype"],
"value_dtype": table_info["value_dtype"],
+ "keys": keys,
+ "values": values,
}
return None
def convert_hashtable_find(self, op):
- """Reject HASHTABLE_FIND until Relax can represent TFLite string
tensors."""
- raise tvm.error.OpNotImplemented(
- "HASHTABLE_FIND requires TensorType.STRING support in Relax TFLite
frontend"
+ """Convert the constant-foldable string-to-int64 HASHTABLE_FIND
subset."""
+ from tflite.TensorType import TensorType
+
+ 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, query, and default inputs with
one output"
+ )
+
+ table_tensor, query_tensor, default_tensor = input_tensors
+ output_tensor = output_tensors[0]
+ table_info = self._get_hashtable_info_for_handle(table_tensor,
"HASHTABLE_FIND")
+ table_key = table_info["table_key"]
+ 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_values = hashtable_values[table_key]
+
+ if (
+ query_tensor.tensor.Type() != table_values["key_dtype"]
+ or default_tensor.tensor.Type() != table_values["value_dtype"]
+ or output_tensor.tensor.Type() != table_values["value_dtype"]
+ ):
+ raise tvm.error.OpNotImplemented("HASHTABLE_FIND key/value dtypes
mismatch")
+
+ if not (
+ self._is_tflite_string_type(table_values["key_dtype"])
+ and table_values["value_dtype"] == TensorType.INT64
+ ):
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_FIND only supports constant string -> int64 tables"
+ )
+ if not self._has_tensor_buffer_data(query_tensor):
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_FIND with runtime string queries is not supported"
+ )
+ if not self._has_tensor_buffer_data(default_tensor):
+ raise tvm.error.OpNotImplemented("HASHTABLE_FIND requires constant
default values")
+
+ query_shape = self._get_tensor_shape_tuple(query_tensor)
+ output_shape = self._get_tensor_shape_tuple(output_tensor)
+ if output_shape != query_shape:
+ raise tvm.error.OpNotImplemented("HASHTABLE_FIND output shape must
match query shape")
+
+ query_values = self._get_string_tensor_value(query_tensor,
"HASHTABLE_FIND")
+ default_values = self.get_tensor_value(default_tensor)
+ default_shape = self._get_tensor_shape_tuple(default_tensor)
+ if default_shape == ():
+ result = np.full(output_shape, int(default_values.item()),
dtype=np.int64)
Review Comment:

In TFLite, default values are sometimes wrapped in a 1-element tensor (e.g.,
shape `(1,)`) instead of being a strict scalar (shape `()`). We can support
both cases seamlessly by checking if `default_values.size == 1`.
```suggestion
if default_shape == () or default_values.size == 1:
result = np.full(output_shape, int(default_values.item()),
dtype=np.int64)
```
##########
python/tvm/relax/frontend/tflite/tflite_frontend.py:
##########
@@ -759,21 +793,103 @@ def convert_hashtable_import(self, op):
):
raise tvm.error.OpNotImplemented("HASHTABLE_IMPORT requires
constant keys and values")
+ if self._is_tflite_string_type(table_info["key_dtype"]):
+ keys = self._get_string_tensor_value(key_tensor,
"HASHTABLE_IMPORT")
+ else:
+ keys = self.get_tensor_value(key_tensor)
+ if self._is_tflite_string_type(table_info["value_dtype"]):
+ values = self._get_string_tensor_value(value_tensor,
"HASHTABLE_IMPORT")
+ else:
+ values = self.get_tensor_value(value_tensor)
+
+ flat_keys = keys.reshape(-1).tolist()
+ if len(set(flat_keys)) != len(flat_keys):
+ raise tvm.error.OpNotImplemented(
+ "HASHTABLE_IMPORT with duplicate keys is not supported"
+ )
Review Comment:

Converting the numpy array to a flat list and then a set to check for
duplicate keys is inefficient and less idiomatic. We can use `np.unique`
directly on the numpy array to perform this check more efficiently.
```python
if np.unique(keys).size != keys.size:
raise tvm.error.OpNotImplemented(
"HASHTABLE_IMPORT with duplicate keys is not supported"
)
```
--
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]