gemini-code-assist[bot] commented on code in PR #19837:
URL: https://github.com/apache/tvm/pull/19837#discussion_r3486270507


##########
python/tvm/relax/frontend/torch/exported_program_translator.py:
##########
@@ -918,6 +918,200 @@ def _gru(self, node: fx.Node) -> relax.Var:
 
         return output
 
+    def _rnn_tanh_cell_unroll(
+        self,
+        input_reshaped,
+        weight_ih,
+        weight_hh,
+        bias_ih,
+        bias_hh,
+        h_prev,
+        seq_len,
+        reverse=False,
+    ):
+        """Unroll vanilla tanh-RNN cells for a single direction."""
+        # Transpose weights for matmul: (hidden_size, in) -> (in, hidden_size)
+        weight_ih_t = self.block_builder.emit(relax.op.permute_dims(weight_ih, 
axes=[1, 0]))
+        weight_hh_t = self.block_builder.emit(relax.op.permute_dims(weight_hh, 
axes=[1, 0]))
+
+        bias = None
+        if bias_ih is not None and bias_hh is not None:
+            bias = self.block_builder.emit(relax.op.add(bias_ih, bias_hh))
+
+        outputs = []
+        time_steps = range(seq_len - 1, -1, -1) if reverse else range(seq_len)
+
+        for t in time_steps:
+            # Input at time t: (batch_size, input_size)
+            x_t = self.block_builder.emit(
+                relax.op.take(input_reshaped, relax.const(t, "int64"), axis=0, 
mode="clip")
+            )
+
+            # h_t = tanh(W_ih @ x_t + W_hh @ h_{t-1} + (b_ih + b_hh))
+            ih = self.block_builder.emit(relax.op.linear_algebra.matmul(x_t, 
weight_ih_t))
+            hh = 
self.block_builder.emit(relax.op.linear_algebra.matmul(h_prev, weight_hh_t))
+            ih_hh = self.block_builder.emit(relax.op.add(ih, hh))
+            if bias is not None:
+                ih_hh = self.block_builder.emit(relax.op.add(ih_hh, bias))
+            h_t = self.block_builder.emit(relax.op.tanh(ih_hh))
+
+            outputs.append(h_t)
+            h_prev = h_t
+
+        if reverse:
+            outputs = outputs[::-1]
+
+        output = self.block_builder.emit(relax.op.stack(outputs, axis=0))
+        # 'h_prev' is the hidden state after the final processed time step 
(this direction' s h_n)
+        # independent of the output-sequence ordering above.
+        return output, h_prev
+
+    def _rnn_tanh(self, node: fx.Node) -> relax.Var:
+        args = self.retrieve_args(node)
+        input_tensor = args[0]
+        hx = args[1] if len(args) > 1 else None
+        params = args[2] if len(args) > 2 else None
+        has_biases = args[3] if len(args) > 3 else True
+        num_layers = args[4] if len(args) > 4 else 1
+        _dropout = args[5] if len(args) > 5 else 0.0  # Not used in inference
+        _train = args[6] if len(args) > 6 else False  # Not used in inference
+        bidirectional = args[7] if len(args) > 7 else False
+        batch_first = args[8] if len(args) > 8 else False
+
+        if num_layers > 1:
+            raise NotImplementedError("Multi-layer RNN is not yet supported")
+
+        def _node_meta(fx_node):
+            meta = fx_node.meta
+            return meta["val"] if "val" in meta else meta["tensor_meta"]
+
+        input_meta = _node_meta(node.args[0])
+        input_shape = list(input_meta.shape)
+        if batch_first:
+            batch_size, seq_len, input_size = input_shape
+        else:
+            seq_len, batch_size, input_size = input_shape
+
+        if not isinstance(seq_len, int):
+            raise NotImplementedError("Dynamic sequence length is not 
supported for rnn_tanh")
+
+        # params per direction: weight_ih, weight_hh, [bias_ih, bias_hh]
+        params_per_direction = 4 if has_biases else 2
+
+        # A vanilla RNN has a single gate, so weight_ih has shape 
(hidden_size, input_size)
+        if params and len(params) >= 2:
+            hidden_size = int(_node_meta(node.args[2][0]).shape[0])
+        else:
+            hidden_size = 16
+
+        dtype = self._convert_data_type(input_meta.dtype)

Review Comment:
   ![high](https://www.gstatic.com/codereviewagent/high-priority.svg)
   
   The `_node_meta` helper function accesses the raw FX node's metadata 
(`node.args[0].meta`) to retrieve shapes and dtypes. This is fragile and 
inconsistent with the rest of the translator (such as `_lstm` and `_gru`), 
which safely query the translated Relax expressions using 
`self.shape_of(input_tensor)` and `input_tensor.struct_info.dtype`. Querying 
the Relax expressions directly is much more robust and allows us to completely 
remove the `_node_meta` helper.
   
   ```python
           input_shape = self.shape_of(input_tensor)
           if batch_first:
               batch_size, seq_len, input_size = input_shape
           else:
               seq_len, batch_size, input_size = input_shape
   
           seq_len = int(seq_len) if isinstance(seq_len, tvm.tirx.IntImm) else 
seq_len
           batch_size = int(batch_size) if isinstance(batch_size, 
tvm.tirx.IntImm) else batch_size
           input_size = int(input_size) if isinstance(input_size, 
tvm.tirx.IntImm) else input_size
   
           if not isinstance(seq_len, int):
               raise NotImplementedError("Dynamic sequence length is not 
supported for rnn_tanh")
   
           # params per direction: weight_ih, weight_hh, [bias_ih, bias_hh]
           params_per_direction = 4 if has_biases else 2
   
           # A vanilla RNN has a single gate, so weight_ih has shape 
(hidden_size, input_size)
           if params and len(params) >= 2:
               hidden_size = int(self.shape_of(params[0])[0])
           else:
               hidden_size = 16
   
           dtype = input_tensor.struct_info.dtype
   ```



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