Spaces:
Running on Zero
Running on Zero
Download model.py from Bayernator/HORST: direct link, hf CLI and curl.
- Browser
- Download file 11.3 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
11.3 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) | |
| # ---- KV-Cache-Dekodieren (nur eval/no_grad; Dropout entfällt) ---- | |
| def embed_step(self, tok, pos): | |
| """tok: (B,) Long, pos: int -> (B,1,dim), wie embed() an Position pos.""" | |
| return self.drop(self.emb(tok) * math.sqrt(self.emb.embedding_dim) + self.pe[pos])[:, None] | |
| def init_cache(self, memory, src_mask): | |
| """Kreuz-Attention-K/V je Decoder-Layer einmal berechnen; Self-K/V (None = leer) wachsen in decode_step.""" | |
| B, S, d = memory.shape | |
| H = self.cfg["heads"] | |
| cross = [] | |
| for layer in self.tf.decoder.layers: | |
| mha = layer.multihead_attn | |
| k, v = F.linear(memory, mha.in_proj_weight[d:], mha.in_proj_bias[d:]).chunk(2, -1) | |
| cross.append((k.reshape(B, S, H, d // H).transpose(1, 2), v.reshape(B, S, H, d // H).transpose(1, 2))) | |
| return dict(cross=cross, self=[(None, None)] * len(cross), kpad=src_mask.new_zeros(B, 0), | |
| pos=0, src_mask=src_mask) | |
| def decode_step(self, tok, cache): | |
| """Ein Dekodierschritt: tok (B,) -> Logits (B,V) float. Äquivalent zu decode(...)[:, -1].""" | |
| B, d, H = tok.size(0), self.cfg["dim"], self.cfg["heads"] | |
| x = self.embed_step(tok, cache["pos"]) | |
| # PAD im Präfix wird wie in decode() (tgt_key_padding_mask) nicht angeschaut; SDPA: True = darf attendieren | |
| cache["kpad"] = torch.cat([cache["kpad"], (tok == PAD)[:, None]], 1) | |
| self_ok = ~cache["kpad"][:, None, None, :] | |
| cross_ok = ~cache["src_mask"][:, None, None, :] | |
| for i, layer in enumerate(self.tf.decoder.layers): | |
| sa, ca = layer.self_attn, layer.multihead_attn | |
| q, k, v = F.linear(layer.norm1(x), sa.in_proj_weight, sa.in_proj_bias).chunk(3, -1) | |
| q, k, v = (t.reshape(B, 1, H, d // H).transpose(1, 2) for t in (q, k, v)) | |
| sk, sv = cache["self"][i] | |
| if sk is not None: | |
| k, v = torch.cat([sk, k], 2), torch.cat([sv, v], 2) | |
| cache["self"][i] = (k, v) | |
| a = F.scaled_dot_product_attention(q, k, v, attn_mask=self_ok) | |
| x = x + sa.out_proj(a.transpose(1, 2).reshape(B, 1, d)) | |
| q = F.linear(layer.norm2(x), ca.in_proj_weight[:d], ca.in_proj_bias[:d]).reshape(B, 1, H, d // H).transpose(1, 2) | |
| a = F.scaled_dot_product_attention(q, *cache["cross"][i], attn_mask=cross_ok) | |
| x = x + ca.out_proj(a.transpose(1, 2).reshape(B, 1, d)) | |
| x = x + layer.linear2(layer.activation(layer.linear1(layer.norm3(x)))) | |
| cache["pos"] += 1 | |
| return F.linear(self.tf.decoder.norm(x)[:, 0], self.emb.weight).float() | |
| def reorder_cache(self, cache, idx, cross=True): | |
| """Zeilen nach idx umordnen (Beam-Indizes). cross=False spart die Kreuz-K/V, wenn idx Sätze nicht mischt.""" | |
| new = dict(cache) | |
| new["self"] = [(None, None) if k is None else (k.index_select(0, idx), v.index_select(0, idx)) | |
| for k, v in cache["self"]] | |
| new["kpad"] = cache["kpad"].index_select(0, idx) | |
| if cross: | |
| new["cross"] = [(k.index_select(0, idx), v.index_select(0, idx)) for k, v in cache["cross"]] | |
| new["src_mask"] = cache["src_mask"].index_select(0, idx) | |
| return new | |
| def generate(self, srcs, beam=1, max_len=None, alpha=0.6, sample=False, top_k=0, temperature=1.0, | |
| copy_bias=0.0, seed=None): | |
| """srcs: Liste 1D-Long-Tensoren (je inkl. EOS). Greedy/Beam/Sampling mit KV-Cache, batchweise. | |
| Gibt je Quelle die IDs ohne BOS/EOS zurück. copy_bias addiert auf die Logits aller Quell-Token.""" | |
| assert not (sample and beam > 1), "Sampling geht nur mit beam=1" | |
| if not srcs: | |
| return [] | |
| B, V, K = len(srcs), self.cfg["vocab"], beam | |
| src = nn.utils.rnn.pad_sequence(srcs, batch_first=True, padding_value=PAD) | |
| dev = src.device | |
| max_len = max_len or min(int(max(len(s) for s in srcs) * 1.5) + 10, self.cfg["max_len"]) | |
| memory, src_mask = self.encode(src) | |
| cache = self.init_cache(memory, src_mask) | |
| bias = None | |
| if copy_bias: | |
| bias = torch.zeros(B, V, device=dev).scatter_(1, src, copy_bias) | |
| bias[:, [PAD, BOS, EOS]] = 0 | |
| if K == 1: | |
| gen = None | |
| if sample: | |
| gen = torch.Generator(device=dev) | |
| gen.manual_seed(seed) if seed is not None else gen.seed() | |
| tok = torch.full((B,), BOS, dtype=torch.long, device=dev) | |
| fin = torch.zeros(B, dtype=torch.bool, device=dev) | |
| steps = [] | |
| for _ in range(max_len): | |
| logits = self.decode_step(tok, cache) | |
| if bias is not None: | |
| logits = logits + bias | |
| if sample: | |
| logits = logits / max(temperature, 1e-4) | |
| if top_k > 0: | |
| logits = logits.masked_fill(logits < logits.topk(min(top_k, V)).values[:, -1:], -float("inf")) | |
| tok = torch.multinomial(logits.softmax(-1), 1, generator=gen)[:, 0] | |
| else: | |
| tok = logits.argmax(-1) | |
| tok = tok.masked_fill(fin, EOS) # fertige Zeilen laufen mit EOS weiter, Ausgabe wird abgeschnitten | |
| steps.append(tok) | |
| fin = fin | (tok == EOS) | |
| if fin.all(): | |
| break | |
| out = [] | |
| for row in torch.stack(steps, 1).tolist(): | |
| out.append(row[:row.index(EOS)] if EOS in row else row) | |
| return out | |
| # Beam-Search: B*K Zeilen, Zeile b*K+j. Tote Zeilen (fertig/unbelegt) haben Score -inf, damit gleiche | |
| # Auswahl wie im alten Code, wo fertige Hypothesen aus der Liste entfernt werden. | |
| rep = torch.arange(B, device=dev).repeat_interleave(K) | |
| cache = self.reorder_cache(cache, rep) | |
| if bias is not None: | |
| bias = bias[rep] | |
| hyps = torch.full((B * K, 1), BOS, dtype=torch.long, device=dev) | |
| scores = torch.full((B, K), -float("inf"), device=dev) | |
| scores[:, 0] = 0 | |
| base = torch.arange(B, device=dev)[:, None] * K | |
| done = [[] for _ in range(B)] | |
| fin = [False] * B | |
| for _ in range(max_len): | |
| logits = self.decode_step(hyps[:, -1], cache) | |
| if bias is not None: | |
| logits = logits + bias | |
| cand = (scores.view(-1, 1) + logits.log_softmax(-1)).view(B, K * V) | |
| top, idx = cand.topk(K) | |
| par = (base + idx // V).view(-1) | |
| hyps = torch.cat([hyps[par], (idx % V).view(-1, 1)], 1) | |
| cache = self.reorder_cache(cache, par, cross=False) | |
| ends = ((idx % V) == EOS) & torch.isfinite(top) | |
| for b, j in ends.nonzero().tolist(): | |
| if not fin[b]: | |
| done[b].append((top[b, j].item() / ((5 + hyps.size(1) - 1) / 6) ** alpha, hyps[b * K + j, 1:-1].tolist())) | |
| scores = top.masked_fill(ends, -float("inf")) | |
| alive = torch.isfinite(scores).any(1).tolist() | |
| fin = [f or len(done[b]) >= K or not alive[b] for b, f in enumerate(fin)] | |
| if all(fin): | |
| break | |
| out = [] | |
| for b in range(B): | |
| if not done[b]: # kein EOS bis max_len: unnormalisiert die besten lebenden Hypothesen | |
| done[b] = [(scores[b, j].item(), hyps[b * K + j, 1:].tolist()) | |
| for j in range(K) if torch.isfinite(scores[b, j])] | |
| out.append(max(done[b])[1]) | |
| return out | |
| def beam_search(self, src, beam=5, max_len=None, alpha=0.6, copy_bias=0.0): | |
| """src: 1D-Tensor mit Token-IDs (inkl. EOS, ohne BOS). Gibt beste Hypothese als Liste von IDs zurück.""" | |
| return self.generate([src], beam=beam, max_len=max_len, alpha=alpha, copy_bias=copy_bias)[0] | |
| def _beam_search_slow(self, src, beam=5, max_len=None, alpha=0.6): | |
| """Alter Code ohne KV-Cache, nur als Referenz für Tests.""" | |
| 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") | |