"""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] @torch.no_grad() 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) @torch.no_grad() 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 @torch.no_grad() 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] @torch.no_grad() 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")