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:
   ![medium](https://www.gstatic.com/codereviewagent/medium-priority.svg)
   
   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:
   ![medium](https://www.gstatic.com/codereviewagent/medium-priority.svg)
   
   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:
   ![medium](https://www.gstatic.com/codereviewagent/medium-priority.svg)
   
   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]

Reply via email to