Spaces:
Running on Zero
Running on Zero
Download model.py from Bayernator/HORST: direct link, hf CLI and curl.
- Browser
- Download file 3.48 kB
-
https://huggingface.co/spaces/Bayernator/HORST/resolve/main/model.py
- Command line
-
hf download hf://spaces/Bayernator/HORST/model.py
-
curl -L -o model.py https://huggingface.co/spaces/Bayernator/HORST/resolve/main/model.py
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) | |
| 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") | |