thefinalboss commited on
Commit
036ee64
·
verified ·
1 Parent(s): 218f72f

Upload train_1b_final.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. train_1b_final.py +604 -0
train_1b_final.py ADDED
@@ -0,0 +1,604 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ CogNet-1B Training Script — RTX 5090 32GB VRAM
4
+ =================================================
5
+ Full AdamW optimizer, BF16 mixed precision, gradient checkpointing.
6
+ Architecture: ~1.02B parameters
7
+ - hidden_dim=2048, 16 blocks, 8 channels
8
+ - channel_dim=384, ff_dim=8192
9
+ - working_slots=128, episodic_slots=256, semantic_slots=512
10
+ """
11
+
12
+ import os
13
+ import sys
14
+ import time
15
+ import math
16
+ import json
17
+ import logging
18
+ import argparse
19
+ from pathlib import Path
20
+
21
+ import torch
22
+ import torch.nn as nn
23
+ import torch.nn.functional as F
24
+ from torch.utils.data import Dataset, DataLoader
25
+
26
+ # ──────────────────────────────────────────────────────────
27
+ # Logging
28
+ # ──────────────────────────────────────────────────────────
29
+ logging.basicConfig(
30
+ level=logging.INFO,
31
+ format='%(asctime)s [%(levelname)s] %(message)s',
32
+ handlers=[
33
+ logging.StreamHandler(sys.stdout),
34
+ logging.FileHandler('/root/CogNet/train_1b.log', mode='a'),
35
+ ],
36
+ )
37
+ log = logging.getLogger('cognet1b')
38
+
39
+ # ──────────────────────────────────────────────────────────
40
+ # Character Tokenizer
41
+ # ──────────────────────────────────────────────────────────
42
+ class CharTokenizer:
43
+ """Character-level tokenizer with 136 vocab (ASCII + French accents + special tokens)."""
44
+
45
+ def __init__(self, vocab_size=136):
46
+ self.vocab_size = vocab_size
47
+ self.pad_token_id = 0
48
+ self.unk_token_id = 1
49
+ self.bos_token_id = 2
50
+ self.eos_token_id = 3
51
+
52
+ # Build character set
53
+ chars = list(range(32, 127)) # ASCII printable
54
+ # French accented chars
55
+ french = [192,193,194,195,196,197,199,200,201,202,203,204,205,206,207,
56
+ 210,211,212,213,214,217,218,219,220,224,225,226,227,228,229,
57
+ 231,232,233,234,235,236,237,238,239,242,243,244,245,246,249,
58
+ 250,251,252,253,255]
59
+ chars.extend(french)
60
+
61
+ self.char_to_id = {self.pad_token_id: 0, self.unk_token_id: 1,
62
+ self.bos_token_id: 2, self.eos_token_id: 3}
63
+ for i, c in enumerate(chars[:vocab_size - 4]):
64
+ self.char_to_id[c] = i + 4
65
+
66
+ self.id_to_char = {v: k for k, v in self.char_to_id.items()}
67
+
68
+ def encode(self, text):
69
+ ids = [self.bos_token_id]
70
+ for ch in text:
71
+ code = ord(ch)
72
+ ids.append(self.char_to_id.get(code, self.unk_token_id))
73
+ ids.append(self.eos_token_id)
74
+ return ids
75
+
76
+ def decode(self, ids):
77
+ chars = []
78
+ for i in ids:
79
+ if i in (self.pad_token_id, self.bos_token_id):
80
+ continue
81
+ if i == self.eos_token_id:
82
+ break
83
+ code = self.id_to_char.get(i, 0)
84
+ if code > 0:
85
+ chars.append(chr(code))
86
+ return ''.join(chars)
87
+
88
+ def save(self, path):
89
+ with open(path, 'w', encoding='utf-8') as f:
90
+ json.dump({
91
+ 'vocab_size': self.vocab_size,
92
+ 'char_to_id': {str(k): v for k, v in self.char_to_id.items()},
93
+ }, f, ensure_ascii=False, indent=2)
94
+
95
+ @classmethod
96
+ def load(cls, path):
97
+ with open(path, 'r', encoding='utf-8') as f:
98
+ data = json.load(f)
99
+ tok = cls.__new__(cls)
100
+ tok.vocab_size = data['vocab_size']
101
+ tok.char_to_id = {int(k): v for k, v in data['char_to_id'].items()}
102
+ tok.id_to_char = {v: k for k, v in tok.char_to_id.items()}
103
+ tok.pad_token_id = 0
104
+ tok.unk_token_id = 1
105
+ tok.bos_token_id = 2
106
+ tok.eos_token_id = 3
107
+ return tok
108
+
109
+
110
+ # ──────────────────────────────────────────────────────────
111
+ # Dataset
112
+ # ──────────────────────────────────────────────────────────
113
+ class TokenDataset(Dataset):
114
+ def __init__(self, data_path, seq_len=512):
115
+ tokens = torch.load(data_path, map_location='cpu', weights_only=True)
116
+ if not isinstance(tokens, torch.LongTensor):
117
+ tokens = tokens.long()
118
+ self.tokens = tokens
119
+ self.seq_len = seq_len
120
+
121
+ def __len__(self):
122
+ return max(0, (len(self.tokens) - 1) // self.seq_len)
123
+
124
+ def __getitem__(self, idx):
125
+ start = idx * self.seq_len
126
+ end = start + self.seq_len + 1
127
+ chunk = self.tokens[start:end]
128
+ x = chunk[:-1]
129
+ y = chunk[1:]
130
+ return x, y
131
+
132
+
133
+ # ──────────────────────────────────────────────────────────
134
+ # CogNet-1B Architecture
135
+ # ──────────────────────────────────────────────────────────
136
+ class CognitiveMemory(nn.Module):
137
+ """Hierarchical cognitive memory: Working, Episodic, Semantic."""
138
+
139
+ def __init__(self, dim, working_slots=128, episodic_slots=256, semantic_slots=512):
140
+ super().__init__()
141
+ self.working = nn.Parameter(torch.randn(working_slots, dim) * 0.02)
142
+ self.episodic = nn.Parameter(torch.randn(episodic_slots, dim) * 0.02)
143
+ self.semantic = nn.Parameter(torch.randn(semantic_slots, dim) * 0.02)
144
+ self.gate_w = nn.Linear(dim * 3, 3)
145
+ self.work_proj = nn.Linear(dim, dim)
146
+ self.epis_proj = nn.Linear(dim, dim)
147
+ self.seman_proj = nn.Linear(dim, dim)
148
+
149
+ def forward(self, x):
150
+ B, T, D = x.shape
151
+ x_pooled = x.mean(dim=1) # [B, D]
152
+
153
+ w_att = torch.matmul(x, self.working.t()) # [B, T, ws]
154
+ w_att = F.softmax(w_att / math.sqrt(D), dim=-1)
155
+ w_out = torch.matmul(w_att, self.working) # [B, T, D]
156
+ w_out = self.work_proj(w_out)
157
+
158
+ e_att = torch.matmul(x, self.episodic.t())
159
+ e_att = F.softmax(e_att / math.sqrt(D), dim=-1)
160
+ e_out = torch.matmul(e_att, self.episodic)
161
+ e_out = self.epis_proj(e_out)
162
+
163
+ s_att = torch.matmul(x, self.semantic.t())
164
+ s_att = F.softmax(s_att / math.sqrt(D), dim=-1)
165
+ s_out = torch.matmul(s_att, self.semantic)
166
+ s_out = self.seman_proj(s_out)
167
+
168
+ gate_in = torch.cat([w_out.mean(1), e_out.mean(1), s_out.mean(1)], dim=-1)
169
+ gates = F.softmax(self.gate_w(gate_in), dim=-1) # [B, 3]
170
+
171
+ combined = (gates[:, 0:1, None] * w_out +
172
+ gates[:, 1:2, None] * e_out +
173
+ gates[:, 2:3, None] * s_out)
174
+ return x + combined
175
+
176
+
177
+ class CogNetBlock(nn.Module):
178
+ """CogNet block with cognitive routing across channels."""
179
+
180
+ def __init__(self, dim, n_channels=8, channel_dim=384, ff_dim=8192,
181
+ working_slots=128, episodic_slots=256, semantic_slots=512):
182
+ super().__init__()
183
+ self.n_channels = n_channels
184
+ self.channel_dim = channel_dim
185
+
186
+ # Channel projections
187
+ self.channel_projs = nn.ModuleList([
188
+ nn.Linear(dim, channel_dim) for _ in range(n_channels)
189
+ ])
190
+ self.channel_merge = nn.Linear(n_channels * channel_dim, dim)
191
+ self.norm1 = nn.LayerNorm(dim)
192
+
193
+ # Feed-forward
194
+ self.ff = nn.Sequential(
195
+ nn.Linear(dim, ff_dim),
196
+ nn.GELU(),
197
+ nn.Linear(ff_dim, dim),
198
+ )
199
+ self.norm2 = nn.LayerNorm(dim)
200
+
201
+ # Cognitive memory
202
+ self.cog_mem = CognitiveMemory(dim, working_slots, episodic_slots, semantic_slots)
203
+ self.norm3 = nn.LayerNorm(dim)
204
+
205
+ # Routing gate
206
+ self.router = nn.Linear(dim, n_channels)
207
+
208
+ def forward(self, x):
209
+ B, T, D = x.shape
210
+
211
+ # Channel routing
212
+ route_logits = self.router(x) # [B, T, n_ch]
213
+ route_weights = F.softmax(route_logits, dim=-1) # [B, T, n_ch]
214
+
215
+ # Process each channel
216
+ channel_outs = []
217
+ for i, proj in enumerate(self.channel_projs):
218
+ ch_x = proj(x) # [B, T, ch_dim]
219
+ ch_x = F.gelu(ch_x)
220
+ # Apply route weight
221
+ w = route_weights[:, :, i:i+1] # [B, T, 1]
222
+ channel_outs.append(ch_x * w)
223
+
224
+ merged = torch.cat(channel_outs, dim=-1) # [B, T, n_ch * ch_dim]
225
+ merged = self.channel_merge(merged) # [B, T, D]
226
+
227
+ # Residual + norm
228
+ x = self.norm1(x + merged)
229
+
230
+ # FFN
231
+ x = self.norm2(x + self.ff(x))
232
+
233
+ # Cognitive memory
234
+ x = self.norm3(self.cog_mem(x))
235
+
236
+ return x
237
+
238
+
239
+ class CogNet1B(nn.Module):
240
+ """CogNet-1B: ~1.06B parameter cognitive language model."""
241
+
242
+ def __init__(self, vocab_size=136, hidden_dim=2048, n_blocks=16,
243
+ n_channels=8, channel_dim=384, ff_dim=8192,
244
+ working_slots=128, episodic_slots=256, semantic_slots=512,
245
+ seq_len=512):
246
+ super().__init__()
247
+ self.vocab_size = vocab_size
248
+ self.hidden_dim = hidden_dim
249
+ self.seq_len = seq_len
250
+
251
+ self.token_emb = nn.Embedding(vocab_size, hidden_dim)
252
+ self.pos_emb = nn.Embedding(seq_len, hidden_dim)
253
+
254
+ self.blocks = nn.ModuleList([
255
+ CogNetBlock(hidden_dim, n_channels, channel_dim, ff_dim,
256
+ working_slots, episodic_slots, semantic_slots)
257
+ for _ in range(n_blocks)
258
+ ])
259
+
260
+ self.final_norm = nn.LayerNorm(hidden_dim)
261
+ self.head = nn.Linear(hidden_dim, vocab_size, bias=False)
262
+
263
+ # Weight tying
264
+ self.head.weight = self.token_emb.weight
265
+
266
+ self._init_weights()
267
+
268
+ def _init_weights(self):
269
+ for module in self.modules():
270
+ if isinstance(module, nn.Linear):
271
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
272
+ if module.bias is not None:
273
+ nn.init.zeros_(module.bias)
274
+ elif isinstance(module, nn.Embedding):
275
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
276
+
277
+ def forward(self, input_ids, targets=None):
278
+ B, T = input_ids.shape
279
+ assert T <= self.seq_len, f"Sequence length {T} > max {self.seq_len}"
280
+
281
+ pos = torch.arange(T, device=input_ids.device).unsqueeze(0)
282
+ x = self.token_emb(input_ids) + self.pos_emb(pos)
283
+
284
+ for block in self.blocks:
285
+ x = block(x)
286
+
287
+ x = self.final_norm(x)
288
+ logits = self.head(x)
289
+
290
+ loss = None
291
+ if targets is not None:
292
+ loss = F.cross_entropy(
293
+ logits.view(-1, self.vocab_size),
294
+ targets.view(-1),
295
+ ignore_index=0,
296
+ )
297
+
298
+ return logits, loss
299
+
300
+ @torch.no_grad()
301
+ def generate(self, input_ids, max_new_tokens=200, temperature=0.8, top_k=50):
302
+ self.eval()
303
+ for _ in range(max_new_tokens):
304
+ idx_cond = input_ids[:, -self.seq_len:]
305
+ logits, _ = self(idx_cond)
306
+ logits = logits[:, -1, :] / temperature
307
+ if top_k > 0:
308
+ v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
309
+ logits[logits < v[:, [-1]]] = float('-inf')
310
+ probs = F.softmax(logits, dim=-1)
311
+ next_id = torch.multinomial(probs, num_samples=1)
312
+ input_ids = torch.cat([input_ids, next_id], dim=1)
313
+ return input_ids
314
+
315
+
316
+ # ──────────────────────────────────────────────────────────
317
+ # Learning Rate Schedule
318
+ # ──────────────────────────────────────────────────────────
319
+ def get_cosine_lr(step, warmup_steps, max_steps, max_lr, min_lr):
320
+ if step < warmup_steps:
321
+ return max_lr * step / max(1, warmup_steps)
322
+ if step >= max_steps:
323
+ return min_lr
324
+ progress = (step - warmup_steps) / max(1, max_steps - warmup_steps)
325
+ return min_lr + 0.5 * (max_lr - min_lr) * (1 + math.cos(math.pi * progress))
326
+
327
+
328
+ # ──────────────────────────────────────────────────────────
329
+ # Main Training Loop
330
+ # ──────────────────────────────────────────────────────────
331
+ def main():
332
+ parser = argparse.ArgumentParser()
333
+ parser.add_argument('--batch-size', type=int, default=8, help='Batch size (RTX 5090 32GB BF16)')
334
+ parser.add_argument('--seq-len', type=int, default=512)
335
+ parser.add_argument('--max-steps', type=int, default=200000)
336
+ parser.add_argument('--warmup-steps', type=int, default=4000)
337
+ parser.add_argument('--max-lr', type=float, default=1e-4)
338
+ parser.add_argument('--min-lr', type=float, default=1e-5)
339
+ parser.add_argument('--weight-decay', type=float, default=0.1)
340
+ parser.add_argument('--grad-clip', type=float, default=1.0)
341
+ parser.add_argument('--save-every', type=int, default=5000)
342
+ parser.add_argument('--eval-every', type=int, default=1000)
343
+ parser.add_argument('--log-every', type=int, default=100)
344
+ parser.add_argument('--data-path', type=str, default='/root/CogNet/data_1b/aicl_10x.pt')
345
+ parser.add_argument('--tokenizer-path', type=str, default='/root/CogNet/tokenizer_v3.json')
346
+ parser.add_argument('--ckpt-dir', type=str, default='/root/CogNet/checkpoints_1b')
347
+ parser.add_argument('--resume', type=str, default=None, help='Path to checkpoint to resume from')
348
+ parser.add_argument('--bf16', action='store_true', default=True, help='Use BF16 mixed precision')
349
+ parser.add_argument('--compile', action='store_true', default=False, help='torch.compile the model')
350
+ args = parser.parse_args()
351
+
352
+ # Setup
353
+ os.makedirs(args.ckpt_dir, exist_ok=True)
354
+ device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
355
+
356
+ log.info(f'=== CogNet-1B Training (241GB VRAM Optimized) ===')
357
+ log.info(f'Device: {device}')
358
+
359
+ if torch.cuda.is_available():
360
+ for i in range(torch.cuda.device_count()):
361
+ props = torch.cuda.get_device_properties(i)
362
+ total_gb = props.total_memory / 1e9
363
+ log.info(f'GPU {i}: {props.name} — {total_gb:.1f} GB VRAM')
364
+
365
+ # Tokenizer
366
+ if os.path.exists(args.tokenizer_path):
367
+ tokenizer = CharTokenizer.load(args.tokenizer_path)
368
+ log.info(f'Loaded tokenizer from {args.tokenizer_path} (vocab={tokenizer.vocab_size})')
369
+ else:
370
+ tokenizer = CharTokenizer()
371
+ tokenizer.save(args.tokenizer_path)
372
+ log.info(f'Created and saved tokenizer (vocab={tokenizer.vocab_size})')
373
+
374
+ # Dataset
375
+ log.info(f'Loading data from {args.data_path}...')
376
+ dataset = TokenDataset(args.data_path, args.seq_len)
377
+ log.info(f'Dataset: {len(dataset):,} sequences of length {args.seq_len}')
378
+
379
+ dataloader = DataLoader(
380
+ dataset,
381
+ batch_size=args.batch_size,
382
+ shuffle=True,
383
+ num_workers=4,
384
+ pin_memory=True,
385
+ drop_last=True,
386
+ )
387
+
388
+ # Model
389
+ log.info('Building CogNet-1B model...')
390
+ model = CogNet1B(
391
+ vocab_size=tokenizer.vocab_size,
392
+ hidden_dim=2048,
393
+ n_blocks=16,
394
+ n_channels=8,
395
+ channel_dim=384,
396
+ ff_dim=8192,
397
+ working_slots=128,
398
+ episodic_slots=256,
399
+ semantic_slots=512,
400
+ seq_len=args.seq_len,
401
+ ).to(device)
402
+
403
+ # Count parameters
404
+ total_params = sum(p.numel() for p in model.parameters())
405
+ trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
406
+ log.info(f'Total parameters: {total_params:,} ({total_params/1e9:.2f}B)')
407
+ log.info(f'Trainable parameters: {trainable_params:,}')
408
+
409
+ # Optional: torch.compile
410
+ if args.compile:
411
+ log.info('Compiling model with torch.compile...')
412
+ model = torch.compile(model)
413
+
414
+ # Optimizer — FULL AdamW, no compromises with 241GB VRAM!
415
+ optimizer = torch.optim.AdamW(
416
+ model.parameters(),
417
+ lr=args.max_lr,
418
+ betas=(0.9, 0.95),
419
+ eps=1e-8,
420
+ weight_decay=args.weight_decay,
421
+ fused=True, # Fused CUDA kernel for speed
422
+ )
423
+
424
+ # Mixed precision scaler
425
+ use_bf16 = args.bf16 and torch.cuda.is_bf16_supported()
426
+ scaler = None
427
+ if not use_bf16:
428
+ scaler = torch.amp.GradScaler('cuda')
429
+ log.info(f'Mixed precision: {"BF16" if use_bf16 else "FP16 with GradScaler"}')
430
+
431
+ # Resume from checkpoint
432
+ start_step = 0
433
+ best_loss = float('inf')
434
+ if args.resume and os.path.exists(args.resume):
435
+ log.info(f'Resuming from {args.resume}...')
436
+ ckpt = torch.load(args.resume, map_location=device, weights_only=False)
437
+ model.load_state_dict(ckpt['model_state_dict'])
438
+ optimizer.load_state_dict(ckpt['optimizer_state_dict'])
439
+ start_step = ckpt.get('step', 0)
440
+ best_loss = ckpt.get('best_loss', float('inf'))
441
+ log.info(f'Resumed at step {start_step}, best_loss={best_loss:.4f}')
442
+
443
+ # Also check for latest checkpoint in ckpt_dir
444
+ if start_step == 0:
445
+ latest = os.path.join(args.ckpt_dir, 'cognet_1b_latest.pt')
446
+ if os.path.exists(latest):
447
+ log.info(f'Found latest checkpoint: {latest}')
448
+ ckpt = torch.load(latest, map_location=device, weights_only=False)
449
+ model.load_state_dict(ckpt['model_state_dict'])
450
+ optimizer.load_state_dict(ckpt['optimizer_state_dict'])
451
+ start_step = ckpt.get('step', 0)
452
+ best_loss = ckpt.get('best_loss', float('inf'))
453
+ log.info(f'Resumed at step {start_step}, best_loss={best_loss:.4f}')
454
+
455
+ # Training loop
456
+ log.info(f'Starting training from step {start_step} to {args.max_steps}')
457
+ log.info(f'Batch size: {args.batch_size}, Seq len: {args.seq_len}')
458
+ log.info(f'LR range: {args.min_lr} -> {args.max_lr} -> {args.min_lr}')
459
+ log.info(f'Warmup: {args.warmup_steps} steps')
460
+
461
+ model.train()
462
+ data_iter = iter(dataloader)
463
+ t0 = time.time()
464
+
465
+ for step in range(start_step, args.max_steps):
466
+ # Get batch
467
+ try:
468
+ batch = next(data_iter)
469
+ except StopIteration:
470
+ data_iter = iter(dataloader)
471
+ batch = next(data_iter)
472
+
473
+ x, y = batch
474
+ x = x.to(device, non_blocking=True)
475
+ y = y.to(device, non_blocking=True)
476
+
477
+ # Learning rate
478
+ lr = get_cosine_lr(step, args.warmup_steps, args.max_steps, args.max_lr, args.min_lr)
479
+ for param_group in optimizer.param_groups:
480
+ param_group['lr'] = lr
481
+
482
+ # Forward + backward
483
+ optimizer.zero_grad(set_to_none=True)
484
+
485
+ if use_bf16:
486
+ with torch.amp.autocast('cuda', dtype=torch.bfloat16):
487
+ logits, loss = model(x, y)
488
+ loss.backward()
489
+ else:
490
+ with torch.amp.autocast('cuda', dtype=torch.float16):
491
+ logits, loss = model(x, y)
492
+ scaler.scale(loss).backward()
493
+
494
+ # Gradient clipping
495
+ if use_bf16:
496
+ grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
497
+ else:
498
+ scaler.unscale_(optimizer)
499
+ grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip)
500
+
501
+ # Optimizer step
502
+ if use_bf16:
503
+ optimizer.step()
504
+ else:
505
+ scaler.step(optimizer)
506
+ scaler.update()
507
+
508
+ # Logging
509
+ if step % args.log_every == 0:
510
+ elapsed = time.time() - t0
511
+ steps_per_sec = args.log_every / max(elapsed, 0.001)
512
+ tokens_per_sec = steps_per_sec * args.batch_size * args.seq_len
513
+
514
+ vram_used = torch.cuda.memory_allocated() / 1e9
515
+ vram_reserved = torch.cuda.memory_reserved() / 1e9
516
+
517
+ log.info(
518
+ f'Step {step:>7d}/{args.max_steps} | '
519
+ f'Loss: {loss.item():.4f} | '
520
+ f'LR: {lr:.2e} | '
521
+ f'Grad: {grad_norm:.2f} | '
522
+ f'VRAM: {vram_used:.1f}/{vram_reserved:.1f} GB | '
523
+ f'Speed: {tokens_per_sec:.0f} tok/s | '
524
+ f'{steps_per_sec:.1f} steps/s'
525
+ )
526
+ t0 = time.time()
527
+
528
+ # Evaluation / sample generation
529
+ if step > 0 and step % args.eval_every == 0:
530
+ model.eval()
531
+ with torch.no_grad():
532
+ prompt = torch.tensor([[tokenizer.bos_token_id]], device=device)
533
+ sample_ids = model.generate(prompt, max_new_tokens=100, temperature=0.8, top_k=50)
534
+ sample_text = tokenizer.decode(sample_ids[0].tolist())
535
+ log.info(f'--- Sample at step {step} ---')
536
+ log.info(sample_text[:300])
537
+ log.info(f'--- End sample ---')
538
+ model.train()
539
+
540
+ # Save checkpoint
541
+ if step > 0 and step % args.save_every == 0:
542
+ ckpt_path = os.path.join(args.ckpt_dir, f'cognet_1b_step_{step}.pt')
543
+ latest_path = os.path.join(args.ckpt_dir, 'cognet_1b_latest.pt')
544
+
545
+ is_best = loss.item() < best_loss
546
+ if is_best:
547
+ best_loss = loss.item()
548
+ best_path = os.path.join(args.ckpt_dir, 'cognet_1b_best.pt')
549
+
550
+ save_dict = {
551
+ 'step': step,
552
+ 'model_state_dict': model.state_dict(),
553
+ 'optimizer_state_dict': optimizer.state_dict(),
554
+ 'loss': loss.item(),
555
+ 'best_loss': best_loss,
556
+ 'config': {
557
+ 'vocab_size': tokenizer.vocab_size,
558
+ 'hidden_dim': 2048,
559
+ 'n_blocks': 16,
560
+ 'n_channels': 8,
561
+ 'channel_dim': 384,
562
+ 'ff_dim': 8192,
563
+ 'working_slots': 128,
564
+ 'episodic_slots': 256,
565
+ 'semantic_slots': 512,
566
+ 'seq_len': args.seq_len,
567
+ },
568
+ }
569
+
570
+ torch.save(save_dict, ckpt_path)
571
+ torch.save(save_dict, latest_path)
572
+ if is_best:
573
+ torch.save(save_dict, best_path)
574
+
575
+ log.info(f'Saved checkpoint: {ckpt_path} (loss={loss.item():.4f}, best={best_loss:.4f})')
576
+
577
+ # Final save
578
+ final_path = os.path.join(args.ckpt_dir, 'cognet_1b_final.pt')
579
+ save_dict = {
580
+ 'step': args.max_steps,
581
+ 'model_state_dict': model.state_dict(),
582
+ 'optimizer_state_dict': optimizer.state_dict(),
583
+ 'loss': loss.item(),
584
+ 'best_loss': best_loss,
585
+ 'config': {
586
+ 'vocab_size': tokenizer.vocab_size,
587
+ 'hidden_dim': 2048,
588
+ 'n_blocks': 16,
589
+ 'n_channels': 8,
590
+ 'channel_dim': 384,
591
+ 'ff_dim': 8192,
592
+ 'working_slots': 128,
593
+ 'episodic_slots': 256,
594
+ 'semantic_slots': 512,
595
+ 'seq_len': args.seq_len,
596
+ },
597
+ }
598
+ torch.save(save_dict, final_path)
599
+ log.info(f'Training complete! Final model saved to {final_path}')
600
+ log.info(f'Final loss: {loss.item():.4f}, Best loss: {best_loss:.4f}')
601
+
602
+
603
+ if __name__ == '__main__':
604
+ main()