File size: 5,708 Bytes
db5e0ee | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 | import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from config import *
class MultiHeadAttention(nn.Module):
def __init__(self, embed_dim, num_heads):
super().__init__()
assert embed_dim % num_heads == 0
self.embed_dim = embed_dim
self.num_heads = num_heads
self.head_dim = embed_dim // num_heads
self.qkv = nn.Linear(embed_dim, 3 * embed_dim)
self.out_proj = nn.Linear(embed_dim, embed_dim)
self.dropout = nn.Dropout(DROPOUT)
def forward(self, x, mask=None):
B, T, C = x.shape
# Q, K, V oluştur
qkv = self.qkv(x)
q, k, v = qkv.split(self.embed_dim, dim=2)
# Multi-head şekline getir: (B, T, C) -> (B, num_heads, T, head_dim)
q = q.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
k = k.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
v = v.view(B, T, self.num_heads, self.head_dim).transpose(1, 2)
# Attention hesapla
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.head_dim)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
attn_weights = F.softmax(scores, dim=-1)
attn_weights = self.dropout(attn_weights)
# Output
out = torch.matmul(attn_weights, v)
out = out.transpose(1, 2).contiguous().view(B, T, C)
out = self.out_proj(out)
return out
class FeedForward(nn.Module):
def __init__(self, embed_dim, hidden_dim):
super().__init__()
self.fc1 = nn.Linear(embed_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, embed_dim)
self.dropout = nn.Dropout(DROPOUT)
def forward(self, x):
x = self.fc1(x)
x = F.gelu(x)
x = self.dropout(x)
x = self.fc2(x)
return x
class TransformerBlock(nn.Module):
def __init__(self, embed_dim, num_heads, hidden_dim):
super().__init__()
self.ln1 = nn.LayerNorm(embed_dim)
self.attn = MultiHeadAttention(embed_dim, num_heads)
self.ln2 = nn.LayerNorm(embed_dim)
self.ff = FeedForward(embed_dim, hidden_dim)
self.dropout = nn.Dropout(DROPOUT)
def forward(self, x, mask=None):
# Self-attention with residual
x = x + self.dropout(self.attn(self.ln1(x), mask))
# Feed-forward with residual
x = x + self.dropout(self.ff(self.ln2(x)))
return x
class PegeModel(nn.Module):
def __init__(self, vocab_size):
super().__init__()
self.token_embedding = nn.Embedding(vocab_size, EMBED_DIM)
self.position_embedding = nn.Embedding(MAX_SEQ_LEN, EMBED_DIM)
self.blocks = nn.ModuleList([
TransformerBlock(EMBED_DIM, NUM_HEADS, HIDDEN_DIM)
for _ in range(NUM_LAYERS)
])
self.ln_f = nn.LayerNorm(EMBED_DIM)
self.head = nn.Linear(EMBED_DIM, vocab_size, bias=False)
# Weight tying
self.token_embedding.weight = self.head.weight
self.dropout = nn.Dropout(DROPOUT)
self.apply(self._init_weights)
def _init_weights(self, module):
if isinstance(module, nn.Linear):
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
torch.nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
elif isinstance(module, nn.LayerNorm):
torch.nn.init.zeros_(module.bias)
torch.nn.init.ones_(module.weight)
def forward(self, input_ids, targets=None):
B, T = input_ids.shape
# Embeddings
tok_emb = self.token_embedding(input_ids)
pos_emb = self.position_embedding(torch.arange(T, device=input_ids.device))
x = self.dropout(tok_emb + pos_emb)
# Causal mask
mask = torch.tril(torch.ones(T, T, device=input_ids.device)).view(1, 1, T, T)
# Transformer blocks
for block in self.blocks:
x = block(x, mask)
x = self.ln_f(x)
logits = self.head(x)
loss = None
if targets is not None:
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
return logits, loss
def generate(self, input_ids, max_new_tokens=100, temperature=1.0, top_k=40,
repetition_penalty=1.3, stop_tokens=None):
self.eval()
stop_tokens = set(stop_tokens or [])
with torch.no_grad():
for _ in range(max_new_tokens):
input_ids_cond = input_ids if input_ids.size(1) <= MAX_SEQ_LEN else input_ids[:, -MAX_SEQ_LEN:]
logits, _ = self(input_ids_cond)
logits = logits[:, -1, :]
if repetition_penalty != 1.0:
for token_id in set(input_ids[0].tolist()):
if logits[0, token_id] > 0:
logits[0, token_id] /= repetition_penalty
else:
logits[0, token_id] *= repetition_penalty
logits = logits / temperature
if top_k is not None:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = float('-inf')
probs = F.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
input_ids = torch.cat([input_ids, next_token], dim=1)
# Stop token gelince dur
if next_token.item() in stop_tokens:
break
return input_ids |