Views
No views yet
| Attribute | Value |
|---|---|
| Model Name | Instinct-1B |
| Organization | AutonomousX |
| Parameters | 1B |
| Vocabulary Size | 50,304 |
| Dataset | PILE |
| Tokenizer | Pythia Tokenizer / BPE |
| Tokens Seen | 85B |
| Training Hardware | TPU v4-8 |
| Optimizer | AdamW |
| Architecture | v-4 (128) pmap |
| Positional Embeddings | No RoPE |
1graph TD
2 %% Dataset and Preparation
3 Data["Dataset: PILE\nRaw Text Data"]
4 Tokenizer["BPE Tokenizer\nVocabulary Construction"]
5 TokenizedData["Tokenized Data\nReady for Training"]
6
7 %% Model Architecture
8 Model["v-4 (128) pmap\nTransformer Decoder\n1B Parameters"]
9 RoPE["No RoPE Positional Embeddings"]
10
11 %% Training Pipeline
12 Optimizer["Optimizer: AdamW"]
13 ForwardPass["Forward Pass\nCompute Loss"]
14 BackwardPass["Backward Pass\nCompute Gradients"]
15 Update["Parameter Update"]
16
17 %% Logging and Checkpoints
18 Checkpoints["Model Checkpoints\nSaved up to 85B tokens"]
19 Logs["Training Logs\nLoss & Perplexity"]
20
21 %% Connections
22 Data --> Tokenizer
23 Tokenizer --> TokenizedData
24 TokenizedData --> ForwardPass
25
26 Model --> ForwardPass
27 RoPE -.-> Model
28
29 ForwardPass --> BackwardPass
30 BackwardPass --> Optimizer
31 Optimizer --> Update
32 Update --> Model
33
34 Update --> Checkpoints
35 Update --> Logstraining_log.txt and val_perplexity.txt. Below is the visualization of the training progress:
N_LAYERS, D_MODEL, N_HEADS, D_HEAD, D_FF, etc., according to this specific model's architecture (1B) to run inference successfully.1#please be patient It may take 20 mins to run the model
2# Install huggingface_hub if not installed
3!pip install -q huggingface_hub
4
5from huggingface_hub import snapshot_download
6
7repo_id = "autonomousX/Instinct-1B"
8
9# Download entire repository
10local_path = snapshot_download(
11 repo_id=repo_id,
12 repo_type="model",
13 local_dir="TPU_1B",
14 local_dir_use_symlinks=False
15)
16
17print("Download complete!")
18print("Saved to:", local_path)
19
20# =========================
21# FAST 1B INFERENCE CELL
22# =========================
23
24import os
25import math
26import jax
27import jax.numpy as jnp
28from flax import linen as nn
29from flax.training import train_state, checkpoints
30import optax
31from transformers import AutoTokenizer
32
33# ---------------- CONFIG ----------------
34SEQ_LEN = 1024
35VOCAB_SIZE = 50304
36
37# NOTE: Adjust these parameters for your specific 1B architecture!
38N_LAYERS = 32
39D_MODEL = 1024
40N_HEADS = 16
41D_HEAD = 64
42D_FF = 4096
43ROTARY_PCT = 0.25
44
45CKPT_PATH = os.path.abspath("TPU_1B/checkpoint_0")
46
47# ---------------- RoPE ----------------
48def build_rope_cache(seq_len, head_dim, rotary_pct):
49 dim = int(head_dim * rotary_pct)
50 freqs = 1.0 / (10000 ** (jnp.arange(0, dim, 2) / dim))
51 pos = jnp.arange(seq_len)
52 angles = jnp.einsum("i,j->ij", pos, freqs)
53 return jnp.sin(angles), jnp.cos(angles)
54
55ROPE_SIN, ROPE_COS = build_rope_cache(SEQ_LEN, D_HEAD, ROTARY_PCT)
56
57def apply_rope(q, k):
58 dim = int(D_HEAD * ROTARY_PCT)
59 T = q.shape[1]
60
61 sin = ROPE_SIN[:T][None, :, None, :]
62 cos = ROPE_COS[:T][None, :, None, :]
63
64 q_rot, q_pass = q[..., :dim], q[..., dim:]
65 k_rot, k_pass = k[..., :dim], k[..., dim:]
66
67 q1, q2 = q_rot[..., ::2], q_rot[..., 1::2]
68 k1, k2 = k_rot[..., ::2], k_rot[..., 1::2]
69
70 q_rot = jnp.concatenate(
71 [q1 * cos - q2 * sin,
72 q1 * sin + q2 * cos],
73 axis=-1
74 )
75
76 k_rot = jnp.concatenate(
77 [k1 * cos - k2 * sin,
78 k1 * sin + k2 * cos],
79 axis=-1
80 )
81
82 return (
83 jnp.concatenate([q_rot, q_pass], axis=-1),
84 jnp.concatenate([k_rot, k_pass], axis=-1),
85 )
86
87# ---------------- MODEL ----------------
88class RMSNorm(nn.Module):
89 dim: int
90 eps: float = 1e-6
91 @nn.compact
92 def __call__(self, x):
93 scale = self.param("scale", nn.initializers.ones, (self.dim,))
94 norm = jnp.sqrt(jnp.mean(x**2, axis=-1, keepdims=True) + self.eps)
95 return x * (scale / norm)
96
97class Attention(nn.Module):
98 @nn.compact
99 def __call__(self, x, mask):
100 B, T, C = x.shape
101 qkv = nn.Dense(3 * C, use_bias=False, dtype=jnp.bfloat16)(x)
102 qkv = qkv.reshape(B, T, 3, N_HEADS, D_HEAD)
103
104 q = qkv[:, :, 0]
105 k = qkv[:, :, 1]
106 v = qkv[:, :, 2]
107
108 q, k = apply_rope(q, k)
109
110 att = jnp.einsum("bthd,bshd->bhts", q, k)
111 att = att / math.sqrt(D_HEAD)
112
113 mask = mask.astype(jnp.float32)
114 mask = (1.0 - mask) * -1e10
115 att = att + mask
116
117 att = nn.softmax(att.astype(jnp.float32), axis=-1)
118 att = att.astype(jnp.bfloat16)
119
120 out = jnp.einsum("bhts,bshd->bthd", att, v)
121 out = out.reshape(B, T, C)
122
123 return nn.Dense(C, use_bias=False, dtype=jnp.bfloat16)(out)
124
125class Block(nn.Module):
126 @nn.compact
127 def __call__(self, x, mask):
128 h = RMSNorm(D_MODEL)(x)
129 h = Attention()(h, mask)
130 x = x + h
131
132 h = RMSNorm(D_MODEL)(x)
133 h = nn.Dense(D_FF, dtype=jnp.bfloat16)(h)
134 h = nn.gelu(h)
135 h = nn.Dense(D_MODEL, dtype=jnp.bfloat16)(h)
136
137 return x + h
138
139class GPT(nn.Module):
140 @nn.compact
141 def __call__(self, input_ids):
142 batch, seq_len = input_ids.shape
143 mask = nn.attention.make_causal_mask(
144 jnp.ones((batch, seq_len), dtype=jnp.bool_)
145 )
146
147 x = nn.Embed(
148 VOCAB_SIZE,
149 D_MODEL,
150 embedding_init=nn.initializers.normal(0.02),
151 dtype=jnp.bfloat16,
152 )(input_ids)
153
154 RematBlock = nn.remat(Block)
155
156 for _ in range(N_LAYERS):
157 x = RematBlock()(x, mask)
158
159 x = RMSNorm(D_MODEL)(x)
160
161 return nn.Dense(
162 VOCAB_SIZE,
163 use_bias=False,
164 dtype=jnp.bfloat16
165 )(x)
166# ---------------- LOAD CHECKPOINT ----------------
167def create_state():
168 model = GPT()
169 rng = jax.random.PRNGKey(0)
170 params = model.init(rng, jnp.ones((1, SEQ_LEN), dtype=jnp.int32))
171 return train_state.TrainState.create(
172 apply_fn=model.apply,
173 params=params,
174 tx=optax.adamw(1e-4),
175 )
176
177state = create_state()
178state = checkpoints.restore_checkpoint(CKPT_PATH, state)
179params = state.params
180model = GPT()
181
182print("Checkpoint loaded.")
183
184@jax.jit
185def forward(params, input_ids):
186 return model.apply(params, input_ids)
187
188import jax.random as random
189
190def generate(params, input_ids, max_new_tokens=30, temperature=0.9, top_k=40):
191 rng = random.PRNGKey(0)
192
193 for _ in range(max_new_tokens):
194
195 logits = model.apply(params, input_ids)
196 logits = logits[:, -1, :]
197 logits = logits.astype(jnp.float32)
198
199 logits = logits / temperature
200
201 top_k_logits, top_k_indices = jax.lax.top_k(logits, top_k)
202 probs = jax.nn.softmax(top_k_logits, axis=-1)
203
204 rng, subkey = random.split(rng)
205 next_token_idx = random.categorical(subkey, jnp.log(probs))
206
207 next_token = jnp.take_along_axis(
208 top_k_indices,
209 next_token_idx[:, None],
210 axis=-1
211 )
212
213 input_ids = jnp.concatenate([input_ids, next_token], axis=1)
214
215 return input_ids
216# ---------------- RUN ----------------
217tokenizer = AutoTokenizer.from_pretrained("autonomousX/Instinct-1B")
218
219prompt = "I am John,"
220tokens = tokenizer(prompt, return_tensors="np")
221input_ids = jnp.array(tokens["input_ids"], dtype=jnp.int32)
222
223output_ids = generate(params, input_ids, 200)
224
225print("\n=== GENERATED TEXT ===\n")
226print(tokenizer.decode(output_ids[0].tolist()))