Spaces:
Runtime error
Runtime error
File size: 5,744 Bytes
efebfa9 c75b153 01f78d9 efebfa9 5235f23 efebfa9 e9f4769 5235f23 efebfa9 e4e4dae efebfa9 9dca230 efebfa9 2e92c3a efebfa9 1917f6e efebfa9 2e92c3a efebfa9 e22ec83 1d6d9b7 efebfa9 d3a0713 efebfa9 0832c01 85ef85f c75b153 efebfa9 0832c01 c75b153 efebfa9 c75b153 efebfa9 | 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 | import torch
import gradio as gr
import torch
import torch.nn as nn
from torch.nn import functional as F
embed_size = int(384*1.5)
block_size = 256
dropout = 0.2
n_layer = 9
n_head = 12
device = "cuda" if torch.cuda.is_available() else "cpu"
vocab = [' ', '!', '"', '#', '$', '%', '&', "'", '(', ')', '*', '+', ',', '-', '.', '/', '0', '1', '2', '3', '4', '5', '6', '7', '8', '9', ':', ';', '<', '=', '>', '?', '@', 'A', 'B', 'C', 'D', 'E', 'F', 'G', 'H', 'I', 'J', 'K', 'L', 'M', 'N', 'O', 'P', 'Q', 'R', 'S', 'T', 'U', 'V', 'W', 'X', 'Y', 'Z', '[', '\\', ']', '^', '`', 'a', 'b', 'c', 'd', 'e', 'f', 'g', 'h', 'i', 'j', 'k', 'l', 'm', 'n', 'o', 'p', 'q', 'r', 's', 't', 'u', 'v', 'w', 'x', 'y', 'z', '{', '|', '}', '~', '\x80', '\x81', '\x82', '\x83', '\x84', '\x85', '\x86', '\x87', '\x88', '\x89', '\x8a', '\x8b', '\x8c', '\x8d', '\x8e', '\x8f', '\x90', '\x91', '\x92', '\x93', '\x94', '\x95', '\x96', '\x97', '\x98', '\x99', '\x9a', '\x9b', '\x9c', '\x9d', '\x9e', '\x9f', '\xa0', '¡', '¢', '£', '¤', '¥', '¦', '§', '¨', '©', 'ª', '«', '¬', '\xad', '®', '¯', '°', '±', '²', '³', '´', 'µ', '¶', '·', '¸', '¹', 'º', '»', '¼', '½', '¾', '¿', 'Â', 'Ã', 'Ä', 'Å', 'Æ', 'Ç', 'È', 'É', 'Ê', 'Ë', 'Ì', 'Í', 'Î', 'Ï', 'Ð', 'Ñ', 'Ö', '×', 'Ø', 'Ù', 'Ú', 'Û', 'à', 'á', 'ã', 'ä', 'å', 'æ', 'ç', 'è', 'é', 'ê', 'ï', 'ð']
vocab_size = len(vocab)
encode = lambda x: [vocab.index(i) for i in x]
decode = lambda x: ''.join([vocab[i] for i in x])
class trans_block(nn.Module):
def __init__(self,embed_size,heads):
super().__init__()
head_size = embed_size // heads
self.attention = Heads(heads,head_size)
self.ff_layer = FF_Layer(embed_size)
self.lnorm1 = nn.LayerNorm(embed_size)
self.lnorm2 = nn.LayerNorm(embed_size)
def forward(self,x):
x = x + self.attention(self.lnorm1(x))
x = x + self.ff_layer(self.lnorm2(x))
return x
class Head(nn.Module):
def __init__(self,headsize):
super().__init__()
self.key = nn.Linear(embed_size,headsize,bias=False)
self.query = nn.Linear(embed_size,headsize,bias=False)
self.value = nn.Linear(embed_size,headsize,bias=False)
self.register_buffer('tril',torch.tril(torch.ones(block_size,block_size)))
self.dropout = nn.Dropout(dropout)
def forward(self,x):
Batches, Time, Channels = x.shape
k = self.key(x)
q = self.query(x)
wei = q @ k.transpose(-2,-1) * Channels**-0.5
wei = wei.masked_fill(self.tril[:Time,:Time] == 0,float('-inf'))
wei = F.softmax(wei,dim=-1)
wei = self.dropout(wei)
v = self.value(x)
out = wei @ v
return out
class Heads(nn.Module):
def __init__(self,n_head,head_size):
super().__init__()
self.heads = nn.ModuleList([Head(head_size) for i in range(n_head)])
self.projection = nn.Linear(embed_size, embed_size)
self.dropout = nn.Dropout(dropout)
def forward(self,x):
out = torch.cat([head(x) for head in self.heads],dim=-1)
out = self.dropout(self.projection(out))
return out
class FF_Layer(nn.Module):
def __init__(self,embed_size):
super().__init__()
self.net = nn.Sequential(
nn.Linear(embed_size,4*embed_size),
nn.ReLU(),
nn.Linear(4*embed_size,embed_size),
nn.Dropout(dropout)
)
def forward(self,x):
return self.net(x)
class BigramLM(nn.Module):
def __init__(self):
super().__init__()
self.embedding_table = nn.Embedding(vocab_size,embed_size)
self.position_embedding_table = nn.Embedding(block_size,embed_size)
self.lm_head = nn.Linear(embed_size,vocab_size)
self.blocks = nn.Sequential(*[trans_block(embed_size,heads = n_head) for _ in range(n_layer)])
self.ln_f = nn.LayerNorm(embed_size)
def forward(self,idx,targets=None):
Branch,Time = idx.shape
token_embed = self.embedding_table(idx)
position_embed = self.position_embedding_table(torch.arange(Time,device=device))
added = token_embed + position_embed
added = self.blocks(added)
added = self.ln_f(added)
logits = self.lm_head(added)
if targets is None:
loss = None
else:
Batch, Time, Channel = logits.shape
logits = logits.view(Batch*Time,Channel)
targets = targets.view(Batch*Time)
loss = F.cross_entropy(logits,targets)
return logits,loss
def generate(self, idx, max_tokens):
for i in range(max_tokens):
idx_condition = idx[:, -block_size:]
logits, loss = self(idx_condition)
logits = logits[:, -1, :]
probs = F.softmax(logits, dim=-1)
idx_next = torch.multinomial(probs, num_samples=1)
idx = torch.cat((idx, idx_next), dim=1)
return idx
print("loading")
model2 = BigramLM()
model2.load_state_dict(torch.load("coder(1).txt",map_location = torch.device(device)))
def generate_text(contextc, tokens):
print("generating")
context = torch.tensor([encode(contextc)])
a = model2.generate(context,max_tokens = tokens)
return a
iface = gr.Interface(
fn=model2.generate,
inputs=[
gr.Textbox(label="Prompt", placeholder="Hullo, said the mysterious man standing on the door"),
gr.Slider(minimum=1, maximum=2000, step=1, label="Number of characters to generate", value=100)
],
outputs=gr.Textbox(label="Generated Text"),
title="HoLLMes",
description="A janky LLM trained on Detective Novels."
)
# Launch the interface
if __name__ == "__main__":
iface.launch()
|