| import torch
|
| import torch.nn as nn
|
| from torch.autograd import Variable
|
| import torch.nn.functional as F
|
| import numpy as np
|
| from torch.nn.functional import cross_entropy, softmax
|
| from transformers import BertModel, BertConfig
|
|
|
|
|
| class Eyettention(nn.Module):
|
| def __init__(self, cf):
|
| super(Eyettention, self).__init__()
|
| self.cf = cf
|
| self.window_width = 1
|
| self.atten_type = cf["atten_type"]
|
| self.hidden_size = 128
|
|
|
|
|
| encoder_config = BertConfig.from_pretrained(self.cf["model_pretrained"])
|
| encoder_config.output_hidden_states = True
|
|
|
| print("keeping Bert with pre-trained weights")
|
| self.encoder = BertModel.from_pretrained(self.cf["model_pretrained"], config=encoder_config)
|
| self.encoder.eval()
|
|
|
| for param in self.encoder.parameters():
|
| param.requires_grad = False
|
|
|
| self.embedding_dropout = nn.Dropout(0.4)
|
| self.encoder_lstm = nn.LSTM(
|
| input_size=768,
|
| hidden_size=int(self.hidden_size / 2),
|
| num_layers=8,
|
| batch_first=True,
|
| bidirectional=True,
|
| dropout=0.2,
|
| )
|
|
|
|
|
| self.position_embeddings = nn.Embedding(
|
| encoder_config.max_position_embeddings, encoder_config.hidden_size
|
| )
|
| self.LayerNorm = nn.LayerNorm(encoder_config.hidden_size, eps=encoder_config.layer_norm_eps)
|
|
|
|
|
|
|
|
|
| self.decoder_cell1 = nn.LSTMCell(
|
| 768 + 2, self.hidden_size
|
| )
|
| self.decoder_cell2 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell3 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell4 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell5 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell6 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell7 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell8 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.dropout_LSTM = nn.Dropout(0.2)
|
|
|
|
|
| self.attn = nn.Linear(
|
| self.hidden_size, self.hidden_size + 1
|
| )
|
|
|
|
|
|
|
| self.dropout_dense = nn.Dropout(0.2)
|
| self.decoder_dense = nn.Sequential(
|
| self.dropout_dense,
|
| nn.Linear(self.hidden_size * 2 + 1, 512),
|
| nn.ReLU(),
|
| self.dropout_dense,
|
| nn.Linear(512, 256),
|
| nn.ReLU(),
|
| self.dropout_dense,
|
| nn.Linear(256, 256),
|
| nn.ReLU(),
|
| self.dropout_dense,
|
| nn.Linear(256, 256),
|
| nn.ReLU(),
|
| nn.Linear(256, self.cf["max_sn_len"] * 2 - 3),
|
| )
|
|
|
|
|
| self.softmax = nn.Softmax(dim=1)
|
|
|
| def pool_subword_to_word(self, subword_emb, word_ids_sn, target, pool_method="sum"):
|
|
|
|
|
| if target == "sn":
|
| max_len = self.cf["max_sn_len"]
|
| elif target == "sp":
|
| max_len = self.cf["max_sp_len"] - 1
|
|
|
| merged_word_emb = torch.empty(subword_emb.shape[0], 0, 768).to(subword_emb.device)
|
| for word_idx in range(max_len):
|
| word_mask = (word_ids_sn == word_idx).unsqueeze(2).repeat(1, 1, 768)
|
|
|
| if pool_method == "sum":
|
| pooled_word_emb = torch.sum(subword_emb * word_mask, 1).unsqueeze(
|
| 1
|
| )
|
| elif pool_method == "mean":
|
| pooled_word_emb = torch.mean(subword_emb * word_mask, 1).unsqueeze(
|
| 1
|
| )
|
| merged_word_emb = torch.cat([merged_word_emb, pooled_word_emb], dim=1)
|
|
|
| mask_word = torch.sum(merged_word_emb, 2).bool()
|
| return merged_word_emb, mask_word
|
|
|
| def encode(self, sn_emd, sn_mask, word_ids_sn, sn_word_len):
|
|
|
| outputs = self.encoder(input_ids=sn_emd, attention_mask=sn_mask)
|
| hidden_rep_orig, pooled_rep = outputs[0], outputs[1]
|
| if word_ids_sn != None:
|
|
|
| merged_word_emb, sn_mask_word = self.pool_subword_to_word(
|
| hidden_rep_orig, word_ids_sn, target="sn", pool_method="sum"
|
| )
|
| else:
|
| merged_word_emb, sn_mask_word = hidden_rep_orig, None
|
|
|
| hidden_rep = self.embedding_dropout(merged_word_emb)
|
| x, (hn, hc) = self.encoder_lstm(hidden_rep, None)
|
|
|
|
|
| x = torch.cat((x, sn_word_len[:, :, None]), dim=2)
|
| return x, sn_mask_word
|
|
|
| def cross_attention(self, ht, hs, sn_mask, cur_word_index):
|
|
|
|
|
|
|
|
|
|
|
| attn_prod = torch.matmul(
|
| self.attn(ht.unsqueeze(1)), hs.permute(0, 2, 1)
|
| )
|
| if self.atten_type == "global":
|
| attn_prod += (~sn_mask).unsqueeze(1) * -1e9
|
| att_weight = softmax(attn_prod, dim=2)
|
|
|
| else:
|
|
|
| aligned_position = cur_word_index
|
|
|
| left = torch.where(
|
| aligned_position - self.window_width >= 0, aligned_position - self.window_width, 0
|
| )
|
| right = torch.where(
|
| aligned_position + self.window_width <= self.cf["max_sn_len"] - 1,
|
| aligned_position + self.window_width,
|
| self.cf["max_sn_len"] - 1,
|
| )
|
|
|
|
|
|
|
| sen_seq = (
|
| torch.arange(self.cf["max_sn_len"])[None, :]
|
| .expand(sn_mask.shape[0], self.cf["max_sn_len"])
|
| .to(sn_mask.device)
|
| )
|
| outside_win_mask = (sen_seq < left.unsqueeze(1)) + (sen_seq > right.unsqueeze(1))
|
|
|
| attn_prod += (~sn_mask + outside_win_mask).unsqueeze(1) * -1e9
|
| att_weight = softmax(attn_prod, dim=2)
|
|
|
| if self.atten_type == "local-g":
|
| gauss = lambda s: torch.exp(
|
| -torch.square(s - aligned_position.unsqueeze(1))
|
| / (2 * torch.square(torch.tensor(self.window_width / 2)))
|
| )
|
| gauss_factor = gauss(sen_seq)
|
| att_weight = att_weight * gauss_factor.unsqueeze(1)
|
|
|
| return att_weight
|
|
|
| def decode(self, sp_emd, sn_mask, sp_pos, enc_out, sp_fix_dur, sp_landing_pos, word_ids_sp):
|
|
|
|
|
| hn = torch.zeros(8, sp_emd.shape[0], self.hidden_size).to(sp_emd.device)
|
| hc = torch.zeros(8, sp_emd.shape[0], self.hidden_size).to(sp_emd.device)
|
| hx, cx = hn[0, :, :], hc[0, :, :]
|
| hx2, cx2 = hn[1, :, :], hc[1, :, :]
|
| hx3, cx3 = hn[2, :, :], hc[2, :, :]
|
| hx4, cx4 = hn[3, :, :], hc[3, :, :]
|
| hx5, cx5 = hn[4, :, :], hc[4, :, :]
|
| hx6, cx6 = hn[5, :, :], hc[5, :, :]
|
| hx7, cx7 = hn[6, :, :], hc[6, :, :]
|
| hx8, cx8 = hn[7, :, :], hc[7, :, :]
|
|
|
| dec_emb_in = self.encoder.embeddings.word_embeddings(sp_emd[:, :-1])
|
| if word_ids_sp is not None:
|
|
|
| sp_merged_word_emd, sp_mask_word = self.pool_subword_to_word(
|
| dec_emb_in, word_ids_sp[:, :-1], target="sp", pool_method="sum"
|
| )
|
| else:
|
| sp_merged_word_emd, sp_mask_word = dec_emb_in, None
|
|
|
|
|
| position_embeddings = self.position_embeddings(sp_pos[:, :-1])
|
| dec_emb_in = sp_merged_word_emd + position_embeddings
|
| dec_emb_in = self.LayerNorm(dec_emb_in)
|
| dec_emb_in = dec_emb_in.permute(1, 0, 2)
|
| dec_emb_in = self.embedding_dropout(dec_emb_in)
|
|
|
|
|
| if sp_landing_pos is not None:
|
| dec_emb_in = torch.cat((dec_emb_in, sp_landing_pos.permute(1, 0)[:-1, :, None]), dim=2)
|
|
|
| if sp_fix_dur is not None:
|
| dec_emb_in = torch.cat((dec_emb_in, sp_fix_dur.permute(1, 0)[:-1, :, None]), dim=2)
|
|
|
|
|
| output = []
|
|
|
| atten_weights_batch = torch.empty(sp_emd.shape[0], 0, self.cf["max_sn_len"]).to(
|
| sp_emd.device
|
| )
|
| for i in range(dec_emb_in.shape[0]):
|
| hx, cx = self.decoder_cell1(dec_emb_in[i], (hx, cx))
|
| hx2, cx2 = self.decoder_cell2(self.dropout_LSTM(hx), (hx2, cx2))
|
| hx3, cx3 = self.decoder_cell3(self.dropout_LSTM(hx2), (hx3, cx3))
|
| hx4, cx4 = self.decoder_cell4(self.dropout_LSTM(hx3), (hx4, cx4))
|
| hx5, cx5 = self.decoder_cell5(self.dropout_LSTM(hx4), (hx5, cx5))
|
| hx6, cx6 = self.decoder_cell6(self.dropout_LSTM(hx5), (hx6, cx6))
|
| hx7, cx7 = self.decoder_cell7(self.dropout_LSTM(hx6), (hx7, cx7))
|
| hx8, cx8 = self.decoder_cell8(self.dropout_LSTM(hx7), (hx8, cx8))
|
|
|
| att_weight = self.cross_attention(
|
| ht=hx8, hs=enc_out, sn_mask=sn_mask, cur_word_index=sp_pos[:, i]
|
| )
|
| atten_weights_batch = torch.cat([atten_weights_batch, att_weight], dim=1)
|
|
|
| context = torch.matmul(att_weight, enc_out)
|
|
|
|
|
| hc = torch.cat([context.squeeze(1), hx8], dim=1)
|
| result = self.decoder_dense(hc)
|
| output.append(result)
|
|
|
| output = torch.stack(output, dim=0)
|
|
|
| return output.permute(1, 0, 2), atten_weights_batch
|
|
|
| def forward(
|
| self,
|
| sn_emd,
|
| sn_mask,
|
| sp_emd,
|
| sp_pos,
|
| word_ids_sn,
|
| word_ids_sp,
|
| sp_fix_dur,
|
| sp_landing_pos,
|
| sn_word_len,
|
| ):
|
| x, sn_mask_word = self.encode(
|
| sn_emd, sn_mask, word_ids_sn, sn_word_len
|
| )
|
|
|
| if sn_mask_word is None:
|
| sn_mask = torch.Tensor.bool(sn_mask)
|
| pred, atten_weights = self.decode(
|
| sp_emd, sn_mask, sp_pos, x, sp_fix_dur, sp_landing_pos, word_ids_sp
|
| )
|
|
|
| else:
|
| pred, atten_weights = self.decode(
|
| sp_emd, sn_mask_word, sp_pos, x, sp_fix_dur, sp_landing_pos, word_ids_sp
|
| )
|
|
|
| return pred, atten_weights
|
|
|
| def scanpath_generation(
|
| self, sn_emd, sn_mask, word_ids_sn, sn_word_len, le, max_pred_len=60, previous_scanpath=None
|
| ):
|
|
|
| if max_pred_len <= 0:
|
| raise ValueError("max_pred_len must be positive.")
|
|
|
| prev_scanpath_len = 0
|
|
|
| if previous_scanpath is not None:
|
| if (
|
| len(previous_scanpath) == 0 or previous_scanpath[0] != 0
|
| ):
|
| previous_scanpath = [0] + previous_scanpath
|
|
|
| prev_scanpath_len = len(previous_scanpath)
|
|
|
| previous_scanpath = torch.as_tensor(
|
| previous_scanpath, dtype=torch.long, device=sn_emd.device
|
| )
|
| if previous_scanpath.ndim == 1:
|
| previous_scanpath = previous_scanpath.unsqueeze(0)
|
|
|
| previous_scanpath = previous_scanpath[:, 1:]
|
|
|
|
|
| enc_out, sn_mask_word = self.encode(sn_emd, sn_mask, word_ids_sn, sn_word_len)
|
| if sn_mask_word is None:
|
| sn_mask = torch.Tensor.bool(sn_mask)
|
| else:
|
| sn_mask = sn_mask_word
|
| sn_len = torch.sum(sn_mask, axis=1) - 2
|
|
|
|
|
|
|
| hn = torch.zeros(8, sn_emd.shape[0], self.hidden_size).to(sn_emd.device)
|
| hc = torch.zeros(8, sn_emd.shape[0], self.hidden_size).to(sn_emd.device)
|
| hx, cx = hn[0, :, :], hc[0, :, :]
|
| hx2, cx2 = hn[1, :, :], hc[1, :, :]
|
| hx3, cx3 = hn[2, :, :], hc[2, :, :]
|
| hx4, cx4 = hn[3, :, :], hc[3, :, :]
|
| hx5, cx5 = hn[4, :, :], hc[4, :, :]
|
| hx6, cx6 = hn[5, :, :], hc[5, :, :]
|
| hx7, cx7 = hn[6, :, :], hc[6, :, :]
|
| hx8, cx8 = hn[7, :, :], hc[7, :, :]
|
|
|
|
|
| dec_in_start = (torch.ones(sn_mask.shape[0]) * 101).long().to(sn_mask.device)
|
| dec_emb_in = self.encoder.embeddings.word_embeddings(dec_in_start)
|
|
|
|
|
|
|
|
|
| start_pos = torch.zeros(sn_mask.shape[0]).to(sn_mask.device)
|
| position_embeddings = self.position_embeddings(start_pos.long())
|
| dec_emb_in = dec_emb_in + position_embeddings
|
| dec_emb_in = self.LayerNorm(dec_emb_in)
|
|
|
|
|
| dec_in = torch.cat(
|
| (dec_emb_in, torch.zeros(dec_emb_in.shape[0], 2).to(sn_emd.device)), dim=1
|
| )
|
|
|
|
|
| output = []
|
| density_prediction = []
|
| pred_counter = 0
|
|
|
| output.append(start_pos.long())
|
| for p in range(prev_scanpath_len + max_pred_len - 1):
|
| hx, cx = self.decoder_cell1(dec_in, (hx, cx))
|
| hx2, cx2 = self.decoder_cell2(self.dropout_LSTM(hx), (hx2, cx2))
|
| hx3, cx3 = self.decoder_cell3(self.dropout_LSTM(hx2), (hx3, cx3))
|
| hx4, cx4 = self.decoder_cell4(self.dropout_LSTM(hx3), (hx4, cx4))
|
| hx5, cx5 = self.decoder_cell5(self.dropout_LSTM(hx4), (hx5, cx5))
|
| hx6, cx6 = self.decoder_cell6(self.dropout_LSTM(hx5), (hx6, cx6))
|
| hx7, cx7 = self.decoder_cell7(self.dropout_LSTM(hx6), (hx7, cx7))
|
| hx8, cx8 = self.decoder_cell8(self.dropout_LSTM(hx7), (hx8, cx8))
|
|
|
| att_weight = self.cross_attention(
|
| ht=hx8, hs=enc_out, sn_mask=sn_mask, cur_word_index=output[-1]
|
| )
|
|
|
| context = torch.matmul(att_weight, enc_out)
|
| hc = torch.cat([context.squeeze(1), hx8], dim=1)
|
|
|
| result = self.decoder_dense(hc)
|
| result = self.softmax(result)
|
| density_prediction.append(result)
|
|
|
|
|
|
|
|
|
|
|
| if previous_scanpath is not None and p < previous_scanpath.shape[1]:
|
|
|
| pred_pos = previous_scanpath[:, p].clone()
|
| else:
|
| pred_indx = torch.multinomial(result, 1)
|
| pred_class = [le.classes_[pred_indx[i]] for i in torch.arange(result.shape[0])]
|
| pred_class = torch.from_numpy(np.array(pred_class)).to(sn_emd.device)
|
|
|
| pred_pos = output[-1] + pred_class
|
|
|
|
|
|
|
| input_ids = []
|
| for i in range(pred_pos.shape[0]):
|
| if pred_pos[i] > sn_len[i]:
|
| pred_pos[i] = sn_len[i] + 1
|
| elif pred_pos[i] < 1:
|
| pred_pos[i] = 1
|
|
|
| if word_ids_sn is not None:
|
| input_ids.append(sn_emd[i, word_ids_sn[i, :] == pred_pos[i]])
|
| else:
|
| input_ids.append(sn_emd[i, pred_pos[i]])
|
| output.append(pred_pos)
|
|
|
|
|
| pred_counter += 1
|
| if word_ids_sn is not None:
|
|
|
| dec_emb_in = torch.empty(0, 768).to(sn_emd.device)
|
| for id in input_ids:
|
| dec_emb_in = torch.cat(
|
| [
|
| dec_emb_in,
|
| torch.sum(self.encoder.embeddings.word_embeddings(id), axis=0)[None, :],
|
| ],
|
| dim=0,
|
| )
|
|
|
| else:
|
| input_ids = torch.stack(input_ids)
|
| dec_emb_in = self.encoder.embeddings.word_embeddings(input_ids)
|
|
|
| position_embeddings = self.position_embeddings(output[-1])
|
| dec_emb_in = dec_emb_in + position_embeddings
|
| dec_emb_in = self.LayerNorm(dec_emb_in)
|
|
|
| dec_in = torch.cat(
|
| (dec_emb_in, torch.zeros(dec_emb_in.shape[0], 2).to(sn_emd.device)), dim=1
|
| )
|
|
|
| output = torch.stack(output, dim=0)
|
| return output.permute(1, 0), density_prediction
|
|
|
|
|
| class Eyettention_readerID(nn.Module):
|
| def __init__(self, cf):
|
| super(Eyettention_readerID, self).__init__()
|
| self.cf = cf
|
| self.window_width = 1
|
| self.atten_type = cf["atten_type"]
|
| self.hidden_size = 128
|
| self.sub_emb_size = cf["subid_emb_size"]
|
|
|
|
|
| encoder_config = BertConfig.from_pretrained(self.cf["model_pretrained"])
|
| encoder_config.output_hidden_states = True
|
|
|
| print("keeping Bert with pre-trained weights")
|
| self.encoder = BertModel.from_pretrained(self.cf["model_pretrained"], config=encoder_config)
|
| self.encoder.eval()
|
|
|
| for param in self.encoder.parameters():
|
| param.requires_grad = False
|
|
|
| self.embedding_dropout = nn.Dropout(0.4)
|
| self.encoder_lstm = nn.LSTM(
|
| input_size=768,
|
| hidden_size=int(self.hidden_size / 2),
|
| num_layers=8,
|
| batch_first=True,
|
| bidirectional=True,
|
| dropout=0.2,
|
| )
|
|
|
|
|
| self.position_embeddings = nn.Embedding(
|
| encoder_config.max_position_embeddings, encoder_config.hidden_size
|
| )
|
| self.LayerNorm = nn.LayerNorm(encoder_config.hidden_size, eps=encoder_config.layer_norm_eps)
|
|
|
| self.sub_embeddings = nn.Embedding(400, self.sub_emb_size)
|
|
|
|
|
|
|
|
|
| self.decoder_cell1 = nn.LSTMCell(
|
| 768 + 2 + self.sub_emb_size, self.hidden_size
|
| )
|
| self.decoder_cell2 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell3 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell4 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell5 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell6 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell7 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.decoder_cell8 = nn.LSTMCell(self.hidden_size, self.hidden_size)
|
| self.dropout_LSTM = nn.Dropout(0.2)
|
|
|
|
|
| self.attn = nn.Linear(
|
| self.hidden_size, self.hidden_size + 1
|
| )
|
|
|
|
|
|
|
| self.dropout_dense = nn.Dropout(0.2)
|
| self.decoder_dense = nn.Sequential(
|
| self.dropout_dense,
|
| nn.Linear(self.hidden_size * 2 + 1, 512),
|
| nn.ReLU(),
|
| self.dropout_dense,
|
| nn.Linear(512, 256),
|
| nn.ReLU(),
|
| self.dropout_dense,
|
| nn.Linear(256, 256),
|
| nn.ReLU(),
|
| self.dropout_dense,
|
| nn.Linear(256, 256),
|
| nn.ReLU(),
|
| nn.Linear(256, self.cf["max_sn_len"] * 2 - 3),
|
| )
|
|
|
|
|
| self.softmax = nn.Softmax(dim=1)
|
|
|
| def pool_subword_to_word(self, subword_emb, word_ids_sn, target, pool_method="sum"):
|
|
|
|
|
| if target == "sn":
|
| max_len = self.cf["max_sn_len"]
|
| elif target == "sp":
|
| max_len = self.cf["max_sp_len"] - 1
|
|
|
| merged_word_emb = torch.empty(subword_emb.shape[0], 0, 768).to(subword_emb.device)
|
| for word_idx in range(max_len):
|
| word_mask = (word_ids_sn == word_idx).unsqueeze(2).repeat(1, 1, 768)
|
|
|
| if pool_method == "sum":
|
| pooled_word_emb = torch.sum(subword_emb * word_mask, 1).unsqueeze(
|
| 1
|
| )
|
| elif pool_method == "mean":
|
| pooled_word_emb = torch.mean(subword_emb * word_mask, 1).unsqueeze(
|
| 1
|
| )
|
| merged_word_emb = torch.cat([merged_word_emb, pooled_word_emb], dim=1)
|
|
|
| mask_word = torch.sum(merged_word_emb, 2).bool()
|
| return merged_word_emb, mask_word
|
|
|
| def encode(self, sn_emd, sn_mask, word_ids_sn, sn_word_len):
|
|
|
| outputs = self.encoder(input_ids=sn_emd, attention_mask=sn_mask)
|
| hidden_rep_orig, pooled_rep = outputs[0], outputs[1]
|
| if word_ids_sn != None:
|
|
|
| merged_word_emb, sn_mask_word = self.pool_subword_to_word(
|
| hidden_rep_orig, word_ids_sn, target="sn", pool_method="sum"
|
| )
|
| else:
|
| merged_word_emb, sn_mask_word = hidden_rep_orig, None
|
|
|
| hidden_rep = self.embedding_dropout(merged_word_emb)
|
| x, (hn, hc) = self.encoder_lstm(hidden_rep, None)
|
|
|
|
|
| x = torch.cat((x, sn_word_len[:, :, None]), dim=2)
|
| return x, sn_mask_word
|
|
|
| def cross_attention(self, ht, hs, sn_mask, cur_word_index):
|
|
|
|
|
|
|
|
|
|
|
| attn_prod = torch.matmul(
|
| self.attn(ht.unsqueeze(1)), hs.permute(0, 2, 1)
|
| )
|
| if self.atten_type == "global":
|
| attn_prod += (~sn_mask).unsqueeze(1) * -1e9
|
| att_weight = softmax(attn_prod, dim=2)
|
|
|
| else:
|
|
|
| aligned_position = cur_word_index
|
|
|
| left = torch.where(
|
| aligned_position - self.window_width >= 0, aligned_position - self.window_width, 0
|
| )
|
| right = torch.where(
|
| aligned_position + self.window_width <= self.cf["max_sn_len"] - 1,
|
| aligned_position + self.window_width,
|
| self.cf["max_sn_len"] - 1,
|
| )
|
|
|
|
|
|
|
| sen_seq = (
|
| torch.arange(self.cf["max_sn_len"])[None, :]
|
| .expand(sn_mask.shape[0], self.cf["max_sn_len"])
|
| .to(sn_mask.device)
|
| )
|
| outside_win_mask = (sen_seq < left.unsqueeze(1)) + (sen_seq > right.unsqueeze(1))
|
|
|
| attn_prod += (~sn_mask + outside_win_mask).unsqueeze(1) * -1e9
|
| att_weight = softmax(attn_prod, dim=2)
|
|
|
| if self.atten_type == "local-g":
|
| gauss = lambda s: torch.exp(
|
| -torch.square(s - aligned_position.unsqueeze(1))
|
| / (2 * torch.square(torch.tensor(self.window_width / 2)))
|
| )
|
| gauss_factor = gauss(sen_seq)
|
| att_weight = att_weight * gauss_factor.unsqueeze(1)
|
|
|
| return att_weight
|
|
|
| def decode(
|
| self, sp_emd, sn_mask, sp_pos, enc_out, sp_fix_dur, sp_landing_pos, word_ids_sp, sub_id
|
| ):
|
|
|
|
|
| hn = torch.zeros(8, sp_emd.shape[0], self.hidden_size).to(sp_emd.device)
|
| hc = torch.zeros(8, sp_emd.shape[0], self.hidden_size).to(sp_emd.device)
|
| hx, cx = hn[0, :, :], hc[0, :, :]
|
| hx2, cx2 = hn[1, :, :], hc[1, :, :]
|
| hx3, cx3 = hn[2, :, :], hc[2, :, :]
|
| hx4, cx4 = hn[3, :, :], hc[3, :, :]
|
| hx5, cx5 = hn[4, :, :], hc[4, :, :]
|
| hx6, cx6 = hn[5, :, :], hc[5, :, :]
|
| hx7, cx7 = hn[6, :, :], hc[6, :, :]
|
| hx8, cx8 = hn[7, :, :], hc[7, :, :]
|
|
|
| dec_emb_in = self.encoder.embeddings.word_embeddings(sp_emd[:, :-1])
|
| if word_ids_sp is not None:
|
|
|
| sp_merged_word_emd, sp_mask_word = self.pool_subword_to_word(
|
| dec_emb_in, word_ids_sp[:, :-1], target="sp", pool_method="sum"
|
| )
|
| else:
|
| sp_merged_word_emd, sp_mask_word = dec_emb_in, None
|
|
|
|
|
| position_embeddings = self.position_embeddings(sp_pos[:, :-1])
|
| dec_emb_in = sp_merged_word_emd + position_embeddings
|
| dec_emb_in = self.LayerNorm(dec_emb_in)
|
| dec_emb_in = dec_emb_in.permute(1, 0, 2)
|
| dec_emb_in = self.embedding_dropout(dec_emb_in)
|
|
|
|
|
| if sp_landing_pos is not None:
|
| dec_emb_in = torch.cat((dec_emb_in, sp_landing_pos.permute(1, 0)[:-1, :, None]), dim=2)
|
|
|
| if sp_fix_dur is not None:
|
| dec_emb_in = torch.cat((dec_emb_in, sp_fix_dur.permute(1, 0)[:-1, :, None]), dim=2)
|
|
|
|
|
| if sub_id is not None:
|
| dec_emb_in = torch.cat(
|
| (dec_emb_in, self.sub_embeddings(sub_id).repeat(dec_emb_in.shape[0], 1, 1)), dim=2
|
| )
|
|
|
|
|
| output = []
|
|
|
| atten_weights_batch = torch.empty(sp_emd.shape[0], 0, self.cf["max_sn_len"]).to(
|
| sp_emd.device
|
| )
|
| for i in range(dec_emb_in.shape[0]):
|
| hx, cx = self.decoder_cell1(dec_emb_in[i], (hx, cx))
|
| hx2, cx2 = self.decoder_cell2(self.dropout_LSTM(hx), (hx2, cx2))
|
| hx3, cx3 = self.decoder_cell3(self.dropout_LSTM(hx2), (hx3, cx3))
|
| hx4, cx4 = self.decoder_cell4(self.dropout_LSTM(hx3), (hx4, cx4))
|
| hx5, cx5 = self.decoder_cell5(self.dropout_LSTM(hx4), (hx5, cx5))
|
| hx6, cx6 = self.decoder_cell6(self.dropout_LSTM(hx5), (hx6, cx6))
|
| hx7, cx7 = self.decoder_cell7(self.dropout_LSTM(hx6), (hx7, cx7))
|
| hx8, cx8 = self.decoder_cell8(self.dropout_LSTM(hx7), (hx8, cx8))
|
|
|
| att_weight = self.cross_attention(
|
| ht=hx8, hs=enc_out, sn_mask=sn_mask, cur_word_index=sp_pos[:, i]
|
| )
|
| atten_weights_batch = torch.cat([atten_weights_batch, att_weight], dim=1)
|
|
|
| context = torch.matmul(att_weight, enc_out)
|
|
|
|
|
| hc = torch.cat([context.squeeze(1), hx8], dim=1)
|
| result = self.decoder_dense(hc)
|
| output.append(result)
|
|
|
| output = torch.stack(output, dim=0)
|
|
|
| return output.permute(1, 0, 2), atten_weights_batch
|
|
|
| def forward(
|
| self,
|
| sn_emd,
|
| sn_mask,
|
| sp_emd,
|
| sp_pos,
|
| word_ids_sn,
|
| word_ids_sp,
|
| sp_fix_dur,
|
| sp_landing_pos,
|
| sn_word_len,
|
| sub_id,
|
| ):
|
| x, sn_mask_word = self.encode(
|
| sn_emd, sn_mask, word_ids_sn, sn_word_len
|
| )
|
|
|
| if sn_mask_word is None:
|
| sn_mask = torch.Tensor.bool(sn_mask)
|
| pred, atten_weights = self.decode(
|
| sp_emd, sn_mask, sp_pos, x, sp_fix_dur, sp_landing_pos, word_ids_sp, sub_id
|
| )
|
|
|
| else:
|
| pred, atten_weights = self.decode(
|
| sp_emd, sn_mask_word, sp_pos, x, sp_fix_dur, sp_landing_pos, word_ids_sp, sub_id
|
| )
|
| return pred, atten_weights
|
|
|