HORST / model.py
Bayernator's picture
HORST: eigenes Korrekturmodell auf ZeroGPU
6424eee verified
Raw History Blame Contribute Delete
3.48 kB
"""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")