cbspace commited on
Commit ·
83b97e5
1
Parent(s): 3adad86
Updated for new model
Browse files- model.py +7 -2
- requirements.txt +2 -1
model.py
CHANGED
|
@@ -1,5 +1,7 @@
|
|
| 1 |
import torch
|
| 2 |
from torch import nn
|
|
|
|
|
|
|
| 3 |
|
| 4 |
class TransformerBlock(nn.Module):
|
| 5 |
def __init__(self, device, n_heads, embed_dim, ffn_dim, dropout):
|
|
@@ -28,7 +30,7 @@ class GPTModel(nn.Module):
|
|
| 28 |
super().__init__()
|
| 29 |
self.max_seq_len = max_seq_len
|
| 30 |
self.device = device
|
| 31 |
-
self.embedding = nn.
|
| 32 |
nn.init.normal_(self.embedding.weight, mean=0.0, std=0.02)
|
| 33 |
self.positional_embedding = nn.Embedding(max_seq_len, embed_dim)
|
| 34 |
|
|
@@ -42,8 +44,11 @@ class GPTModel(nn.Module):
|
|
| 42 |
input_embed = self.embedding(input_tokens)
|
| 43 |
positions = torch.arange(0, input_tokens.size(1), device=input_tokens.device).unsqueeze(0)
|
| 44 |
input_embed = input_embed + self.positional_embedding(positions)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
|
| 46 |
-
x = self.transformer(input_embed)
|
| 47 |
x = self.layer_norm(x)
|
| 48 |
x = self.output_projection(x)
|
| 49 |
return x
|
|
|
|
| 1 |
import torch
|
| 2 |
from torch import nn
|
| 3 |
+
import bitsandbytes as bnb
|
| 4 |
+
from torch.utils.checkpoint import checkpoint
|
| 5 |
|
| 6 |
class TransformerBlock(nn.Module):
|
| 7 |
def __init__(self, device, n_heads, embed_dim, ffn_dim, dropout):
|
|
|
|
| 30 |
super().__init__()
|
| 31 |
self.max_seq_len = max_seq_len
|
| 32 |
self.device = device
|
| 33 |
+
self.embedding = bnb.nn.StableEmbedding(n_vocab, embed_dim)
|
| 34 |
nn.init.normal_(self.embedding.weight, mean=0.0, std=0.02)
|
| 35 |
self.positional_embedding = nn.Embedding(max_seq_len, embed_dim)
|
| 36 |
|
|
|
|
| 44 |
input_embed = self.embedding(input_tokens)
|
| 45 |
positions = torch.arange(0, input_tokens.size(1), device=input_tokens.device).unsqueeze(0)
|
| 46 |
input_embed = input_embed + self.positional_embedding(positions)
|
| 47 |
+
x = input_embed
|
| 48 |
+
|
| 49 |
+
for block in self.transformer_blocks:
|
| 50 |
+
x = checkpoint(block, x, use_reentrant=False)
|
| 51 |
|
|
|
|
| 52 |
x = self.layer_norm(x)
|
| 53 |
x = self.output_projection(x)
|
| 54 |
return x
|
requirements.txt
CHANGED
|
@@ -1,4 +1,5 @@
|
|
| 1 |
torch
|
| 2 |
tiktoken
|
| 3 |
huggingface_hub
|
| 4 |
-
safetensors
|
|
|
|
|
|
| 1 |
torch
|
| 2 |
tiktoken
|
| 3 |
huggingface_hub
|
| 4 |
+
safetensors
|
| 5 |
+
bitsandbytes
|