Aharrypotter commented on code in PR #19867:
URL: https://github.com/apache/tvm/pull/19867#discussion_r3458545074


##########
python/tvm/relax/frontend/tflite/tflite_frontend.py:
##########
@@ -1541,38 +1547,86 @@ def convert_range(self, op):
 
         start, limit, delta = input_tensors[0], input_tensors[1], 
input_tensors[2]
 
-        def get_scalar_value(tensor):
+        # out type inference
+        if delta.tensor.Type() == TensorType.FLOAT32:
+            out_type = self.get_tensor_type_str(delta.tensor.Type())
+        else:
+            out_type = self.get_tensor_type_str(start.tensor.Type())
+
+        def is_dynamic(tensor):
+            return self.has_expr(tensor.tensor_idx) and not isinstance(
+                self.get_expr(tensor.tensor_idx), relax.Constant
+            )
+
+        def static_scalar(tensor):
             if self.has_expr(tensor.tensor_idx):
-                expr = self.get_expr(tensor.tensor_idx)
-                if isinstance(expr, relax.Constant):
-                    value = expr.data.numpy()
-                else:
-                    # relax.op.arange currently expects scalar-like values 
here.
-                    # Keep dynamic scalar RANGE explicit until frontend 
support is added.
-                    raise tvm.error.OpNotImplemented(
-                        "TFLite RANGE with dynamic scalar inputs is not 
supported in"
-                        "Relax frontend yet."
-                    )
+                value = self.get_expr(tensor.tensor_idx).data.numpy()
             else:
                 value = self.get_tensor_value(tensor)
-
             # TFLite RANGE operands are scalar tensors in the flatbuffer.
             assert value.size == 1, "RANGE scalar input must have exactly one 
element"
             return value.item()
 
-        start_value = get_scalar_value(start)
-        limit_value = get_scalar_value(limit)
-        delta_value = get_scalar_value(delta)
+        if not (is_dynamic(start) or is_dynamic(limit) or is_dynamic(delta)):
+            return relax.op.arange(
+                static_scalar(start), static_scalar(limit), 
static_scalar(delta), out_type
+            )
+
+        return self._convert_dynamic_range(start, limit, delta, out_type)
 
-        # out type inference
-        if delta.tensor.Type() == TensorType.FLOAT32:
-            out_type = self.get_tensor_type_str(delta.tensor.Type())
-        else:
-            out_type = self.get_tensor_type_str(start.tensor.Type())
+    def _scalar_tensor_to_dim(self, expr, name):
+        """Lift a runtime scalar Relax expr to a symbolic ``tirx.Var`` 
dimension.
 
-        out = relax.op.arange(start_value, limit_value, delta_value, out_type)
+        Mirrors the ``tensor_to_shape`` + ``match_cast`` bridge used by
+        ``_get_shape_expr_from_tensor`` so a data-dependent scalar can be used 
as
+        a ``PrimExpr`` (e.g. an output length). The scalar is cast to int64 
first.
+        """
+        expr = self.bb.normalize(relax.op.astype(expr, "int64"))
+        expr = self.bb.normalize(relax.op.reshape(expr, (1,)))
+        expr = self.bb.match_cast(expr, relax.TensorType([1], "int64"))
+        shape_var = self.bb.emit(relax.op.tensor_to_shape(expr))
+        dim = tirx.Var(name, "int64")

Review Comment:
   tirx is intentional here



-- 
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