MCPcopy Create free account
hub / github.com/MotrixLab/FineMoGen / T2MTextEncoder

Class T2MTextEncoder

mogen/models/rnns/t2m_bigru.py:106–158  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

104
105@SUBMODULES.register_module()
106class T2MTextEncoder(nn.Module):
107
108 def __init__(self, word_size, pos_size, hidden_size, output_size,
109 max_text_len):
110 super().__init__()
111 self.text_encoder = TextEncoderBiGRUCo(
112 word_size=word_size,
113 pos_size=pos_size,
114 hidden_size=hidden_size,
115 output_size=output_size,
116 )
117 self.w_vectorizer = WordVectorizer('./data/glove', 'our_vab')
118 self.max_text_len = max_text_len
119
120 def load_pretrained(self, ckpt_path):
121 checkpoint = torch.load(ckpt_path, map_location='cpu')
122 self.text_encoder.load_state_dict(checkpoint['text_encoder'])
123
124 def forward(self, text, token, device):
125 B = len(text)
126 pos_one_hot = []
127 word_emb = []
128 sent_len = []
129 for i in range(B):
130 tokens = token[i].split(" ")
131 if len(tokens) < self.max_text_len:
132 tokens = ['sos/OTHER'] + tokens + ['eos/OTHER']
133 batch_sent_len = len(tokens)
134 tokens = tokens + \
135 ['unk/OTHER'] * (self.max_text_len + 2 - batch_sent_len)
136 else:
137 tokens = tokens[:self.max_text_len]
138 tokens = ['sos/OTHER'] + tokens + ['eos/OTHER']
139 batch_sent_len = len(tokens)
140 sent_len.append(batch_sent_len)
141 batch_word_emb = []
142 batch_pos_one_hot = []
143 for cur_token in tokens:
144 cur_word_emb, cur_pos_one_hot = self.w_vectorizer[cur_token]
145 cur_word_emb = torch.from_numpy(cur_word_emb).float()
146 cur_pos_one_hot = torch.from_numpy(cur_pos_one_hot).float()
147 batch_word_emb.append(cur_word_emb)
148 batch_pos_one_hot.append(cur_pos_one_hot)
149
150 batch_word_emb = torch.stack(batch_word_emb, dim=0)
151 batch_pos_one_hot = torch.stack(batch_pos_one_hot, dim=0)
152 word_emb.append(batch_word_emb)
153 pos_one_hot.append(batch_pos_one_hot)
154 word_emb = torch.stack(word_emb, dim=0).to(device)
155 pos_one_hot = torch.stack(pos_one_hot, dim=0).to(device)
156 sent_len = torch.tensor(sent_len, dtype=torch.long).to(device)
157 text_embedding = self.text_encoder(word_emb, pos_one_hot, sent_len)
158 return text_embedding
159
160
161class TextEncoderBiGRUCo(nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected