from typing import List import torch import torch.nn as nn class Vocab: def __init__(self, char2idx=None, idx2char=None): if char2idx is None: char2idx = { "": 0, "": 1, "": 2, "": 3, } self.char2idx = char2idx if idx2char is None: self.idx2char = {i: c for c, i in self.char2idx.items()} else: self.idx2char = {int(k): v for k, v in idx2char.items()} def encode(self, s: str) -> List[int]: unk = self.char2idx[""] return [self.char2idx.get(ch, unk) for ch in s] def decode(self, ids: List[int]) -> str: out = [] eos_id = self.char2idx[""] for i in ids: if i == eos_id: break if i > eos_id: out.append(self.idx2char.get(int(i), "")) return "".join(out) class LemmaModel(nn.Module): def __init__( self, vocab_size: int, char_emb_dim: int = 96, hidden_size: int = 128, drop_prob: float = 0.30, num_heads: int = 16, max_gen_len: int = 30, ): super().__init__() self.max_gen_len = max_gen_len self.emb = nn.Embedding(vocab_size, char_emb_dim, padding_idx=0) self.dropout_enc = nn.Dropout(drop_prob) self.dropout_dec = nn.Dropout(drop_prob) self.dropout_att = nn.Dropout(drop_prob) self.enc1 = nn.LSTM(char_emb_dim, hidden_size, bidirectional=True, batch_first=True) self.enc2 = nn.LSTM(hidden_size * 2, hidden_size, bidirectional=True, batch_first=True) self.attn = nn.MultiheadAttention(hidden_size * 2, num_heads, batch_first=True) self.dec = nn.LSTM(char_emb_dim + hidden_size * 4, hidden_size * 2, batch_first=True) self.dec_cross_attn = nn.MultiheadAttention( embed_dim=hidden_size * 2, num_heads=num_heads, kdim=hidden_size * 4, vdim=hidden_size * 4, batch_first=True, ) self.out = nn.Linear(hidden_size * 2, vocab_size, bias=True) def encode(self, src, src_lens): emb = self.emb(src) packed1 = nn.utils.rnn.pack_padded_sequence( emb, src_lens.cpu(), batch_first=True, enforce_sorted=False ) enc1_o, _ = self.enc1(packed1) enc1_o, _ = nn.utils.rnn.pad_packed_sequence(enc1_o, batch_first=True) enc1_o = self.dropout_enc(enc1_o) packed2 = nn.utils.rnn.pack_padded_sequence( enc1_o, src_lens.cpu(), batch_first=True, enforce_sorted=False ) enc2_o, _ = self.enc2(packed2) enc2_o, _ = nn.utils.rnn.pad_packed_sequence(enc2_o, batch_first=True) enc2_o = self.dropout_enc(enc2_o) attn_o, _ = self.attn(enc1_o, enc2_o, enc2_o) attn_o = self.dropout_att(attn_o) return torch.cat([enc2_o, attn_o], dim=-1) def forward(self, src, src_lens, tgt): encoder_combined = self.encode(src, src_lens) dt = self.emb(tgt[:, :-1]) target_len = dt.size(1) if encoder_combined.size(1) >= target_len: comb_trim = encoder_combined[:, :target_len, :] else: pad = encoder_combined.new_zeros( encoder_combined.size(0), target_len - encoder_combined.size(1), encoder_combined.size(2), ) comb_trim = torch.cat([encoder_combined, pad], dim=1) dec_inp = torch.cat([dt, comb_trim], dim=-1) dec_o, _ = self.dec(dec_inp) dec_o = self.dropout_dec(dec_o) cross_out, _ = self.dec_cross_attn(dec_o, encoder_combined, encoder_combined) cross_out = self.dropout_att(cross_out) return self.out(cross_out) def generate(self, src, src_lens, vocab, max_len=None): self.eval() if max_len is None: max_len = self.max_gen_len batch_size = src.size(0) with torch.no_grad(): encoder_combined = self.encode(src, src_lens) source_len = encoder_combined.size(1) cur = torch.full( (batch_size, 1), vocab.char2idx[""], device=src.device, dtype=torch.long ) hidden = None hyps = [[] for _ in range(batch_size)] for step in range(max_len): emb_t = self.emb(cur).squeeze(1) if source_len == 0: comb_t = encoder_combined[:, 0, :] else: comb_t = encoder_combined[:, min(step, source_len - 1), :] dec_inp_t = torch.cat([emb_t, comb_t], dim=-1).unsqueeze(1) dec_o, hidden = self.dec(dec_inp_t, hidden) dec_o = self.dropout_dec(dec_o) cross_out, _ = self.dec_cross_attn(dec_o, encoder_combined, encoder_combined) cross_out = self.dropout_att(cross_out) logits = self.out(cross_out) cur = logits.argmax(-1) for i in range(batch_size): hyps[i].append(int(cur[i, 0].item())) if all(int(cur[i, 0].item()) == vocab.char2idx[""] for i in range(batch_size)): break return [vocab.decode(h) for h in hyps]