ML: Рекурентні мережі на PyTorch
Вступ
Encoder-Decoder
Нехай є дві різні рекурентні мережі. Перша називається Encoder, а друга - Decoder. Енкодер на вхід отримує текст однією мовою (source), а декодер повинен на виході видати текст іншою мовою (target). Остання комірка RNN енкодера містить на виході вектор прихованого стану $\mathbf{h}_e^{\text{last}}$ своєї останньої комірки. Цей вектор "зберігає" в собі інформацію про все source-речення (context vector). Його надсилають як початковий прихований стан у першу комірку RNN декодера.
Потім на вхід першої комірки декодера подається службовий токен <BOS> (begin of sentence). На виході комірки повинно з'явитися слово-переклад "кіт". Це означає, що вихід кожної комірки пропускається через лінійний шар з числом нейронів рівним числу слів у словнику. Потім softmax-функція, видає ймовірності слів з яких вибирається номер максимальної (argmax). Отримане слово "кіт" передається на вхід другої комірки і т.д. поки не вийде службовий токен <ЕOS> (end of sentence).
Більш швидкий, але не такий якісний режим тренування називається примусове навчання (teacher forcing). У цьому випадку на всі входи декодера одразу подають правильний target-переклад, а на виході від нього вимагають видати це речення зсунуте на одне слово вліво. Зазвичай між режимами "чесного" і "примусового" навчання відбувається випадкове перемикання.
Реалізація Encoder-Decoder
Наведемо реалізацію архітектури Encoder-Decoder на PyTorch. Нехай у словнику source-мови VOC_SIZE і розмірність векторів ембедінгу цих слів E = VEC_DIM. Тоді модуль енкодера має вигляд:
VEC_DIM = 100
class EncoderRNN(nn.Module):
def __init__(self, VOC_SIZE, E): # розміри словника і ембедінгу
super(EncoderRNN, self).__init__()
self.emb = nn.Embedding(VOC_SIZE, E, scale_grad_by_freq=True)
self.rnn = nn.GRU(E, E, bidirectional=True) # двонаправлена GRU
def forward(self, X):
""" X:(B,L) B - речень з L словами в кожному. Нуль - відсутність слова """
lens = torch.tensor([ len(x)-len(x[x==0]) for x in X ])
emb = self.emb( X.t() ) # (B,L) -> (L,B) -> (L,B,E)
Xp = pack_padded_sequence(emb, lens, enforce_sorted=False)
_, Hn = self.rnn(Xp) # (2,B,E)
Hn = torch.cat([Hn[0],Hn[1]], dim=1) # (B,2*E)
return Hn.view(1,-1, Hn.size(1)) # (1,B,2*E) тільки прихований стан
В енкодер будемо засилати по одному реченню (batch_size=1) змінної довжини L у вигляді вектора L цілих чисел (long). На виході енкодер повертає пару Y - тензор (L,1,E) виходів усіх комірок і вихід Hid останньої комірки.
class DecoderRNN(nn.Module):
def __init__(self, VOC_SIZE, E):
super(DecoderRNN, self).__init__()
self.emb = nn.Embedding(VOC_SIZE, E, scale_grad_by_freq=True)
self.rnn = nn.GRU(E, 2*E) # hidden з 2-направленої
self.out = nn.Linear(2*E, VOC_SIZE)
def forward(self, Hid, X = None, forcing = False): # Hid:(1,B,2*E), X:(B,L)
max_len = MAX_EN_LEN if X is None else len(X[0]) # максимальна довжина речення
W = torch.empty( Hid.size(1), dtype=torch.long ).fill_(BOS_INDEX)
Wrds = torch.zeros( Hid.size(1), max_len, dtype=torch.long ) # передбач. слова
Prbs = torch.ones ( Hid.size(1), max_len, dtype=torch.float ) # ймовірності
for i in range(max_len): #
W = self.emb( W.view(1,-1) ) # (1,B,E)
Y, Hid = self.rnn(W, Hid) # (1,B,2*E)
Y = self.out (Y[0]) # (B,VOC_SIZE)
Y = torch.softmax( Y, dim=1 ) # (B,VOC_SIZE)
_, W = Y.detach().topk(1, dim=1) # (B,1)
Wrds[:,i].copy_(W.squeeze()) # прибираємо 1 і зберігаємо
if not X is None and i < X.size(1):
for b in range(X.size(0)): Prbs[b,i] = Y[b, X[b,i]]
if forcing:
W.copy_( X[:,i].view(-1,1) )
return Wrds, Prbs # (B,L), (B,L)
Сумарна модель:
class EncoderDecoderRNN(nn.Module):
def __init__(self, encoder, decoder):
super(EncoderDecoderRNN, self).__init__()
self.enc = encoder
self.dec = decoder
def forward(self, sourse, target, forcing=False): # (L1,) (L2,)
Hd = self.enc(sourse)
return self.dec(Hd, target, forcing)
gpu = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
cpu = torch.device("cpu")
encoder = EncoderRNN(len(voc_en), VEC_DIM)
decoder = DecoderRNN(len(voc_ru), VEC_DIM)
model = EncoderDecoderRNN(encoder, decoder) # екземпляр мережі
model.to(gpu)