HORST / model.py
Bayernator's picture
Pixel-Design und Trainingsstufe 2 (KV-Cache-Dekodierer)
3345e65 verified
Raw History Blame Contribute Delete
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]
@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")