Skip to content
This repository was archived by the owner on Nov 22, 2022. It is now read-only.

Commit 4d401d3

Browse files
shreydesaifacebook-github-bot
authored andcommitted
Making seq2seq_model more torchscript-friendly
Summary: seq2seq_model.py changes: - removes dependence on _ when unpacking the tensor dicts - explicit None checks on dict feats Differential Revision: D23673319 fbshipit-source-id: a50efc6418af57cb2fa89273ddb14d51ca13b0e5
1 parent 3e7b626 commit 4d401d3

1 file changed

Lines changed: 17 additions & 5 deletions

File tree

pytext/models/seq_models/seq2seq_model.py

Lines changed: 17 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,11 @@ class Seq2SeqModel(Model):
3535
Sequence to sequence model using an encoder-decoder architecture.
3636
"""
3737

38+
SRC_TOKENS_TENSORIZER_INDEX = 0
39+
SRC_LENGTHS_TENSORIZER_INDEX = 1
40+
TRG_TOKENS_TENSORIZER_INDEX = 0
41+
TRG_LENGTHS_TENSORIZER_INDEX = 1
42+
3843
class Config(Model.Config):
3944
class ModelInput(Model.Config.ModelInput):
4045
src_seq_tokens: TokenTensorizer.Config = TokenTensorizer.Config()
@@ -107,8 +112,13 @@ def arrange_model_inputs(
107112
torch.Tensor,
108113
torch.Tensor,
109114
]:
110-
src_tokens, src_lengths, _ = tensor_dict["src_seq_tokens"]
111-
trg_tokens, trg_lengths, _ = tensor_dict["trg_seq_tokens"]
115+
src_seq_tokens = tensor_dict["src_seq_tokens"]
116+
trg_seq_tokens = tensor_dict["trg_seq_tokens"]
117+
118+
src_tokens = src_seq_tokens[self.SRC_TOKENS_TENSORIZER_INDEX]
119+
src_lengths = src_seq_tokens[self.SRC_LENGTHS_TENSORIZER_INDEX]
120+
trg_tokens = trg_seq_tokens[self.TRG_TOKENS_TENSORIZER_INDEX]
121+
trg_lengths = trg_seq_tokens[self.TRG_LENGTHS_TENSORIZER_INDEX]
112122

113123
def _shift_target(in_sequences, seq_lens, eos_idx, pad_idx):
114124
shifted_sequence = GetTensor(
@@ -136,7 +146,9 @@ def _shift_target(in_sequences, seq_lens, eos_idx, pad_idx):
136146
)
137147

138148
def arrange_targets(self, tensor_dict):
139-
trg_tokens, trg_lengths, _ = tensor_dict["trg_seq_tokens"]
149+
trg_seq_tokens = tensor_dict["trg_seq_tokens"]
150+
trg_tokens = trg_seq_tokens[self.TRG_TOKENS_TENSORIZER_INDEX]
151+
trg_lengths = trg_seq_tokens[self.TRG_LENGTHS_TENSORIZER_INDEX]
140152
return (trg_tokens, trg_lengths)
141153

142154
def __init__(
@@ -196,7 +208,7 @@ def forward(
196208
):
197209
additional_features: List[List[torch.Tensor]] = []
198210

199-
if dict_feats:
211+
if dict_feats is not None:
200212
additional_features.append(list(dict_feats))
201213

202214
if contextual_token_embedding is not None:
@@ -206,7 +218,7 @@ def forward(
206218
src_tokens, additional_features, src_lengths, trg_tokens
207219
)
208220

209-
if dict_feats:
221+
if dict_feats is not None:
210222
(
211223
output_dict["dict_tokens"],
212224
output_dict["dict_weights"],

0 commit comments

Comments
 (0)