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:

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]