cbspace commited on
Commit ·
da39f22
1
Parent(s): 0233b91
Fix inference
Browse files
app.py
CHANGED
|
@@ -21,7 +21,7 @@ dropout = 0.4
|
|
| 21 |
tokenizer = tiktoken.encoding_for_model('gpt2')
|
| 22 |
model_path = hf_hub_download('cbspace/gpt', 'model.safetensors')
|
| 23 |
state_dict = load_file(model_path)
|
| 24 |
-
model = GPTModel(n_layers, n_heads, embed_dim, ffn_dim, n_vocab, max_seq_len, dropout)
|
| 25 |
model.load_state_dict(state_dict, strict=False)
|
| 26 |
model.eval()
|
| 27 |
|
|
|
|
| 21 |
tokenizer = tiktoken.encoding_for_model('gpt2')
|
| 22 |
model_path = hf_hub_download('cbspace/gpt', 'model.safetensors')
|
| 23 |
state_dict = load_file(model_path)
|
| 24 |
+
model = GPTModel(device, n_layers, n_heads, embed_dim, ffn_dim, n_vocab, max_seq_len, dropout)
|
| 25 |
model.load_state_dict(state_dict, strict=False)
|
| 26 |
model.eval()
|
| 27 |
|
model.py
CHANGED
|
@@ -24,9 +24,10 @@ class TransformerBlock(nn.Module):
|
|
| 24 |
|
| 25 |
|
| 26 |
class GPTModel(nn.Module):
|
| 27 |
-
def __init__(self, n_layers, n_heads, embed_dim, ffn_dim, n_vocab, max_seq_len, dropout):
|
| 28 |
super().__init__()
|
| 29 |
self.max_seq_len = max_seq_len
|
|
|
|
| 30 |
self.embedding = nn.Embedding(n_vocab, embed_dim)
|
| 31 |
nn.init.normal_(self.embedding.weight, mean=0.0, std=0.02)
|
| 32 |
self.positional_embedding = nn.Embedding(max_seq_len, embed_dim)
|
|
@@ -55,7 +56,7 @@ class GPTModel(nn.Module):
|
|
| 55 |
context_list = [i for i in input_ctx]
|
| 56 |
with torch.no_grad():
|
| 57 |
while len(context_list) < max_length:
|
| 58 |
-
context = torch.tensor(context_list, dtype=torch.long, device=device).unsqueeze(0)
|
| 59 |
logits = self(context)[0,-1,:]
|
| 60 |
|
| 61 |
if greedy:
|
|
|
|
| 24 |
|
| 25 |
|
| 26 |
class GPTModel(nn.Module):
|
| 27 |
+
def __init__(self, device, n_layers, n_heads, embed_dim, ffn_dim, n_vocab, max_seq_len, dropout):
|
| 28 |
super().__init__()
|
| 29 |
self.max_seq_len = max_seq_len
|
| 30 |
+
self.device = device
|
| 31 |
self.embedding = nn.Embedding(n_vocab, embed_dim)
|
| 32 |
nn.init.normal_(self.embedding.weight, mean=0.0, std=0.02)
|
| 33 |
self.positional_embedding = nn.Embedding(max_seq_len, embed_dim)
|
|
|
|
| 56 |
context_list = [i for i in input_ctx]
|
| 57 |
with torch.no_grad():
|
| 58 |
while len(context_list) < max_length:
|
| 59 |
+
context = torch.tensor(context_list, dtype=torch.long, device=self.device).unsqueeze(0)
|
| 60 |
logits = self(context)[0,-1,:]
|
| 61 |
|
| 62 |
if greedy:
|