cbspace commited on
Commit
da39f22
·
1 Parent(s): 0233b91

Fix inference

Browse files
Files changed (2) hide show
  1. app.py +1 -1
  2. model.py +3 -2
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: