cchung100m commented on code in PR #19837:
URL: https://github.com/apache/tvm/pull/19837#discussion_r3486324235
##########
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:
Thanks for the suggestion. I looked into this and kept `_node_meta`
intentionally - reverting to `self_shape_of(input_tensor)` /
`input_tensor.struct_info.dtype` would re-break the now-passing `test_rnn_tanh`.
>[2026-06-22T16:46:09.119Z] =================================== FAILURES
===================================
[2026-06-22T16:46:09.119Z] ________________________________ test_rnn_tanh
_________________________________
[2026-06-22T16:46:09.119Z] [gw0] linux -- Python 3.10.19
/venv/apache-tvm-py3.10/bin/python3
[2026-06-22T16:46:09.119Z]
tests/python/relax/test_frontend_from_exported_program.py:8604: in test_rnn_tanh
[2026-06-22T16:46:09.119Z] _check(
[2026-06-22T16:46:09.119Z]
tests/python/relax/test_frontend_from_exported_program.py:8587: in _check
[2026-06-22T16:46:09.119Z] mod = from_exported_program(exported_program,
run_ep_decomposition=False)
[2026-06-22T16:46:09.119Z]
python/tvm/relax/frontend/torch/exported_program_translator.py:2389: in
from_exported_program
[2026-06-22T16:46:09.119Z] return
ExportedProgramImporter().from_exported_program(
[2026-06-22T16:46:09.119Z]
python/tvm/relax/frontend/torch/exported_program_translator.py:2224: in
from_exported_program
[2026-06-22T16:46:09.119Z] output_args = self._translate_fx_graph(
[2026-06-22T16:46:09.119Z]
python/tvm/relax/frontend/torch/exported_program_translator.py:1533: in
_translate_fx_graph
[2026-06-22T16:46:09.119Z] self.env[node] =
self.convert_map[func_name](node)
[2026-06-22T16:46:09.119Z]
python/tvm/relax/frontend/torch/exported_program_translator.py:1007: in
_rnn_tanh
[2026-06-22T16:46:09.119Z] dtype = input_tensor.struct_info.dtype
[2026-06-22T16:46:09.119Z] E AttributeError: 'Var' object has no attribute
'struct_info'
[2026-06-22T16:46:09.119Z] =============================== warnings summary
===============================
--
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]