"""Transformer-Encoder-Decoder für Fehlerkorrektur, von Grund auf (keine vortrainierten Gewichte).""" import math import torch import torch.nn as nn import torch.nn.functional as F PAD, UNK, BOS, EOS = 0, 1, 2, 3 class GEC(nn.Module): def __init__(self, vocab, dim=512, layers=6, heads=8, ffn=2048, dropout=0.1, max_len=256): super().__init__() self.cfg = dict(vocab=vocab, dim=dim, layers=layers, heads=heads, ffn=ffn, dropout=dropout, max_len=max_len) self.emb = nn.Embedding(vocab, dim, padding_idx=PAD) # geteilt: Encoder, Decoder, Ausgabe nn.init.normal_(self.emb.weight, std=dim ** -0.5) pos = torch.arange(max_len)[:, None] * torch.exp(torch.arange(0, dim, 2) * (-math.log(10000.0) / dim)) pe = torch.zeros(max_len, dim) pe[:, 0::2], pe[:, 1::2] = torch.sin(pos), torch.cos(pos) self.register_buffer("pe", pe, persistent=False) self.drop = nn.Dropout(dropout) self.tf = nn.Transformer(dim, heads, layers, layers, ffn, dropout, batch_first=True, norm_first=True) def embed(self, x): return self.drop(self.emb(x) * math.sqrt(self.emb.embedding_dim) + self.pe[: x.size(1)]) def encode(self, src): mask = src == PAD return self.tf.encoder(self.embed(src), src_key_padding_mask=mask), mask def decode(self, tgt_in, memory, src_mask): causal = nn.Transformer.generate_square_subsequent_mask(tgt_in.size(1), device=tgt_in.device, dtype=torch.bool) h = self.tf.decoder(self.embed(tgt_in), memory, tgt_mask=causal, tgt_is_causal=True, tgt_key_padding_mask=tgt_in == PAD, memory_key_padding_mask=src_mask) return F.linear(h, self.emb.weight) def forward(self, src, tgt_in): memory, src_mask = self.encode(src) return self.decode(tgt_in, memory, src_mask) @torch.no_grad() def beam_search(self, src, beam=5, max_len=None, alpha=0.6): """src: 1D-Tensor mit Token-IDs (ohne BOS/EOS). Gibt beste Hypothese als Liste von IDs zurück.""" max_len = max_len or min(int(len(src) * 1.5) + 10, self.cfg["max_len"]) memory, src_mask = self.encode(src[None]) hyps = torch.full((1, 1), BOS, dtype=torch.long, device=src.device) scores = torch.zeros(1, device=src.device) done = [] for step in range(max_len): n = hyps.size(0) logp = self.decode(hyps, memory.expand(n, -1, -1), src_mask.expand(n, -1))[:, -1].float().log_softmax(-1) cand = (scores[:, None] + logp).view(-1) top, idx = cand.topk(min(beam, cand.numel())) v = logp.size(-1) hyps = torch.cat([hyps[idx // v], (idx % v)[:, None]], 1) scores = top fin = hyps[:, -1] == EOS for h, s in zip(hyps[fin], scores[fin]): done.append((s.item() / ((5 + len(h) - 1) / 6) ** alpha, h[1:-1].tolist())) hyps, scores = hyps[~fin], scores[~fin] if len(done) >= beam or hyps.size(0) == 0: break if not done: done = [(s.item(), h[1:].tolist()) for h, s in zip(hyps, scores)] return max(done)[1] if __name__ == "__main__": m = GEC(32000) print(f"{sum(p.numel() for p in m.parameters()) / 1e6:.1f} Mio. Parameter") src = torch.randint(4, 32000, (2, 7)) assert m(src, src).shape == (2, 7, 32000) m.eval() assert len(m.beam_search(src[0], beam=3)) > 0 print("model.py OK")